1use std::{
2 collections::{BTreeMap, BTreeSet},
3 error::Error,
4 fmt,
5 sync::atomic::Ordering,
6};
7
8use sim_expr_tree_core::MountEpoch;
9use sim_incremental_core::{
10 ContinuationToken, GraphSnapshot, Observation, ObservationKind, Revision, SnapshotNode,
11 ValueFingerprint,
12};
13use sim_kernel::{Cx, Expr, Symbol, Table, Value};
14
15use super::*;
16
17mod codec;
18use codec::{DecodeError, decode_persisted, encode_persisted, restore_value};
19mod identity;
20use identity::state_identities;
21
22pub const GRAPH_SCHEMA_VERSION: u64 = 1;
24
25pub const DERIVED_SNAPSHOT_KEY: &str = "expr-tree-graph";
27
28const MAX_PERSISTED_GRAPH_NODES: usize = 100_000;
29const MAX_PERSISTED_GRAPH_EDGES: usize = 1_000_000;
30
31pub struct DerivedTableAdapter<'a> {
37 table: &'a dyn Table,
38 cx: &'a mut Cx,
39 key: Symbol,
40}
41
42impl<'a> DerivedTableAdapter<'a> {
43 pub fn new(table: &'a dyn Table, cx: &'a mut Cx) -> Self {
45 Self {
46 table,
47 cx,
48 key: Symbol::new(DERIVED_SNAPSHOT_KEY),
49 }
50 }
51
52 pub fn with_key(table: &'a dyn Table, cx: &'a mut Cx, key: Symbol) -> Self {
54 Self { table, cx, key }
55 }
56
57 fn load(&mut self) -> Result<Option<Expr>, DerivedSnapshotError> {
58 if !self
59 .table
60 .has(self.cx, self.key.clone())
61 .map_err(backend_error)?
62 {
63 return Ok(None);
64 }
65 let value = self
66 .table
67 .get(self.cx, self.key.clone())
68 .map_err(backend_error)?;
69 value
70 .object()
71 .as_expr(self.cx)
72 .map(Some)
73 .map_err(backend_error)
74 }
75
76 fn store(&mut self, expr: Expr) -> Result<(), DerivedSnapshotError> {
77 let value = self.cx.factory().expr(expr).map_err(backend_error)?;
78 self.table
79 .set(self.cx, self.key.clone(), value)
80 .map_err(backend_error)
81 }
82
83 pub fn delete(&mut self) -> Result<(), DerivedSnapshotError> {
85 self.table
86 .del(self.cx, self.key.clone())
87 .map(|_| ())
88 .map_err(backend_error)
89 }
90}
91
92#[derive(Clone, Copy, Debug, Default, Eq, PartialEq)]
94pub struct DerivedPersistReport {
95 pub nodes: usize,
97 pub reverse_edges: usize,
99 pub receipts: usize,
101 pub pending_continuations: usize,
103}
104
105#[derive(Clone, Copy, Debug, Eq, PartialEq)]
107pub enum DerivedRestoreDisposition {
108 Rehydrated,
110 RebuiltMissing,
112 RebuiltCorrupt,
114 RebuiltIncompatible,
116 RebuiltGenerationMismatch,
118}
119
120#[derive(Clone, Copy, Debug, Eq, PartialEq)]
122pub struct DerivedRestoreReport {
123 pub disposition: DerivedRestoreDisposition,
125 pub restored_nodes: usize,
127 pub recovered_dirty: usize,
129 pub restored_receipts: usize,
131 pub pending_continuations: usize,
133}
134
135impl DerivedRestoreReport {
136 fn rebuilt(disposition: DerivedRestoreDisposition) -> Self {
137 Self {
138 disposition,
139 restored_nodes: 0,
140 recovered_dirty: 0,
141 restored_receipts: 0,
142 pending_continuations: 0,
143 }
144 }
145}
146
147#[derive(Clone, Debug, Eq, PartialEq)]
149pub enum DerivedSnapshotError {
150 Backend {
152 message: String,
154 },
155 Graph {
157 message: String,
159 },
160}
161
162impl fmt::Display for DerivedSnapshotError {
163 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
164 match self {
165 Self::Backend { message } => write!(f, "derived Table operation failed: {message}"),
166 Self::Graph { message } => write!(f, "cannot snapshot derived graph: {message}"),
167 }
168 }
169}
170
171impl Error for DerivedSnapshotError {}
172
173struct PersistedCalc {
174 schema: u64,
175 source_generation: u64,
176 control_generation: u64,
177 source_identity: u64,
178 control_identity: u64,
179 graph: GraphSnapshot<CalcQuery, MemoValue>,
180 reverse: BTreeMap<CalcQuery, BTreeSet<CalcQuery>>,
181 receipts: BTreeMap<String, CalcReceipt>,
182 queue: AutomaticQueueSnapshot,
183 next_request_id: u64,
184 next_logical_tick: u64,
185 next_volatile: u64,
186 refresh_samples: BTreeMap<String, BackendRefreshSample>,
187 last_good: BTreeMap<String, Expr>,
188}
189
190impl ExprTreeCalc {
191 pub fn persist_derived(
194 &mut self,
195 derived: &mut DerivedTableAdapter<'_>,
196 ) -> Result<DerivedPersistReport, DerivedSnapshotError> {
197 let roots = self
198 .state
199 .read()
200 .expect("calc state poisoned")
201 .cells
202 .keys()
203 .cloned()
204 .map(CalcQuery::Cell)
205 .collect::<Vec<_>>();
206 let graph = self
207 .engine
208 .snapshot(
209 roots,
210 SnapshotBudgets::new(MAX_PERSISTED_GRAPH_NODES, MAX_PERSISTED_GRAPH_EDGES),
211 )
212 .map_err(|error| DerivedSnapshotError::Graph {
213 message: error.to_string(),
214 })?;
215 let reverse = derive_reverse(&graph);
216 let (source_identity, control_identity) =
217 state_identities(&self.state, &self.context_factory)?;
218 let (source_generation, control_generation, receipts, next_logical_tick, last_good_values) = {
219 let state = self.state.read().expect("calc state poisoned");
220 (
221 state.source_generation,
222 state.control_generation,
223 state.receipts.clone(),
224 state.next_logical_tick,
225 state.last_good.clone(),
226 )
227 };
228 let mut value_cx = (self.context_factory)();
229 let last_good = last_good_values
230 .into_iter()
231 .filter_map(|(cell, value)| {
232 value
233 .object()
234 .as_expr(&mut value_cx)
235 .ok()
236 .map(|expr| (cell, expr))
237 })
238 .collect();
239 let queue = self.automatic_queue_snapshot();
240 let pending_continuations = queue
241 .entries
242 .iter()
243 .filter(|entry| entry.incremental_continuation.is_some())
244 .count();
245 let persisted = PersistedCalc {
246 schema: GRAPH_SCHEMA_VERSION,
247 source_generation,
248 control_generation,
249 source_identity,
250 control_identity,
251 graph,
252 reverse,
253 receipts,
254 queue,
255 next_request_id: self.next_request_id,
256 next_logical_tick,
257 next_volatile: self.next_volatile.load(Ordering::Acquire),
258 refresh_samples: self.refresh_samples.clone(),
259 last_good,
260 };
261 let report = DerivedPersistReport {
262 nodes: persisted.graph.nodes.len(),
263 reverse_edges: persisted.reverse.values().map(BTreeSet::len).sum(),
264 receipts: persisted.receipts.len(),
265 pending_continuations,
266 };
267 let mut cx = (self.context_factory)();
268 let encoded = encode_persisted(&persisted, &mut cx)?;
269 derived.store(encoded)?;
270 Ok(report)
271 }
272
273 pub fn restore_derived(
279 &mut self,
280 derived: &mut DerivedTableAdapter<'_>,
281 ) -> Result<DerivedRestoreReport, DerivedSnapshotError> {
282 let Some(encoded) = derived.load()? else {
283 self.rebuild_derived_state();
284 return Ok(DerivedRestoreReport::rebuilt(
285 DerivedRestoreDisposition::RebuiltMissing,
286 ));
287 };
288 let mut cx = (self.context_factory)();
289 let persisted = match decode_persisted(&encoded, &mut cx) {
290 Ok(snapshot) => snapshot,
291 Err(DecodeError::Incompatible) => {
292 self.rebuild_derived_state();
293 derived.delete()?;
294 return Ok(DerivedRestoreReport::rebuilt(
295 DerivedRestoreDisposition::RebuiltIncompatible,
296 ));
297 }
298 Err(DecodeError::Corrupt(_)) => {
299 self.rebuild_derived_state();
300 derived.delete()?;
301 return Ok(DerivedRestoreReport::rebuilt(
302 DerivedRestoreDisposition::RebuiltCorrupt,
303 ));
304 }
305 };
306 let (source_identity, control_identity) =
307 state_identities(&self.state, &self.context_factory)?;
308 let generations_match = {
309 let state = self.state.read().expect("calc state poisoned");
310 persisted.source_generation == state.source_generation
311 && persisted.control_generation == state.control_generation
312 };
313 if !generations_match
314 || persisted.source_identity != source_identity
315 || persisted.control_identity != control_identity
316 {
317 self.rebuild_derived_state();
318 derived.delete()?;
319 return Ok(DerivedRestoreReport::rebuilt(
320 DerivedRestoreDisposition::RebuiltGenerationMismatch,
321 ));
322 }
323 if persisted.reverse != derive_reverse(&persisted.graph) {
324 self.rebuild_derived_state();
325 derived.delete()?;
326 return Ok(DerivedRestoreReport::rebuilt(
327 DerivedRestoreDisposition::RebuiltCorrupt,
328 ));
329 }
330
331 let current = current_from_graph(&persisted.graph);
332 let restore = match self.engine.restore_snapshot(persisted.graph) {
333 Ok(report) => report,
334 Err(_) => {
335 self.rebuild_derived_state();
336 derived.delete()?;
337 return Ok(DerivedRestoreReport::rebuilt(
338 DerivedRestoreDisposition::RebuiltCorrupt,
339 ));
340 }
341 };
342 let pending_tokens = persisted
343 .queue
344 .entries
345 .iter()
346 .filter_map(|entry| entry.incremental_continuation)
347 .collect::<BTreeSet<_>>();
348 if self.restore_automatic_queue(persisted.queue).is_err() {
349 self.rebuild_derived_state();
350 derived.delete()?;
351 return Ok(DerivedRestoreReport::rebuilt(
352 DerivedRestoreDisposition::RebuiltCorrupt,
353 ));
354 }
355 let restored_receipts = persisted.receipts.len();
356 {
357 let mut state = self.state.write().expect("calc state poisoned");
358 state.receipts = persisted.receipts;
359 state.next_logical_tick = persisted.next_logical_tick.max(1);
360 state.current = current.0;
361 state.failed_cells = current.1;
362 state.volatile = current.2;
363 state.last_good = restore_last_good(persisted.last_good, &mut cx);
364 }
365 self.next_request_id = self.next_request_id.max(persisted.next_request_id);
366 self.next_volatile
367 .store(persisted.next_volatile.max(1), Ordering::Release);
368 self.refresh_samples = persisted.refresh_samples;
369 self.restored_continuations = pending_tokens;
370 Ok(DerivedRestoreReport {
371 disposition: DerivedRestoreDisposition::Rehydrated,
372 restored_nodes: restore.nodes,
373 recovered_dirty: restore.recovered_dirty,
374 restored_receipts,
375 pending_continuations: self.restored_continuations.len(),
376 })
377 }
378
379 fn rebuild_derived_state(&mut self) {
380 self.engine = IncrementalEngine::new();
381 let (cells, mount_epochs) = {
382 let state = self.state.read().expect("calc state poisoned");
383 (
384 state.cells.keys().cloned().collect::<Vec<_>>(),
385 state
386 .mounts
387 .iter()
388 .map(|(path, mount)| (path.clone(), mount.epoch))
389 .collect::<Vec<_>>(),
390 )
391 };
392 {
393 let mut state = self.state.write().expect("calc state poisoned");
394 state.active_request = None;
395 state.attempts.clear();
396 state.receipts.clear();
397 state.next_logical_tick = 1;
398 state.current.clear();
399 state.last_good.clear();
400 state.volatile.clear();
401 state.failed_cells.clear();
402 }
403 self.next_request_id = 1;
404 self.automatic_queue.clear();
405 self.automatic_generation = 1;
406 self.next_queue_sequence = 1;
407 self.next_volatile.store(1, Ordering::Release);
408 self.restored_continuations.clear();
409 self.refresh_samples = mount_epochs
410 .into_iter()
411 .map(|(path, epoch)| (path, BackendRefreshSample::new(epoch)))
412 .collect();
413 for cell in cells {
414 self.register_cell_query(cell);
415 }
416 self.schedule_dirty_automatic();
417 }
418}
419
420fn backend_error(error: impl fmt::Display) -> DerivedSnapshotError {
421 DerivedSnapshotError::Backend {
422 message: error.to_string(),
423 }
424}
425
426fn derive_reverse(
427 graph: &GraphSnapshot<CalcQuery, MemoValue>,
428) -> BTreeMap<CalcQuery, BTreeSet<CalcQuery>> {
429 let mut reverse = BTreeMap::new();
430 for node in &graph.nodes {
431 for observation in &node.dependencies {
432 reverse
433 .entry(observation.key().clone())
434 .or_insert_with(BTreeSet::new)
435 .insert(node.key.clone());
436 }
437 }
438 reverse
439}
440
441type CurrentState = (
442 BTreeMap<String, Result<Value, CalcError>>,
443 BTreeSet<String>,
444 BTreeSet<String>,
445);
446
447fn current_from_graph(graph: &GraphSnapshot<CalcQuery, MemoValue>) -> CurrentState {
448 let mut current = BTreeMap::new();
449 let mut failed = BTreeSet::new();
450 let mut volatile = BTreeSet::new();
451 for node in &graph.nodes {
452 let CalcQuery::Cell(cell) = &node.key else {
453 continue;
454 };
455 let Some(memo) = &node.value else {
456 continue;
457 };
458 match &memo.outcome {
459 MemoOutcome::Value(value) => {
460 current.insert(cell.clone(), Ok(value.clone()));
461 if memo.is_volatile() {
462 volatile.insert(cell.clone());
463 }
464 }
465 MemoOutcome::Failure(failure) => {
466 current.insert(cell.clone(), Err(CalcError::Cell(failure.clone())));
467 failed.insert(cell.clone());
468 }
469 }
470 }
471 (current, failed, volatile)
472}
473
474fn restore_last_good(values: BTreeMap<String, Expr>, cx: &mut Cx) -> BTreeMap<String, Value> {
475 values
476 .into_iter()
477 .filter_map(|(cell, expr)| restore_value(cx, expr).ok().map(|(value, _)| (cell, value)))
478 .collect()
479}