Skip to main content

weavatrix_memory/context/
compiler.rs

1use super::{
2    BytesTokenEstimator, ContextBundle, ContextReceipt, ContextRequest, TokenEstimator,
3    scope::{node_allowed, relation_allowed},
4    token::{fact_tokens, node_tokens},
5};
6use crate::{
7    EntityId, MemoryError, MemoryFact, MemoryNode, MemoryProjection, MemoryView, ProjectionClock,
8    Result, project_graph,
9};
10use std::{
11    cmp::Reverse,
12    collections::{BTreeMap, BTreeSet, VecDeque},
13};
14
15#[derive(Debug, Clone)]
16pub struct ContextCompiler<T = BytesTokenEstimator> {
17    estimator: T,
18}
19
20impl<T> ContextCompiler<T>
21where
22    T: TokenEstimator,
23{
24    #[must_use]
25    pub const fn new(estimator: T) -> Self {
26        Self { estimator }
27    }
28
29    /// Compiles a deterministic, evidence-carrying graph within a hard budget.
30    ///
31    /// # Errors
32    ///
33    /// Returns an error for missing seeds, graph failures, or an undersized
34    /// budget that cannot hold the requested seed nodes.
35    pub fn compile(
36        &self,
37        projection: &MemoryProjection,
38        request: &ContextRequest,
39    ) -> Result<ContextBundle> {
40        if request.seeds.is_empty() {
41            return Err(MemoryError::InvalidValue {
42                field: "context.seeds",
43                reason: "at least one seed is required",
44            });
45        }
46        let (nodes_by_id, candidate_facts, distances, excluded_by_scope) =
47            neighborhood(projection, request)?;
48        let (mut selected_nodes, mut used) = self.select_seeds(request, &nodes_by_id)?;
49
50        let mut candidates = candidate_facts.values().copied().collect::<Vec<_>>();
51        candidates.sort_by_key(|fact| {
52            (
53                distances[&fact.source].min(distances[&fact.target]),
54                Reverse(fact.confidence),
55                Reverse(fact.recorded_at),
56                fact.id.clone(),
57            )
58        });
59
60        let mut facts = Vec::new();
61        let mut omitted = 0;
62        for fact in &candidates {
63            let missing_nodes = [&fact.source, &fact.target]
64                .into_iter()
65                .filter(|id| !selected_nodes.contains(*id))
66                .collect::<Vec<_>>();
67            let node_cost = missing_nodes
68                .iter()
69                .map(|id| node_tokens(&self.estimator, &nodes_by_id[id]))
70                .sum::<usize>();
71            let cost = fact_tokens(&self.estimator, fact) + node_cost;
72            if used.saturating_add(cost) > request.token_budget {
73                omitted += 1;
74                continue;
75            }
76            used += cost;
77            for id in missing_nodes {
78                selected_nodes.insert(id.clone());
79            }
80            facts.push((*fact).clone());
81        }
82        let nodes = selected_nodes
83            .iter()
84            .map(|id| nodes_by_id[id].clone())
85            .collect::<Vec<_>>();
86        facts.sort_by(|left, right| left.id.cmp(&right.id));
87        let view = MemoryView { nodes, facts };
88        let graph = project_graph(&view)?;
89        Ok(ContextBundle {
90            receipt: ContextReceipt {
91                valid_at: request.valid_at,
92                known_at: request.known_at,
93                source_position: projection.last_global_position(),
94                estimator: self.estimator.name().to_owned(),
95                token_budget: request.token_budget,
96                estimated_tokens: used,
97                examined_facts: candidates.len(),
98                selected_facts: view.facts.len(),
99                omitted_by_budget: omitted,
100                excluded_by_scope,
101            },
102            view,
103            graph,
104        })
105    }
106
107    fn select_seeds(
108        &self,
109        request: &ContextRequest,
110        nodes: &BTreeMap<EntityId, MemoryNode>,
111    ) -> Result<(BTreeSet<crate::EntityId>, usize)> {
112        let mut selected = BTreeSet::new();
113        let mut used = 0;
114        let mut seeds = request.seeds.clone();
115        seeds.sort();
116        seeds.dedup();
117        for seed in &seeds {
118            let node = nodes.get(seed).ok_or_else(|| MemoryError::MissingEntity {
119                id: seed.to_string(),
120            })?;
121            used += node_tokens(&self.estimator, node);
122            selected.insert(seed.clone());
123        }
124        if used > request.token_budget {
125            return Err(MemoryError::BudgetTooSmall {
126                required: used,
127                available: request.token_budget,
128            });
129        }
130        Ok((selected, used))
131    }
132}
133
134type Neighborhood<'a> = (
135    BTreeMap<EntityId, MemoryNode>,
136    BTreeMap<crate::FactId, &'a MemoryFact>,
137    BTreeMap<EntityId, usize>,
138    usize,
139);
140
141fn neighborhood<'a>(
142    projection: &'a MemoryProjection,
143    request: &ContextRequest,
144) -> Result<Neighborhood<'a>> {
145    let clock = ProjectionClock::new(request.valid_at, request.known_at);
146    let mut nodes = BTreeMap::new();
147    let mut facts = BTreeMap::new();
148    let mut distances = BTreeMap::new();
149    let mut queue = VecDeque::new();
150    let mut excluded = BTreeSet::new();
151    let mut seeds = request.seeds.clone();
152    seeds.sort();
153    seeds.dedup();
154    for seed in seeds {
155        let node = scoped_node(projection, request, &seed)?;
156        nodes.insert(seed.clone(), node.clone());
157        distances.insert(seed.clone(), 0);
158        queue.push_back(seed);
159    }
160    while let Some(current) = queue.pop_front() {
161        let depth = distances[&current];
162        for fact_id in projection.incident_fact_ids(&current) {
163            let fact = projection.fact(fact_id).expect("projection index is valid");
164            if !projection.fact_is_active(fact, clock) || !relation_allowed(request, &fact.relation)
165            {
166                continue;
167            }
168            let other = if fact.source == current {
169                &fact.target
170            } else {
171                &fact.source
172            };
173            let Some(other_node) = projection.visible_node(other, request.known_at) else {
174                continue;
175            };
176            if !node_allowed(request, other_node) {
177                excluded.insert(fact.id.clone());
178                continue;
179            }
180            if !distances.contains_key(other) {
181                if depth >= request.max_depth {
182                    continue;
183                }
184                distances.insert(other.clone(), depth + 1);
185                nodes.insert(other.clone(), other_node.clone());
186                queue.push_back(other.clone());
187            }
188            facts.insert(fact.id.clone(), fact);
189        }
190    }
191    facts.retain(|_, fact| {
192        distances.contains_key(&fact.source) && distances.contains_key(&fact.target)
193    });
194    Ok((nodes, facts, distances, excluded.len()))
195}
196
197fn scoped_node<'a>(
198    projection: &'a MemoryProjection,
199    request: &ContextRequest,
200    id: &EntityId,
201) -> Result<&'a MemoryNode> {
202    projection
203        .visible_node(id, request.known_at)
204        .filter(|node| node_allowed(request, node))
205        .ok_or_else(|| MemoryError::MissingEntity { id: id.to_string() })
206}
207
208impl Default for ContextCompiler<BytesTokenEstimator> {
209    fn default() -> Self {
210        Self::new(BytesTokenEstimator::default())
211    }
212}