1use std::collections::{BTreeMap, BTreeSet};
20
21use uqa_core::{DocId, EdgeId, Payload, PostingEntry, PostingList, Value, VertexId};
22
23use crate::posting_list::{GraphPayload, GraphPostingList};
24use crate::store::{GraphStore, GraphStoreError, GraphStoreResult};
25
26pub const MAX_GRAPH_EMBEDDING_DIMENSIONS: usize = 4_096;
27pub const MAX_GRAPH_EMBEDDING_LAYERS: u32 = 256;
28const MAX_EXACT_F64_INTEGER: u64 = 9_007_199_254_740_992;
29
30pub struct GraphEmbedding<'a> {
31 pub graph: &'a str,
32 pub dimensions: usize,
33 pub k_layers: u32,
34}
35
36impl<'a> GraphEmbedding<'a> {
37 pub fn new(graph: &'a str) -> Self {
38 Self {
39 graph,
40 dimensions: 32,
41 k_layers: 2,
42 }
43 }
44
45 pub fn dimensions(mut self, dims: usize) -> Self {
46 self.dimensions = dims;
47 self
48 }
49
50 pub fn k_layers(mut self, k: u32) -> Self {
51 self.k_layers = k;
52 self
53 }
54
55 pub fn execute<G: GraphStore>(&self, store: &G) -> GraphStoreResult<GraphPostingList> {
56 if self.dimensions == 0 || self.dimensions > MAX_GRAPH_EMBEDDING_DIMENSIONS {
57 return Err(GraphStoreError::InvalidQuery(format!(
58 "graph embedding dimensions must be in 1..={MAX_GRAPH_EMBEDDING_DIMENSIONS}, got {}",
59 self.dimensions
60 )));
61 }
62 if self.k_layers > MAX_GRAPH_EMBEDDING_LAYERS {
63 return Err(GraphStoreError::InvalidQuery(format!(
64 "graph embedding layer count {} exceeds limit {MAX_GRAPH_EMBEDDING_LAYERS}",
65 self.k_layers
66 )));
67 }
68 let vertex_ids = store.vertex_ids_in_graph(self.graph)?;
69 let mut vertices: Vec<VertexId> = Vec::new();
70 vertices
71 .try_reserve_exact(vertex_ids.len())
72 .map_err(|error| allocation_error("vertex id list", vertex_ids.len(), &error))?;
73 vertices.extend(vertex_ids.iter().copied());
74 if vertices.is_empty() {
75 return Ok(GraphPostingList::new());
76 }
77 for vid in &vertices {
78 if store.get_vertex(*vid).is_none() {
79 return Err(GraphStoreError::CorruptGraph(format!(
80 "graph embedding graph {:?} references missing vertex {vid}",
81 self.graph
82 )));
83 }
84 }
85
86 let mut all_labels: BTreeSet<String> = BTreeSet::new();
88 for vid in &vertices {
89 for eid in store.out_edge_ids(*vid, self.graph)? {
90 let edge = store.get_edge(eid).ok_or_else(|| {
91 GraphStoreError::CorruptGraph(format!("missing embedding edge {eid}"))
92 })?;
93 all_labels.insert(edge.label.clone());
94 }
95 }
96 let label_to_idx: BTreeMap<String, usize> = all_labels
97 .iter()
98 .enumerate()
99 .map(|(i, label)| (label.clone(), i))
100 .collect();
101 let n_labels = all_labels.len();
102
103 let mut entries: Vec<PostingEntry> = Vec::new();
104 entries
105 .try_reserve_exact(vertices.len())
106 .map_err(|error| allocation_error("posting entries", vertices.len(), &error))?;
107 let mut graph_payloads: BTreeMap<DocId, GraphPayload> = BTreeMap::new();
108 for vid in &vertices {
109 let embedding =
110 self.compute_embedding(store, *vid, &vertex_ids, &label_to_idx, n_labels)?;
111 let mut fields: BTreeMap<String, Value> = BTreeMap::new();
112 let mut embedding_values = Vec::new();
113 embedding_values
114 .try_reserve_exact(embedding.len())
115 .map_err(|error| allocation_error("embedding payload", embedding.len(), &error))?;
116 embedding_values.extend(embedding.iter().map(|x| Value::Float(*x)));
117 fields.insert("_embedding".into(), Value::List(embedding_values));
118 let payload = Payload {
119 positions: Vec::new(),
120 score: 0.0,
121 fields,
122 };
123 entries.push(PostingEntry::new(*vid, payload));
124 graph_payloads.insert(
125 *vid,
126 GraphPayload {
127 subgraph_vertices: vec![*vid],
128 subgraph_edges: Vec::new(),
129 graph_name: self.graph.to_string(),
130 score_override: None,
131 },
132 );
133 }
134 GraphPostingList::try_from_parts(
135 PostingList::from_sorted_unchecked(entries),
136 graph_payloads,
137 )
138 .map_err(Into::into)
139 }
140
141 fn compute_embedding<G: GraphStore>(
142 &self,
143 store: &G,
144 vid: VertexId,
145 graph_vertices: &BTreeSet<VertexId>,
146 label_to_idx: &BTreeMap<String, usize>,
147 n_labels: usize,
148 ) -> GraphStoreResult<Vec<f64>> {
149 let out_edges: BTreeSet<EdgeId> = store.out_edge_ids(vid, self.graph)?;
150 let in_edges: BTreeSet<EdgeId> = store.in_edge_ids(vid, self.graph)?;
151 let out_degree = out_edges.len();
152 let in_degree = in_edges.len();
153
154 let label_dims = n_labels.min(self.dimensions / 2);
155 let mut label_dist = Vec::new();
156 label_dist
157 .try_reserve_exact(label_dims)
158 .map_err(|error| allocation_error("edge-label distribution", label_dims, &error))?;
159 label_dist.resize(label_dims, 0.0_f64);
160 for eid in &out_edges {
161 let edge = store.get_edge(*eid).ok_or_else(|| {
162 GraphStoreError::CorruptGraph(format!("missing embedding edge {eid}"))
163 })?;
164 if let Some(&idx) = label_to_idx.get(&edge.label) {
165 if idx < label_dims {
166 label_dist[idx] += 1.0;
167 }
168 }
169 }
170 let total: f64 = label_dist.iter().sum();
171 if total > 0.0 {
172 for v in &mut label_dist {
173 *v /= total;
174 }
175 }
176
177 let layer_capacity = usize::try_from(self.k_layers).map_err(|_| {
178 GraphStoreError::InvalidQuery(format!(
179 "graph embedding layer count {} does not fit usize",
180 self.k_layers
181 ))
182 })?;
183 let mut hop_counts: Vec<f64> = Vec::new();
184 hop_counts
185 .try_reserve_exact(layer_capacity)
186 .map_err(|error| allocation_error("hop counts", layer_capacity, &error))?;
187 let mut visited: BTreeSet<VertexId> = BTreeSet::from([vid]);
188 let mut frontier: BTreeSet<VertexId> = BTreeSet::from([vid]);
189 for _ in 0..self.k_layers {
190 let mut next_frontier: BTreeSet<VertexId> = BTreeSet::new();
191 for v in &frontier {
192 for eid in store.out_edge_ids(*v, self.graph)? {
193 let edge = store.get_edge(eid).ok_or_else(|| {
194 GraphStoreError::CorruptGraph(format!("missing embedding edge {eid}"))
195 })?;
196 if !graph_vertices.contains(&edge.target_id) {
197 return Err(GraphStoreError::CorruptGraph(format!(
198 "embedding edge {eid} targets vertex {} outside graph {:?}",
199 edge.target_id, self.graph
200 )));
201 }
202 if !visited.contains(&edge.target_id) {
203 next_frontier.insert(edge.target_id);
204 visited.insert(edge.target_id);
205 }
206 }
207 }
208 hop_counts.push(usize_to_f64_exact(
209 next_frontier.len(),
210 "graph embedding hop frontier",
211 )?);
212 frontier = next_frontier;
213 }
214
215 let generated_features = 2_usize
216 .checked_add(label_dims)
217 .and_then(|value| value.checked_add(layer_capacity))
218 .ok_or_else(|| {
219 GraphStoreError::InvalidQuery(
220 "graph embedding feature count overflows usize".to_string(),
221 )
222 })?;
223 let raw_capacity = self.dimensions.max(generated_features);
224 let mut raw: Vec<f64> = Vec::new();
225 raw.try_reserve_exact(raw_capacity)
226 .map_err(|error| allocation_error("raw embedding", raw_capacity, &error))?;
227 raw.push(usize_to_f64_exact(
228 out_degree,
229 "graph embedding out-degree",
230 )?);
231 raw.push(usize_to_f64_exact(in_degree, "graph embedding in-degree")?);
232 raw.extend(label_dist);
233 raw.extend(hop_counts);
234 normalize_embedding(raw, self.dimensions, vid)
235 }
236}
237
238fn usize_to_f64_exact(value: usize, context: &str) -> GraphStoreResult<f64> {
239 if !u64::try_from(value).is_ok_and(|value| value <= MAX_EXACT_F64_INTEGER) {
240 return Err(GraphStoreError::InvalidQuery(format!(
241 "{context} {value} exceeds the exact f64 integer range"
242 )));
243 }
244 Ok(value as f64)
245}
246
247fn allocation_error(
248 context: &str,
249 elements: usize,
250 error: &std::collections::TryReserveError,
251) -> GraphStoreError {
252 GraphStoreError::InvalidQuery(format!(
253 "cannot allocate {elements} elements for graph embedding {context}: {error}"
254 ))
255}
256
257fn normalize_embedding(
258 mut raw: Vec<f64>,
259 dimensions: usize,
260 vertex_id: VertexId,
261) -> GraphStoreResult<Vec<f64>> {
262 if raw.len() < dimensions {
263 raw.resize(dimensions, 0.0);
264 } else {
265 raw.truncate(dimensions);
266 }
267 let norm = raw.iter().map(|value| value * value).sum::<f64>().sqrt();
268 if !norm.is_finite() {
269 return Err(GraphStoreError::InvalidQuery(format!(
270 "graph embedding norm is not finite for vertex {vertex_id}"
271 )));
272 }
273 if norm > 0.0 {
274 for value in &mut raw {
275 *value /= norm;
276 }
277 }
278 Ok(raw)
279}
280
281#[cfg(test)]
282mod tests {
283 use uqa_core::{Edge, Vertex};
284
285 use super::*;
286 use crate::memory_store::MemoryGraphStore;
287
288 #[test]
289 fn missing_edge_record_is_corruption_not_an_empty_label_bucket() {
290 let mut store = MemoryGraphStore::new();
291 store.create_graph("g");
292 store.add_vertex(Vertex::new(1, "n"), "g").unwrap();
293 store.add_vertex(Vertex::new(2, "n"), "g").unwrap();
294 store.add_edge(Edge::new(10, 1, 2, "edge"), "g").unwrap();
295 store.remove_edge_record_for_corruption_test(10);
296
297 let error = GraphEmbedding::new("g").execute(&store).unwrap_err();
298 assert!(matches!(error, GraphStoreError::CorruptGraph(_)));
299 }
300}