Skip to main content

sim_incremental_core/
engine.rs

1//! Incremental query engine and query frames.
2
3use std::collections::{BTreeMap, BTreeSet};
4
5use crate::{
6    BudgetKind, ContinuationToken, FingerprintValue, IncrementalError, Observation,
7    ObservationKind, Query, QueryBudgets, QueryResult, Revision, ValueFingerprint,
8    state::{Node, RunState},
9};
10
11/// A dependency-light incremental query engine.
12pub struct IncrementalEngine<K, V> {
13    pub(crate) queries: BTreeMap<K, Query<K, V>>,
14    pub(crate) nodes: BTreeMap<K, Node<K, V>>,
15    pub(crate) reverse: BTreeMap<K, BTreeSet<K>>,
16    pub(crate) source_revisions: BTreeMap<K, Revision>,
17    pub(crate) continuations: BTreeMap<ContinuationToken, K>,
18    pub(crate) next_revision: u64,
19    pub(crate) next_token: u64,
20}
21
22impl<K, V> Default for IncrementalEngine<K, V>
23where
24    K: Ord + Clone,
25{
26    fn default() -> Self {
27        Self::new()
28    }
29}
30
31impl<K, V> IncrementalEngine<K, V>
32where
33    K: Ord + Clone,
34{
35    /// Creates an empty incremental query engine.
36    #[must_use]
37    pub fn new() -> Self {
38        Self {
39            queries: BTreeMap::new(),
40            nodes: BTreeMap::new(),
41            reverse: BTreeMap::new(),
42            source_revisions: BTreeMap::new(),
43            continuations: BTreeMap::new(),
44            next_revision: 1,
45            next_token: 1,
46        }
47    }
48
49    /// Registers or replaces a query callback for `key`.
50    pub fn register_query(&mut self, key: K, query: Query<K, V>) {
51        self.queries.insert(key.clone(), query);
52        self.nodes.entry(key.clone()).or_default().dirty = true;
53        self.mark_dirty_cascade(&key);
54    }
55
56    /// Registers a query callback function for `key`.
57    pub fn register_fn<F>(&mut self, key: K, query: F)
58    where
59        F: for<'a> Fn(&K, &mut QueryFrame<'a, K, V>) -> QueryResult<K, V> + Send + Sync + 'static,
60    {
61        self.register_query(key, Query::new(query));
62    }
63
64    /// Removes a query and marks its dependents dirty.
65    pub fn remove_query(&mut self, key: &K) -> bool {
66        let removed = self.queries.remove(key).is_some();
67        if removed {
68            self.mark_dirty_cascade(key);
69            self.detach_node(key);
70        }
71        removed
72    }
73
74    /// Advances an external observation stamp and invalidates reverse
75    /// dependents.
76    pub fn invalidate(&mut self, key: &K) -> Revision {
77        let revision = self.alloc_revision();
78        self.source_revisions.insert(key.clone(), revision);
79        self.mark_dirty_cascade(key);
80        revision
81    }
82
83    /// Returns the current external observation revision for `key`.
84    #[must_use]
85    pub fn source_revision(&self, key: &K) -> Revision {
86        self.source_revisions
87            .get(key)
88            .copied()
89            .unwrap_or(Revision::ZERO)
90    }
91
92    /// Returns dirty memo keys in deterministic order.
93    #[must_use]
94    pub fn dirty_keys(&self) -> Vec<K> {
95        self.nodes
96            .iter()
97            .filter(|(_, node)| node.dirty)
98            .map(|(key, _)| key.clone())
99            .collect()
100    }
101
102    /// Returns the current memo revision for `key`, when present.
103    #[must_use]
104    pub fn memo_revision(&self, key: &K) -> Option<Revision> {
105        self.nodes.get(key).map(|node| node.revision)
106    }
107
108    /// Returns the current memo fingerprint for `key`, when present.
109    #[must_use]
110    pub fn memo_fingerprint(&self, key: &K) -> Option<ValueFingerprint> {
111        self.nodes.get(key).and_then(|node| node.fingerprint)
112    }
113
114    pub(crate) fn alloc_revision(&mut self) -> Revision {
115        let revision = Revision::new(self.next_revision);
116        self.next_revision += 1;
117        revision
118    }
119
120    pub(crate) fn alloc_continuation(&mut self, root: K) -> ContinuationToken {
121        let token = ContinuationToken::new(self.next_token);
122        self.next_token += 1;
123        self.continuations.insert(token, root);
124        token
125    }
126}
127
128impl<K, V> IncrementalEngine<K, V>
129where
130    K: Ord + Clone,
131    V: Clone + FingerprintValue,
132{
133    /// Verifies a root query using unbounded budgets.
134    pub fn verify(&mut self, key: K) -> QueryResult<K, V> {
135        self.verify_with_budgets(key, QueryBudgets::default())
136    }
137
138    /// Verifies a root query using explicit budgets.
139    pub fn verify_with_budgets(&mut self, key: K, budgets: QueryBudgets) -> QueryResult<K, V> {
140        let mut run = RunState::new(key.clone(), budgets);
141        self.evaluate(key, &mut run)
142    }
143
144    /// Resumes the root query represented by a continuation token.
145    pub fn resume(&mut self, token: ContinuationToken, budgets: QueryBudgets) -> QueryResult<K, V> {
146        let root = self
147            .continuations
148            .get(&token)
149            .cloned()
150            .ok_or(IncrementalError::UnknownContinuation { token })?;
151        let value = self.verify_with_budgets(root, budgets)?;
152        self.continuations.remove(&token);
153        Ok(value)
154    }
155
156    /// Verifies roots in stable key order using unbounded budgets.
157    pub fn verify_many<I>(&mut self, keys: I) -> QueryResult<K, Vec<(K, V)>>
158    where
159        I: IntoIterator<Item = K>,
160    {
161        self.verify_many_with_budgets(keys, QueryBudgets::default())
162    }
163
164    /// Verifies roots in stable key order using explicit budgets for each root.
165    pub fn verify_many_with_budgets<I>(
166        &mut self,
167        keys: I,
168        budgets: QueryBudgets,
169    ) -> QueryResult<K, Vec<(K, V)>>
170    where
171        I: IntoIterator<Item = K>,
172    {
173        let ordered = keys.into_iter().collect::<BTreeSet<_>>();
174        let mut out = Vec::new();
175        for key in ordered {
176            let value = self.verify_with_budgets(key.clone(), budgets)?;
177            out.push((key, value));
178        }
179        Ok(out)
180    }
181
182    pub(crate) fn evaluate(&mut self, key: K, run: &mut RunState<K>) -> QueryResult<K, V> {
183        self.check_cancelled(run)?;
184        if !self.queries.contains_key(&key) {
185            return Err(IncrementalError::UnknownQuery { key });
186        }
187        if let Some(index) = run.stack.iter().position(|item| item == &key) {
188            let mut path = run.stack[index..].to_vec();
189            path.push(key);
190            return Err(IncrementalError::Cycle { path });
191        }
192        self.charge_depth(run, &key)?;
193
194        run.stack.push(key.clone());
195        let result = self.evaluate_pushed(key, run);
196        run.stack.pop();
197        result
198    }
199
200    fn evaluate_pushed(&mut self, key: K, run: &mut RunState<K>) -> QueryResult<K, V> {
201        if self.try_reuse_memo(&key, run)? {
202            let value = self
203                .nodes
204                .get(&key)
205                .and_then(|node| node.value.clone())
206                .ok_or_else(|| IncrementalError::UnknownQuery { key: key.clone() })?;
207            return Ok(value);
208        }
209
210        self.charge_work(run, &key, 1)?;
211        let query = self
212            .queries
213            .get(&key)
214            .cloned()
215            .ok_or_else(|| IncrementalError::UnknownQuery { key: key.clone() })?;
216        let mut observations = Vec::new();
217        let value = {
218            let mut frame = QueryFrame {
219                engine: self,
220                run,
221                observations: &mut observations,
222            };
223            query.run(&key, &mut frame)
224        };
225
226        let value = value?;
227        self.charge_output(run, &key, 1)?;
228        self.commit_value(key.clone(), value, observations);
229        self.nodes
230            .get(&key)
231            .and_then(|node| node.value.clone())
232            .ok_or(IncrementalError::UnknownQuery { key })
233    }
234
235    fn try_reuse_memo(
236        &mut self,
237        key: &K,
238        run: &mut RunState<K>,
239    ) -> Result<bool, IncrementalError<K>> {
240        let Some(node) = self.nodes.get(key) else {
241            return Ok(false);
242        };
243        if node.value.is_none() {
244            return Ok(false);
245        }
246        let dependencies = node.dependencies.clone();
247        let needs_refresh = node.dirty
248            || dependencies
249                .iter()
250                .any(|observation| !self.observation_is_current(observation));
251        if needs_refresh {
252            for observation in dependencies
253                .iter()
254                .filter(|observation| matches!(observation.kind(), ObservationKind::Read))
255            {
256                self.evaluate(observation.key().clone(), run)?;
257            }
258        }
259        if dependencies
260            .iter()
261            .all(|observation| self.observation_is_current(observation))
262        {
263            if let Some(node) = self.nodes.get_mut(key) {
264                node.dirty = false;
265            }
266            Ok(true)
267        } else {
268            Ok(false)
269        }
270    }
271
272    fn commit_value(&mut self, key: K, value: V, dependencies: Vec<Observation<K>>) {
273        let fingerprint = value.incremental_fingerprint();
274        let old_dependencies = self
275            .nodes
276            .get(&key)
277            .map(|node| node.dependencies.clone())
278            .unwrap_or_default();
279        for observation in old_dependencies {
280            if let Some(dependents) = self.reverse.get_mut(observation.key()) {
281                dependents.remove(&key);
282            }
283        }
284
285        let same_value = self
286            .nodes
287            .get(&key)
288            .and_then(|node| node.fingerprint)
289            .is_some_and(|old| old == fingerprint);
290        let revision = if same_value {
291            self.nodes
292                .get(&key)
293                .map(|node| node.revision)
294                .unwrap_or_else(|| self.alloc_revision())
295        } else {
296            self.alloc_revision()
297        };
298
299        for observation in &dependencies {
300            self.reverse
301                .entry(observation.key().clone())
302                .or_default()
303                .insert(key.clone());
304        }
305        self.nodes.insert(
306            key,
307            Node {
308                revision,
309                dirty: false,
310                value: Some(value),
311                fingerprint: Some(fingerprint),
312                dependencies,
313            },
314        );
315    }
316
317    fn observation_is_current(&self, observation: &Observation<K>) -> bool {
318        match observation.kind() {
319            ObservationKind::Read => self.nodes.get(observation.key()).is_some_and(|node| {
320                !node.dirty
321                    && node.revision == observation.revision()
322                    && node.fingerprint == observation.fingerprint()
323            }),
324            ObservationKind::Missing
325            | ObservationKind::Listing
326            | ObservationKind::Policy
327            | ObservationKind::Epoch
328            | ObservationKind::Custom(_) => {
329                self.source_revision(observation.key()) == observation.revision()
330            }
331        }
332    }
333
334    pub(crate) fn memo_observation(&self, key: &K) -> Result<Observation<K>, IncrementalError<K>> {
335        let node = self
336            .nodes
337            .get(key)
338            .ok_or_else(|| IncrementalError::UnknownQuery { key: key.clone() })?;
339        let fingerprint = node
340            .fingerprint
341            .ok_or_else(|| IncrementalError::UnknownQuery { key: key.clone() })?;
342        Ok(Observation::read(key.clone(), node.revision, fingerprint))
343    }
344
345    pub(crate) fn record_observation(
346        &mut self,
347        run: &mut RunState<K>,
348        observations: &mut Vec<Observation<K>>,
349        observation: Observation<K>,
350    ) -> Result<(), IncrementalError<K>> {
351        self.charge_observation(run, observation.key())?;
352        observations.push(observation);
353        Ok(())
354    }
355
356    pub(crate) fn charge_work(
357        &mut self,
358        run: &mut RunState<K>,
359        key: &K,
360        units: usize,
361    ) -> Result<(), IncrementalError<K>> {
362        self.check_cancelled(run)?;
363        if run.work.saturating_add(units) > run.budgets.max_work {
364            return Err(self.budget_error(
365                run,
366                key,
367                BudgetKind::Work,
368                run.budgets.max_work,
369                run.work.saturating_add(units),
370            ));
371        }
372        run.work += units;
373        Ok(())
374    }
375
376    pub(crate) fn charge_output(
377        &mut self,
378        run: &mut RunState<K>,
379        key: &K,
380        units: usize,
381    ) -> Result<(), IncrementalError<K>> {
382        self.check_cancelled(run)?;
383        if run.output.saturating_add(units) > run.budgets.max_output {
384            return Err(self.budget_error(
385                run,
386                key,
387                BudgetKind::Output,
388                run.budgets.max_output,
389                run.output.saturating_add(units),
390            ));
391        }
392        run.output += units;
393        Ok(())
394    }
395
396    fn charge_depth(&mut self, run: &mut RunState<K>, key: &K) -> Result<(), IncrementalError<K>> {
397        if run.stack.len().saturating_add(1) > run.budgets.max_depth {
398            return Err(self.budget_error(
399                run,
400                key,
401                BudgetKind::Depth,
402                run.budgets.max_depth,
403                run.stack.len().saturating_add(1),
404            ));
405        }
406        Ok(())
407    }
408
409    fn charge_observation(
410        &mut self,
411        run: &mut RunState<K>,
412        key: &K,
413    ) -> Result<(), IncrementalError<K>> {
414        if run.observations.saturating_add(1) > run.budgets.max_observations {
415            return Err(self.budget_error(
416                run,
417                key,
418                BudgetKind::Observations,
419                run.budgets.max_observations,
420                run.observations.saturating_add(1),
421            ));
422        }
423        run.observations += 1;
424        Ok(())
425    }
426
427    fn check_cancelled(&self, run: &RunState<K>) -> Result<(), IncrementalError<K>> {
428        if run.cancelled {
429            Err(IncrementalError::Cancelled)
430        } else {
431            Ok(())
432        }
433    }
434
435    fn budget_error(
436        &mut self,
437        run: &RunState<K>,
438        _key: &K,
439        kind: BudgetKind,
440        limit: usize,
441        consumed: usize,
442    ) -> IncrementalError<K> {
443        let continuation = Some(self.alloc_continuation(run.root.clone()));
444        IncrementalError::BudgetExceeded {
445            kind,
446            limit,
447            consumed,
448            continuation,
449        }
450    }
451}
452
453impl<K, V> IncrementalEngine<K, V>
454where
455    K: Ord + Clone,
456{
457    fn mark_dirty_cascade(&mut self, key: &K) {
458        let mut pending = BTreeSet::from([key.clone()]);
459        let mut seen = BTreeSet::new();
460        while let Some(next) = pending.iter().next().cloned() {
461            pending.remove(&next);
462            if !seen.insert(next.clone()) {
463                continue;
464            }
465            if let Some(node) = self.nodes.get_mut(&next) {
466                node.dirty = true;
467            }
468            if let Some(dependents) = self.reverse.get(&next) {
469                pending.extend(dependents.iter().cloned());
470            }
471        }
472    }
473
474    fn detach_node(&mut self, key: &K) {
475        let Some(node) = self.nodes.remove(key) else {
476            return;
477        };
478        for observation in node.dependencies {
479            if let Some(dependents) = self.reverse.get_mut(observation.key()) {
480                dependents.remove(key);
481            }
482        }
483        self.reverse.remove(key);
484    }
485}
486
487/// Execution context handed to a query callback.
488pub struct QueryFrame<'a, K, V> {
489    pub(crate) engine: &'a mut IncrementalEngine<K, V>,
490    pub(crate) run: &'a mut RunState<K>,
491    pub(crate) observations: &'a mut Vec<Observation<K>>,
492}