weavatrix_memory/context/
compiler.rs1use 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 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[¤t];
162 for fact_id in projection.incident_fact_ids(¤t) {
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}