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