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 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 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[¤t];
165 for fact_id in projection.incident_fact_ids(¤t) {
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}