1use std::collections::{BTreeMap, BTreeSet};
16
17use uqa_core::types::{
18 GraphPhiEnvelope, GraphPhiPayload, GRAPH_PHI_EDGES_FIELD, GRAPH_PHI_FIELD,
19 GRAPH_PHI_VERTICES_FIELD,
20};
21use uqa_core::{DocId, EdgeId, PostingEntry, PostingList, Value, VertexId};
22
23#[derive(Debug, Clone, Default, PartialEq)]
26pub struct GraphPayload {
27 pub subgraph_vertices: Vec<VertexId>,
28 pub subgraph_edges: Vec<EdgeId>,
29 pub graph_name: String,
30 pub score_override: Option<f64>,
33}
34
35impl GraphPayload {
36 pub fn new() -> Self {
37 Self::default()
38 }
39}
40
41#[derive(Debug, Clone, Copy, PartialEq, Eq)]
46pub enum SubgraphMergePolicy {
47 Union,
48 Intersection,
49 PreferLeft,
50 PreferRight,
51}
52
53#[derive(Debug, Clone, PartialEq, Eq, thiserror::Error)]
54pub enum GraphPostingListError {
55 #[error("graph payload for document {doc_id} is outside the posting-list support")]
56 PayloadOutsideSupport { doc_id: DocId },
57 #[error(
58 "cannot combine graph payloads for document {doc_id}: graph names {left:?} and {right:?} conflict"
59 )]
60 ConflictingGraphNames {
61 doc_id: DocId,
62 left: String,
63 right: String,
64 },
65}
66
67pub type GraphPostingListResult<T> = Result<T, GraphPostingListError>;
68
69#[derive(Debug, Clone, Default, PartialEq)]
73pub struct GraphPostingList {
74 inner: PostingList,
75 graph_payloads: BTreeMap<DocId, GraphPayload>,
76}
77
78impl GraphPostingList {
79 pub fn new() -> Self {
80 Self::default()
81 }
82
83 pub fn try_from_parts(
86 inner: PostingList,
87 graph_payloads: BTreeMap<DocId, GraphPayload>,
88 ) -> GraphPostingListResult<Self> {
89 if let Some(doc_id) = graph_payloads
90 .keys()
91 .find(|doc_id| inner.get_entry(**doc_id).is_none())
92 {
93 return Err(GraphPostingListError::PayloadOutsideSupport { doc_id: *doc_id });
94 }
95 Ok(Self {
96 inner,
97 graph_payloads,
98 })
99 }
100
101 pub fn to_posting_list(&self) -> PostingList {
110 let mut converted = Vec::with_capacity(self.inner.len());
111 for entry in self.inner.entries() {
112 let mut payload = entry.payload.clone();
113 let graph_payload = self.graph_payloads.get(&entry.doc_id);
114 let needs_envelope = graph_payload.is_some()
115 || payload.fields.contains_key(GRAPH_PHI_FIELD)
116 || payload.fields.contains_key(GRAPH_PHI_VERTICES_FIELD)
117 || payload.fields.contains_key(GRAPH_PHI_EDGES_FIELD);
118 if needs_envelope {
119 let original_reserved = payload.fields.remove(GRAPH_PHI_FIELD);
120 let original_vertices = payload.fields.remove(GRAPH_PHI_VERTICES_FIELD);
121 let original_edges = payload.fields.remove(GRAPH_PHI_EDGES_FIELD);
122 let encoded_graph = graph_payload.map(|gp| GraphPhiPayload {
123 vertices: gp.subgraph_vertices.clone(),
124 edges: gp.subgraph_edges.clone(),
125 graph_name: gp.graph_name.clone(),
126 });
127
128 if let Some(graph) = &encoded_graph {
129 payload.fields.insert(
130 GRAPH_PHI_VERTICES_FIELD.to_string(),
131 graph.encoded_vertices(),
132 );
133 payload
134 .fields
135 .insert(GRAPH_PHI_EDGES_FIELD.to_string(), graph.encoded_edges());
136 } else {
137 restore_payload_field(
138 &mut payload.fields,
139 GRAPH_PHI_VERTICES_FIELD,
140 original_vertices.clone(),
141 );
142 restore_payload_field(
143 &mut payload.fields,
144 GRAPH_PHI_EDGES_FIELD,
145 original_edges.clone(),
146 );
147 }
148 if let Some(score) = graph_payload.and_then(|gp| gp.score_override) {
149 payload.score = score;
150 }
151 payload.fields.insert(
152 GRAPH_PHI_FIELD.to_string(),
153 GraphPhiEnvelope {
154 base_score: entry.payload.score,
155 graph_payload: encoded_graph,
156 score_override: graph_payload.and_then(|gp| gp.score_override),
157 original_reserved,
158 original_vertices,
159 original_edges,
160 }
161 .encode(),
162 );
163 }
164 converted.push(PostingEntry::new(entry.doc_id, payload));
165 }
166 PostingList::from_sorted_unchecked(converted)
167 }
168
169 pub fn from_posting_list(pl: &PostingList) -> Self {
175 let mut graph_payloads: BTreeMap<DocId, GraphPayload> = BTreeMap::new();
176 let mut entries = Vec::with_capacity(pl.len());
177 for entry in pl.entries() {
178 let envelope_value = entry.payload.fields.get(GRAPH_PHI_FIELD);
179 if let Some(envelope) = GraphPhiEnvelope::decode(envelope_value) {
180 let mut p = entry.payload.clone();
181 p.score = envelope.base_score;
182 p.fields.remove(GRAPH_PHI_FIELD);
183 p.fields.remove(GRAPH_PHI_VERTICES_FIELD);
184 p.fields.remove(GRAPH_PHI_EDGES_FIELD);
185 restore_payload_field(&mut p.fields, GRAPH_PHI_FIELD, envelope.original_reserved);
186 restore_payload_field(
187 &mut p.fields,
188 GRAPH_PHI_VERTICES_FIELD,
189 envelope.original_vertices,
190 );
191 restore_payload_field(
192 &mut p.fields,
193 GRAPH_PHI_EDGES_FIELD,
194 envelope.original_edges,
195 );
196 if let Some(graph) = envelope.graph_payload {
197 graph_payloads.insert(
198 entry.doc_id,
199 GraphPayload {
200 subgraph_vertices: graph.vertices,
201 subgraph_edges: graph.edges,
202 graph_name: graph.graph_name,
203 score_override: envelope.score_override,
204 },
205 );
206 }
207 entries.push(PostingEntry::new(entry.doc_id, p));
208 continue;
209 }
210
211 if GraphPhiEnvelope::is_recognized(envelope_value) {
212 entries.push(entry.clone());
213 continue;
214 }
215
216 let vertices = decode_id_list(entry.payload.fields.get(GRAPH_PHI_VERTICES_FIELD));
217 let edges = decode_id_list(entry.payload.fields.get(GRAPH_PHI_EDGES_FIELD));
218 let has_graph_fields = entry.payload.fields.contains_key(GRAPH_PHI_VERTICES_FIELD)
219 || entry.payload.fields.contains_key(GRAPH_PHI_EDGES_FIELD);
220 let mut payload = entry.payload.clone();
221 if has_graph_fields {
222 payload.fields.remove(GRAPH_PHI_VERTICES_FIELD);
223 payload.fields.remove(GRAPH_PHI_EDGES_FIELD);
224 graph_payloads.insert(
225 entry.doc_id,
226 GraphPayload {
227 subgraph_vertices: vertices,
228 subgraph_edges: edges,
229 graph_name: String::new(),
230 score_override: Some(entry.payload.score),
231 },
232 );
233 }
234 entries.push(PostingEntry::new(entry.doc_id, payload));
235 }
236 Self {
237 inner: PostingList::from_sorted_unchecked(entries),
238 graph_payloads,
239 }
240 }
241
242 pub fn try_set_graph_payload(
244 &mut self,
245 doc_id: DocId,
246 payload: GraphPayload,
247 ) -> GraphPostingListResult<()> {
248 if self.inner.get_entry(doc_id).is_none() {
249 return Err(GraphPostingListError::PayloadOutsideSupport { doc_id });
250 }
251 self.graph_payloads.insert(doc_id, payload);
252 Ok(())
253 }
254
255 pub fn get_graph_payload(&self, doc_id: DocId) -> Option<&GraphPayload> {
256 self.graph_payloads.get(&doc_id)
257 }
258
259 pub fn inner(&self) -> &PostingList {
260 &self.inner
261 }
262
263 pub fn len(&self) -> usize {
264 self.inner.len()
265 }
266
267 pub fn is_empty(&self) -> bool {
268 self.inner.is_empty()
269 }
270
271 pub fn merge_union(&self, other: &Self) -> GraphPostingListResult<Self> {
274 self.merge_union_with(other, SubgraphMergePolicy::Union)
275 }
276
277 pub fn merge_intersection(&self, other: &Self) -> GraphPostingListResult<Self> {
280 self.merge_intersection_with(other, SubgraphMergePolicy::Intersection)
281 }
282
283 pub fn merge_union_with(
285 &self,
286 other: &Self,
287 policy: SubgraphMergePolicy,
288 ) -> GraphPostingListResult<Self> {
289 let inner = self.inner.merge_union(&other.inner);
290 let graph_payloads = self.merge_graph_payloads(other, &inner, policy)?;
291 Self::try_from_parts(inner, graph_payloads)
292 }
293
294 pub fn merge_intersection_with(
296 &self,
297 other: &Self,
298 policy: SubgraphMergePolicy,
299 ) -> GraphPostingListResult<Self> {
300 let inner = self.inner.merge_intersection(&other.inner);
301 let graph_payloads = self.merge_graph_payloads(other, &inner, policy)?;
302 Self::try_from_parts(inner, graph_payloads)
303 }
304
305 pub fn exclude(&self, other: &Self) -> Self {
306 let inner = self.inner.exclude(&other.inner);
307 let graph_payloads = self
308 .graph_payloads
309 .iter()
310 .filter(|(doc_id, _)| inner.get_entry(**doc_id).is_some())
311 .map(|(doc_id, payload)| (*doc_id, payload.clone()))
312 .collect();
313 Self {
314 inner,
315 graph_payloads,
316 }
317 }
318
319 fn merge_graph_payloads(
320 &self,
321 other: &Self,
322 merged_inner: &PostingList,
323 policy: SubgraphMergePolicy,
324 ) -> GraphPostingListResult<BTreeMap<DocId, GraphPayload>> {
325 let mut merged = BTreeMap::new();
326 for entry in merged_inner {
327 let doc_id = entry.doc_id;
328 let left = self.graph_payloads.get(&doc_id);
329 let right = other.graph_payloads.get(&doc_id);
330 let payload = match (left, right) {
331 (None, None) => continue,
332 (Some(payload), None) | (None, Some(payload)) => payload.clone(),
333 (Some(left), Some(right)) => merge_graph_payload(doc_id, left, right, policy)?,
334 };
335
336 let overlaps =
337 self.inner.get_entry(doc_id).is_some() && other.inner.get_entry(doc_id).is_some();
338 let mut payload = payload;
339 if overlaps
340 && (left.and_then(|value| value.score_override).is_some()
341 || right.and_then(|value| value.score_override).is_some())
342 {
343 let left_score = effective_score(self, doc_id);
344 let right_score = effective_score(other, doc_id);
345 payload.score_override = Some(left_score + right_score);
346 }
347 merged.insert(doc_id, payload);
348 }
349 Ok(merged)
350 }
351}
352
353fn effective_score(list: &GraphPostingList, doc_id: DocId) -> f64 {
354 list.graph_payloads
355 .get(&doc_id)
356 .and_then(|payload| payload.score_override)
357 .unwrap_or_else(|| {
358 list.inner
359 .get_entry(doc_id)
360 .map_or(0.0, |entry| entry.payload.score)
361 })
362}
363
364fn merge_graph_payload(
365 doc_id: DocId,
366 left: &GraphPayload,
367 right: &GraphPayload,
368 policy: SubgraphMergePolicy,
369) -> GraphPostingListResult<GraphPayload> {
370 let graph_name = match policy {
371 SubgraphMergePolicy::PreferLeft => left.graph_name.clone(),
372 SubgraphMergePolicy::PreferRight => right.graph_name.clone(),
373 SubgraphMergePolicy::Union | SubgraphMergePolicy::Intersection => {
374 compatible_graph_name(doc_id, &left.graph_name, &right.graph_name)?
375 }
376 };
377 let (subgraph_vertices, subgraph_edges) = match policy {
378 SubgraphMergePolicy::Union => (
379 set_union(&left.subgraph_vertices, &right.subgraph_vertices),
380 set_union(&left.subgraph_edges, &right.subgraph_edges),
381 ),
382 SubgraphMergePolicy::Intersection => (
383 set_intersection(&left.subgraph_vertices, &right.subgraph_vertices),
384 set_intersection(&left.subgraph_edges, &right.subgraph_edges),
385 ),
386 SubgraphMergePolicy::PreferLeft => {
387 (left.subgraph_vertices.clone(), left.subgraph_edges.clone())
388 }
389 SubgraphMergePolicy::PreferRight => (
390 right.subgraph_vertices.clone(),
391 right.subgraph_edges.clone(),
392 ),
393 };
394 Ok(GraphPayload {
395 subgraph_vertices,
396 subgraph_edges,
397 graph_name,
398 score_override: None,
399 })
400}
401
402fn compatible_graph_name(doc_id: DocId, left: &str, right: &str) -> GraphPostingListResult<String> {
403 if left == right || right.is_empty() {
404 Ok(left.to_string())
405 } else if left.is_empty() {
406 Ok(right.to_string())
407 } else {
408 Err(GraphPostingListError::ConflictingGraphNames {
409 doc_id,
410 left: left.to_string(),
411 right: right.to_string(),
412 })
413 }
414}
415
416fn set_union<T: Copy + Ord>(left: &[T], right: &[T]) -> Vec<T> {
417 left.iter()
418 .chain(right)
419 .copied()
420 .collect::<BTreeSet<_>>()
421 .into_iter()
422 .collect()
423}
424
425fn set_intersection<T: Copy + Ord>(left: &[T], right: &[T]) -> Vec<T> {
426 let right: BTreeSet<_> = right.iter().copied().collect();
427 left.iter()
428 .copied()
429 .filter(|value| right.contains(value))
430 .collect::<BTreeSet<_>>()
431 .into_iter()
432 .collect()
433}
434
435fn restore_payload_field(fields: &mut BTreeMap<String, Value>, key: &str, value: Option<Value>) {
436 if let Some(value) = value {
437 fields.insert(key.to_string(), value);
438 }
439}
440
441fn decode_id_list(value: Option<&Value>) -> Vec<u64> {
442 let Some(Value::List(items)) = value else {
443 return Vec::new();
444 };
445 items
446 .iter()
447 .filter_map(|v| match v {
448 Value::Int(n) => u64::try_from(*n).ok(),
449 Value::Bytes(bytes) if bytes.len() == size_of::<u64>() => {
450 let mut encoded = [0_u8; size_of::<u64>()];
451 encoded.copy_from_slice(bytes);
452 Some(u64::from_be_bytes(encoded))
453 }
454 _ => None,
455 })
456 .collect()
457}
458
459#[cfg(test)]
460mod tests;