1use std::cmp::Ordering;
7use std::collections::HashMap;
8
9use omgbase_search::EmbeddingProvider;
10use omgbase_store::Store;
11use oqx::ast::{BinaryOp, Consumer, FollowDestination, SelectItem, Where};
12use oqx::{build, print_query};
13use serde_json::{Map, Value as Json, json};
14
15use crate::error::{Result, SurfaceError};
16use crate::query::{QueryOptions, query};
17use crate::read::find_doc_by_ref;
18
19const DEFAULT_DEGREES: i64 = 1;
20const DEFAULT_MAX_DOCUMENTS: i64 = 200;
21const MAX_DEPTH: i64 = 8;
22
23#[derive(Clone, Debug, Default)]
25pub struct GraphArgs {
26 pub roots: Vec<String>,
27 pub degrees: Option<i64>,
28 pub direction: Option<String>,
30 pub predicate: Option<String>,
31 pub select: Vec<String>,
32 pub max_documents: Option<i64>,
33}
34
35const EDGE_FIELDS: [(&str, &str); 10] = [
40 ("id", "$id"),
41 ("src", "$src"),
42 ("dst", "$dst"),
43 ("dst_path", "$dst_path"),
44 ("dst_uri", "$dst_uri"),
45 ("dst_kind", "dst_kind"),
46 ("predicate", "predicate"),
47 ("provenance", "provenance"),
48 ("anchor", "anchor"),
49 ("src_field", "src_field"),
50];
51
52fn build_user_select(select: &[String]) -> Result<(Vec<SelectItem>, Vec<Option<String>>)> {
59 let mut items = Vec::new();
60 let mut out_names: Vec<Option<String>> = vec![None; select.len()];
61 let mut used: Vec<String> = Vec::new();
62 for (i, expr) in select.iter().enumerate() {
63 let trimmed = expr.trim();
64 if trimmed.is_empty() {
65 continue;
66 }
67 let ident = trimmed.strip_prefix('$').unwrap_or(trimmed);
68 let is_ident = ident
69 .chars()
70 .next()
71 .is_some_and(|c| c.is_ascii_alphabetic() || c == '_')
72 && ident.chars().all(|c| c.is_ascii_alphanumeric() || c == '_');
73 let mut name = if is_ident {
74 ident.to_owned()
75 } else {
76 format!("sel_{i}")
77 };
78 while used.contains(&name) {
79 name = format!("{name}_{i}");
80 }
81 used.push(name.clone());
82 out_names[i] = Some(name);
83 let parsed = oqx::parse_string(&format!("select _u{i}: {trimmed} from docs"))?;
84 items.extend(parsed.select);
85 }
86 Ok((items, out_names))
87}
88
89fn build_query(root_ids: &[String], dir: &str, depth: u32, user_select: &[SelectItem]) -> String {
98 let mut edges_body = build::subquery();
99 edges_body.select = EDGE_FIELDS
100 .iter()
101 .map(|(name, expr)| build::field(name, build::ident(expr)))
102 .collect();
103 let edges = build::op(
104 build::path(&["doc", &format!("{dir}_edges")]),
105 Consumer::Collect,
106 edges_body,
107 );
108 let mut seed: Vec<Where> = root_ids
109 .iter()
110 .map(|id| {
111 build::scalar(build::binary(
112 BinaryOp::Eq,
113 build::ident("$id"),
114 build::lit(id.as_str()),
115 ))
116 })
117 .collect();
118 let mut q = build::query(build::ident("docs"));
119 q.select = vec![
120 build::field("_depth", build::ident("$depth")),
121 build::field("_stop", build::ident("$stop")),
122 build::collect("_edges", edges),
123 ];
124 q.select.extend(user_select.iter().cloned());
125 q.r#where = Some(if seed.len() == 1 {
126 seed.pop().expect("one seed")
127 } else {
128 build::or(seed)
129 });
130 let mut walk = build::follow(vec![FollowDestination::Relation(build::path(&[
131 "doc", dir,
132 ]))]);
133 walk.distinct = true;
134 walk.depth = Some(depth);
135 q.follow = Some(walk);
136 print_query(&q).expect("a built query has no bindings")
137}
138
139fn locale_compare(a: &str, b: &str) -> Ordering {
141 a.cmp(b)
142}
143
144struct Doc {
145 id: String,
146 path: String,
147 degree: i64,
148 extra: Vec<(String, Json)>,
149}
150
151pub fn graph_neighborhood(
153 store: &Store,
154 repo_id: &str,
155 args: &GraphArgs,
156 provider: Option<&dyn EmbeddingProvider>,
157) -> Result<Json> {
158 if args.roots.is_empty() {
159 return Err(SurfaceError::new(
160 "target_missing",
161 "graph requires at least one root (path or id)",
162 ));
163 }
164 let mut root_ids: Vec<String> = Vec::new();
165 for r in &args.roots {
166 let Some(info) = find_doc_by_ref(store.conn(), repo_id, r)? else {
167 return Err(SurfaceError::with_data(
168 "doc_missing",
169 format!("no document for {}", Json::String(r.clone())),
170 json!({ "root": r }),
171 ));
172 };
173 if !root_ids.contains(&info.doc_id) {
174 root_ids.push(info.doc_id);
175 }
176 }
177 let degrees = args.degrees.unwrap_or(DEFAULT_DEGREES).max(0);
178 let depth = MAX_DEPTH.min(degrees + 1);
179 let walk_depth = u32::try_from(depth).expect("1..=8");
180 let effective_degrees = depth - 1;
181 let direction = args.direction.clone().unwrap_or_else(|| "both".to_owned());
182 let max_documents =
183 usize::try_from(args.max_documents.unwrap_or(DEFAULT_MAX_DOCUMENTS).max(1)).unwrap_or(1);
184 let dirs: Vec<&str> = match direction.as_str() {
185 "both" => vec!["out", "in"],
186 "in" => vec!["in"],
187 _ => vec!["out"],
188 };
189 let (user_select, out_names) = build_user_select(&args.select)?;
190
191 let mut queries = Vec::new();
192 let mut docs: Vec<Doc> = Vec::new();
193 let mut edges: Vec<Json> = Vec::new();
194 let mut query_truncated = false;
195 for dir in &dirs {
196 let q = build_query(&root_ids, dir, walk_depth, &user_select);
197 queries.push(q.clone());
198 let res = query(
199 store,
200 repo_id,
201 &q,
202 QueryOptions {
203 limit: Some(max_documents + 1),
204 cursor: None,
205 provider,
206 in_memory: false,
207 },
208 )?;
209 if res.truncated {
210 query_truncated = true;
211 }
212 for hit in &res.hits {
213 let id = hit["id"].as_str().unwrap_or_default().to_owned();
214 let path = hit["path"].as_str().unwrap_or_default().to_owned();
215 let hop = hit["_depth"].as_f64().unwrap_or(f64::NAN);
216 let degree = (hop - 1.0) as i64;
217 let replace = docs
218 .iter()
219 .position(|d| d.id == id)
220 .map(|i| (i, degree < docs[i].degree));
221 match replace {
222 Some((_, false)) => {}
223 found => {
224 let extra: Vec<(String, Json)> = out_names
225 .iter()
226 .enumerate()
227 .filter_map(|(i, n)| {
228 n.as_ref().map(|name| {
229 (
230 name.clone(),
231 hit.get(format!("_u{i}")).cloned().unwrap_or(Json::Null),
232 )
233 })
234 })
235 .collect();
236 let doc = Doc {
237 id: id.clone(),
238 path,
239 degree,
240 extra,
241 };
242 match found {
243 Some((i, true)) => docs[i] = doc,
244 _ => docs.push(doc),
245 }
246 }
247 }
248 if let Some(raw) = hit["_edges"].as_array() {
249 for e in raw {
250 let eid = e["id"].as_str().unwrap_or_default();
251 if !edges.iter().any(|x| x["id"].as_str() == Some(eid)) {
252 edges.push(e.clone());
253 }
254 }
255 }
256 }
257 }
258
259 if let Some(p) = &args.predicate {
260 let mut adj: HashMap<String, Vec<String>> = HashMap::new();
261 for e in &edges {
262 if e["predicate"].as_str() != Some(p.as_str()) {
263 continue;
264 }
265 let src = e["src"].as_str().unwrap_or_default().to_owned();
266 let dst = e["dst"].as_str().unwrap_or_default().to_owned();
267 if dirs.contains(&"out") {
268 adj.entry(src.clone()).or_default().push(dst.clone());
269 }
270 if dirs.contains(&"in") {
271 adj.entry(dst).or_default().push(src);
272 }
273 }
274 let mut depth_of: Vec<(String, i64)> = root_ids.iter().map(|id| (id.clone(), 0)).collect();
275 let mut wave: Vec<String> = root_ids.clone();
276 let mut lvl = 1;
277 while lvl <= effective_degrees && !wave.is_empty() {
278 let mut next = Vec::new();
279 for from in &wave {
280 for to in adj.get(from).map_or(&[][..], Vec::as_slice) {
281 if docs.iter().any(|d| &d.id == to) && !depth_of.iter().any(|(id, _)| id == to)
282 {
283 depth_of.push((to.clone(), lvl));
284 next.push(to.clone());
285 }
286 }
287 }
288 wave = next;
289 lvl += 1;
290 }
291 let mut restricted: Vec<Doc> = Vec::new();
292 for (id, d) in depth_of {
293 if let Some(orig) = docs.iter().find(|x| x.id == id) {
294 restricted.push(Doc {
295 id: orig.id.clone(),
296 path: orig.path.clone(),
297 degree: d,
298 extra: orig.extra.clone(),
299 });
300 }
301 }
302 docs = restricted;
303 }
304
305 docs.sort_by(|a, b| {
306 a.degree
307 .cmp(&b.degree)
308 .then_with(|| locale_compare(&a.path, &b.path))
309 });
310 let capped = docs.len() > max_documents;
311 docs.truncate(max_documents);
312 let truncated = query_truncated || capped;
313
314 let mut frontier = Vec::new();
315 let documents: Vec<Json> = docs
316 .iter()
317 .map(|d| {
318 let is_frontier = d.degree == effective_degrees;
319 if is_frontier {
320 frontier.push(json!({ "id": d.id, "path": d.path, "degree": d.degree }));
321 }
322 let mut m = Map::new();
323 m.insert("id".to_owned(), json!(d.id));
324 m.insert("path".to_owned(), json!(d.path));
325 m.insert("degree".to_owned(), json!(d.degree));
326 m.insert("frontier".to_owned(), json!(is_frontier));
327 for (k, v) in &d.extra {
328 m.insert(k.clone(), v.clone());
329 }
330 Json::Object(m)
331 })
332 .collect();
333
334 let reached = |id: &str| docs.iter().any(|d| d.id == id);
335 let mut kept: Vec<Json> = edges
336 .into_iter()
337 .filter(|e| {
338 if let Some(p) = &args.predicate {
339 if e["predicate"].as_str() != Some(p.as_str()) {
340 return false;
341 }
342 }
343 let src_in = reached(e["src"].as_str().unwrap_or_default());
344 let dst_kind = e["dst_kind"].as_str().unwrap_or_default();
345 let dst_dangling =
346 dst_kind == "external" || (dst_kind == "document" && e["dst_path"].is_null());
347 let dst_in = reached(e["dst"].as_str().unwrap_or_default()) || dst_dangling;
348 src_in && dst_in
349 })
350 .collect();
351 kept.sort_by(|a, b| {
352 let (sa, sb) = (
353 a["src"].as_str().unwrap_or_default(),
354 b["src"].as_str().unwrap_or_default(),
355 );
356 if sa == sb {
357 locale_compare(
358 a["id"].as_str().unwrap_or_default(),
359 b["id"].as_str().unwrap_or_default(),
360 )
361 } else {
362 locale_compare(sa, sb)
363 }
364 });
365
366 Ok(json!({
367 "roots": root_ids,
368 "degrees": effective_degrees,
369 "direction": direction,
370 "documents": documents,
371 "edges": kept,
372 "frontier": frontier,
373 "truncated": truncated,
374 "queries": queries,
375 }))
376}
377
378#[cfg(test)]
379mod tests {
380 use super::*;
381
382 #[test]
383 fn user_select_aliases() {
384 let (items, names) = build_user_select(&[
385 "layer".into(),
386 "$path".into(),
387 "a + 1".into(),
388 "layer".into(),
389 ])
390 .unwrap();
391 let printed: Vec<String> = items
392 .iter()
393 .map(|it| oqx::print(oqx::Node::Select(it)).unwrap())
394 .collect();
395 assert_eq!(
396 printed,
397 ["_u0: layer", "_u1: $path", "_u2: a + 1", "_u3: layer"]
398 );
399 assert_eq!(
400 names,
401 [
402 Some("layer".into()),
403 Some("path".into()),
404 Some("sel_2".into()),
405 Some("layer_3".into())
406 ]
407 );
408 assert!(build_user_select(&[]).unwrap().0.is_empty());
409 let e = build_user_select(&["a +".into()]).unwrap_err();
410 assert_eq!(e.code, "filter_invalid");
411 }
412
413 #[test]
414 fn query_shape() {
415 let q = build_query(&["d_0".to_owned()], "out", 2, &[]);
416 assert!(
417 q.starts_with("select _depth: $depth, _stop: $stop, _edges: doc.out_edges collect {")
418 );
419 assert!(q.ends_with("from docs where $id == \"d_0\" follow distinct doc.out { depth 2 }"));
420 assert!(oqx::parse_string(&q).is_ok());
421 let q = build_query(&["d_14".to_owned(), "d_15".to_owned()], "in", 8, &[]);
425 assert_eq!(
426 q,
427 "select _depth: $depth, _stop: $stop, _edges: doc.in_edges collect { id: $id, src: $src, dst: $dst, dst_path: $dst_path, dst_uri: $dst_uri, dst_kind: dst_kind, predicate: predicate, provenance: provenance, anchor: anchor, src_field: src_field } from docs where $id == \"d_14\" or $id == \"d_15\" follow distinct doc.in { depth 8 }"
428 );
429 }
430
431 #[test]
432 fn order_is_bytewise() {
433 assert_eq!(locale_compare("B", "a"), Ordering::Less);
434 assert_eq!(locale_compare("a", "b"), Ordering::Less);
435 }
436}