diff --git a/.github/workflows/rust-test.yml b/.github/workflows/rust-test.yml index b868068e5..23cde336c 100644 --- a/.github/workflows/rust-test.yml +++ b/.github/workflows/rust-test.yml @@ -92,6 +92,10 @@ jobs: run: cargo test --manifest-path crates/lance-graph/Cargo.toml --no-run - name: Run unit tests run: cargo test --manifest-path crates/lance-graph/Cargo.toml --lib + - name: Run integration tests + # `--test '*'` selects only integration-test targets; `--tests` would + # re-run the lib unit tests the step above already ran. + run: cargo test --manifest-path crates/lance-graph/Cargo.toml --test '*' - name: Run doc tests run: cargo test --manifest-path crates/lance-graph/Cargo.toml --doc # lance-graph-contract is the zero-dep trait crate every workspace @@ -526,8 +530,10 @@ jobs: - name: Install cargo-llvm-cov uses: taiki-e/install-action@cargo-llvm-cov - name: Run tests with coverage + env: + CARGO_BUILD_JOBS: "2" run: | - cargo llvm-cov --manifest-path crates/lance-graph/Cargo.toml --lcov --output-path lcov.info + cargo llvm-cov --manifest-path crates/lance-graph/Cargo.toml --lib --lcov --output-path lcov.info - name: Upload coverage to Codecov uses: codecov/codecov-action@v4 with: diff --git a/AGENTS.md b/AGENTS.md index 1db62c087..941f1cb3b 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -2,13 +2,15 @@ ## Project Structure & Module Organization - `crates/lance-graph/` hosts the Rust Cypher engine; keep new modules under `src/` and co-locate helpers inside `query/` or feature-specific submodules. +- `crates/lance-graph-catalog/` provides catalog and namespace utilities, including Unity Catalog and Delta Lake integration; feature-gated via `unity-catalog`. +- `crates/lance-graph-benches/` contains performance benchmarks in a dedicated unpublished crate; add new benchmarks here rather than in the main crate. - `crates/lance-graph-python/src/` contains the PyO3 bridge; `python/python/lance_graph/` holds the pure-Python facade and packaging metadata. - `python/python/tests/` stores functional tests; mirror new features with targeted cases here and in the corresponding Rust module. - `examples/` demonstrates Cypher usage; update or add examples when introducing new public APIs. ## Build, Test, and Development Commands - `cargo check` / `cargo test --all` (run inside `crates/lance-graph`) validate Rust code paths. -- `cargo bench --bench graph_execution` measures performance-critical changes; include shortened runs with `--warm-up-time 1`. +- `cargo bench --bench graph_execution -p lance-graph-benches` measures performance-critical changes; include shortened runs with `--warm-up-time 1`. - `uv venv --python 3.11 .venv` and `uv pip install -e '.[tests]'` bootstrap the Python workspace. - `maturin develop` rebuilds the extension after Rust edits; `pytest python/python/tests/ -v` exercises Python bindings. - `make lint` (in `python/`) runs `ruff`, formatting checks, and `pyright`. diff --git a/crates/lance-graph/src/csr_index.rs b/crates/lance-graph/src/csr_index.rs new file mode 100644 index 000000000..188f410c9 --- /dev/null +++ b/crates/lance-graph/src/csr_index.rs @@ -0,0 +1,790 @@ +// SPDX-License-Identifier: Apache-2.0 +// SPDX-FileCopyrightText: Copyright The Lance Authors + +//! CSR (Compressed Sparse Row) adjacency index for native graph traversal +//! +//! Instead of translating graph traversals into SQL joins, this module provides +//! a CSR-based adjacency index that enables O(1) neighbor lookup. Inspired by +//! [GraphAr](https://graphar.apache.org)'s approach of encoding CSR offset +//! tables alongside columnar edge data. +//! +//! # Layout +//! +//! ```text +//! Offset Array: [0, 3, 5, 5, 9, ...] (one entry per vertex + 1 sentinel) +//! Neighbor Array: [2, 5, 7, 1, 4, 0, 3, 6, 8, ...] (destination vertex IDs) +//! ``` +//! +//! For vertex `v`, its neighbors are `neighbors[offsets[v]..offsets[v+1]]`. + +use arrow_array::{Array, RecordBatch, UInt64Array}; +use arrow_schema::{DataType, Field, Schema}; +use std::collections::HashMap; +use std::sync::Arc; + +use crate::error::{GraphError, Result}; + +/// In-memory CSR adjacency index for fast neighbor lookup. +/// +/// Stores graph topology in two arrays: +/// - `offsets[v]` = start position of vertex v's neighbors in the neighbor array +/// - `neighbors[offsets[v]..offsets[v+1]]` = destination vertex IDs +#[derive(Debug, Clone)] +pub struct CsrIndex { + offsets: Vec, + neighbors: Vec, + num_vertices: u64, +} + +impl CsrIndex { + /// Look up all neighbors of a vertex. Returns an empty slice for vertices + /// with no outgoing edges or vertex IDs beyond the index range. + pub fn neighbors(&self, vertex_id: u64) -> &[u64] { + let v = vertex_id as usize; + if v >= self.offsets.len() - 1 { + return &[]; + } + let start = self.offsets[v] as usize; + let end = self.offsets[v + 1] as usize; + &self.neighbors[start..end] + } + + /// Return the out-degree of a vertex (number of outgoing edges). + pub fn degree(&self, vertex_id: u64) -> u32 { + self.neighbors(vertex_id).len() as u32 + } + + /// Return the total number of vertices in the index. + pub fn num_vertices(&self) -> u64 { + self.num_vertices + } + + /// Return the total number of edges in the index. + pub fn num_edges(&self) -> u64 { + self.neighbors.len() as u64 + } + + /// Export the CSR index as an Arrow RecordBatch (offset table). + /// + /// Schema: `vertex_id: u64, offset: u64, degree: u64` + /// + /// This can be persisted as a Lance dataset for later loading. + pub fn to_record_batch(&self) -> Result { + let n = self.num_vertices as usize; + let vertex_ids: Vec = (0..n as u64).collect(); + let offsets: Vec = self.offsets[..n].to_vec(); + let degrees: Vec = (0..n) + .map(|i| self.offsets[i + 1] - self.offsets[i]) + .collect(); + + let schema = Arc::new(Schema::new(vec![ + Field::new("vertex_id", DataType::UInt64, false), + Field::new("offset", DataType::UInt64, false), + Field::new("degree", DataType::UInt64, false), + ])); + + RecordBatch::try_new( + schema, + vec![ + Arc::new(UInt64Array::from(vertex_ids)), + Arc::new(UInt64Array::from(offsets)), + Arc::new(UInt64Array::from(degrees)), + ], + ) + .map_err(|e| GraphError::PlanError { + message: format!("Failed to create CSR offset RecordBatch: {}", e), + location: snafu::Location::new(file!(), line!(), column!()), + }) + } + + /// Export the neighbor (destination) array as an Arrow RecordBatch. + /// + /// Schema: `dst_id: u64` + /// + /// This is the edge list sorted by source vertex, suitable for sequential scans. + pub fn neighbors_to_record_batch(&self) -> Result { + let schema = Arc::new(Schema::new(vec![Field::new( + "dst_id", + DataType::UInt64, + false, + )])); + + RecordBatch::try_new( + schema, + vec![Arc::new(UInt64Array::from(self.neighbors.clone()))], + ) + .map_err(|e| GraphError::PlanError { + message: format!("Failed to create CSR neighbors RecordBatch: {}", e), + location: snafu::Location::new(file!(), line!(), column!()), + }) + } + + /// Perform a k-hop BFS traversal from a starting vertex. + /// + /// Returns all vertices reachable within `max_hops` hops, organized by + /// distance. Each entry in the returned Vec is the set of vertices at + /// that hop distance (index 0 = starting vertex, index 1 = 1-hop neighbors, etc.). + pub fn bfs(&self, start: u64, max_hops: u32) -> Vec> { + let mut visited = vec![false; self.num_vertices as usize]; + let mut levels: Vec> = Vec::with_capacity(max_hops as usize + 1); + + if (start as usize) >= self.num_vertices as usize { + return levels; + } + + visited[start as usize] = true; + levels.push(vec![start]); + + for _ in 0..max_hops { + let frontier = levels.last().unwrap(); + let mut next_level = Vec::new(); + + for &v in frontier { + for &neighbor in self.neighbors(v) { + let n = neighbor as usize; + if n < visited.len() && !visited[n] { + visited[n] = true; + next_level.push(neighbor); + } + } + } + + if next_level.is_empty() { + break; + } + levels.push(next_level); + } + + levels + } + + /// Find the shortest path between two vertices using BFS. + /// + /// Returns `None` if no path exists, or `Some(path)` where path is the + /// sequence of vertex IDs from `start` to `end` (inclusive). + pub fn shortest_path(&self, start: u64, end: u64) -> Option> { + // Range check first: a vertex outside the index has no path, not even + // the trivial one to itself. + let n = self.num_vertices as usize; + if start as usize >= n || end as usize >= n { + return None; + } + + if start == end { + return Some(vec![start]); + } + + let mut visited = vec![false; n]; + let mut parent: Vec> = vec![None; n]; + let mut queue = std::collections::VecDeque::new(); + + visited[start as usize] = true; + queue.push_back(start); + + while let Some(current) = queue.pop_front() { + for &neighbor in self.neighbors(current) { + let ni = neighbor as usize; + if ni < n && !visited[ni] { + visited[ni] = true; + parent[ni] = Some(current); + + if neighbor == end { + let mut path = vec![end]; + let mut node = end; + while let Some(p) = parent[node as usize] { + path.push(p); + node = p; + } + path.reverse(); + return Some(path); + } + queue.push_back(neighbor); + } + } + } + + None + } +} + +/// Builder for constructing a CSR index from edge data. +/// +/// Accepts edges as (source, destination) pairs and builds the compressed +/// sparse row representation. +#[derive(Debug)] +pub struct CsrIndexBuilder { + edges: Vec<(u64, u64)>, + num_vertices: Option, +} + +impl CsrIndexBuilder { + pub fn new() -> Self { + Self { + edges: Vec::new(), + num_vertices: None, + } + } + + /// Set the number of vertices explicitly (a lower bound). If not set, it is + /// inferred from the maximum vertex ID seen in the edges. An edge whose + /// endpoint is `>= n` grows the vertex range at [`Self::build`] rather + /// than being stored where no offset can reach it. + pub fn with_num_vertices(mut self, n: u64) -> Self { + self.num_vertices = Some(n); + self + } + + /// Add a single directed edge from `src` to `dst`. + pub fn add_edge(mut self, src: u64, dst: u64) -> Self { + self.edges.push((src, dst)); + self + } + + /// Add edges from an Arrow RecordBatch with `src_id` and `dst_id` columns. + /// + /// A row whose `src_id` or `dst_id` is null contributes no edge, the same + /// way a null key joins nothing on the relational expand path. + pub fn add_edges_from_batch(mut self, batch: &RecordBatch) -> Result { + let src_col = batch + .column_by_name("src_id") + .ok_or_else(|| GraphError::PlanError { + message: "Edge batch missing 'src_id' column".to_string(), + location: snafu::Location::new(file!(), line!(), column!()), + })?; + let dst_col = batch + .column_by_name("dst_id") + .ok_or_else(|| GraphError::PlanError { + message: "Edge batch missing 'dst_id' column".to_string(), + location: snafu::Location::new(file!(), line!(), column!()), + })?; + + let src_array = src_col + .as_any() + .downcast_ref::() + .ok_or_else(|| GraphError::PlanError { + message: "src_id column must be UInt64".to_string(), + location: snafu::Location::new(file!(), line!(), column!()), + })?; + let dst_array = dst_col + .as_any() + .downcast_ref::() + .ok_or_else(|| GraphError::PlanError { + message: "dst_id column must be UInt64".to_string(), + location: snafu::Location::new(file!(), line!(), column!()), + })?; + + for i in 0..batch.num_rows() { + // `value(i)` ignores the validity bitmap, so a null slot would + // otherwise become a fabricated edge (usually to vertex 0). + if src_array.is_null(i) || dst_array.is_null(i) { + continue; + } + self.edges.push((src_array.value(i), dst_array.value(i))); + } + + Ok(self) + } + + /// Build the CSR index. + /// + /// Sorts edges by source vertex, then builds offset and neighbor arrays. + /// + /// # Panics + /// + /// If a vertex ID or the declared vertex count is too large for the + /// offset table to address; [`Self::try_build`] reports that as an error. + pub fn build(self) -> CsrIndex { + match self.try_build() { + Ok(index) => index, + Err(e) => panic!("{e}"), + } + } + + /// Build the CSR index, failing instead of overflowing when a vertex ID + /// (or the declared vertex count) is too large for the offset table — + /// `num_vertices + 1` offsets must fit in both `u64` and `usize`. + pub fn try_build(mut self) -> Result { + let unaddressable = || GraphError::PlanError { + message: "CSR vertex id too large: the offset table cannot address it".to_string(), + location: snafu::Location::new(file!(), line!(), column!()), + }; + + // Every endpoint must be addressable: otherwise an edge from a source + // past the range is stored after the last offset (counted, never + // reachable) and a destination past it names a nonexistent vertex. + let needed = match self.edges.iter().flat_map(|&(s, d)| [s, d]).max() { + Some(m) => m.checked_add(1).ok_or_else(unaddressable)?, + None => 0, + }; + let num_vertices = self.num_vertices.map_or(needed, |n| n.max(needed)); + let offsets_len = num_vertices + .checked_add(1) + .and_then(|len| usize::try_from(len).ok()) + .ok_or_else(unaddressable)?; + + // Sort by source vertex for CSR construction. Stable, so a source's + // neighbors keep insertion order and BFS / shortest_path tie-breaks + // are deterministic. + self.edges.sort_by_key(|&(src, _)| src); + + // Build offset and neighbor arrays + let mut offsets = vec![0u64; offsets_len]; + let mut neighbors = Vec::with_capacity(self.edges.len()); + + // Count degrees + let mut degree_map: HashMap = HashMap::new(); + for &(src, _) in &self.edges { + *degree_map.entry(src).or_insert(0) += 1; + } + + // Build prefix-sum offsets + let mut running = 0u64; + for v in 0..num_vertices { + offsets[v as usize] = running; + running += degree_map.get(&v).copied().unwrap_or(0); + } + offsets[num_vertices as usize] = running; + + // Fill neighbor array + for &(_, dst) in &self.edges { + neighbors.push(dst); + } + + Ok(CsrIndex { + offsets, + neighbors, + num_vertices, + }) + } +} + +impl Default for CsrIndexBuilder { + fn default() -> Self { + Self::new() + } +} + +/// Build both outgoing (CSR) and incoming (CSC) adjacency indices from edge data. +/// +/// Returns `(outgoing_index, incoming_index)`. +pub fn build_bidirectional_index(edges: &[(u64, u64)], num_vertices: u64) -> (CsrIndex, CsrIndex) { + let mut outgoing_builder = CsrIndexBuilder::new().with_num_vertices(num_vertices); + let mut incoming_builder = CsrIndexBuilder::new().with_num_vertices(num_vertices); + + for &(src, dst) in edges { + outgoing_builder = outgoing_builder.add_edge(src, dst); + incoming_builder = incoming_builder.add_edge(dst, src); + } + + (outgoing_builder.build(), incoming_builder.build()) +} + +#[cfg(test)] +mod tests { + use super::*; + + // Test graph: + // 0 -> 1, 2, 3 + // 1 -> 2 + // 2 -> 3 + // 3 -> (none) + fn sample_index() -> CsrIndex { + CsrIndexBuilder::new() + .with_num_vertices(4) + .add_edge(0, 1) + .add_edge(0, 2) + .add_edge(0, 3) + .add_edge(1, 2) + .add_edge(2, 3) + .build() + } + + #[test] + fn test_basic_neighbor_lookup() { + let idx = sample_index(); + + assert_eq!(idx.neighbors(0), &[1, 2, 3]); + assert_eq!(idx.neighbors(1), &[2]); + assert_eq!(idx.neighbors(2), &[3]); + assert_eq!(idx.neighbors(3), &[] as &[u64]); + } + + #[test] + fn test_degree() { + let idx = sample_index(); + assert_eq!(idx.degree(0), 3); + assert_eq!(idx.degree(1), 1); + assert_eq!(idx.degree(2), 1); + assert_eq!(idx.degree(3), 0); + } + + #[test] + fn test_metadata() { + let idx = sample_index(); + assert_eq!(idx.num_vertices(), 4); + assert_eq!(idx.num_edges(), 5); + } + + #[test] + fn test_out_of_range_vertex() { + let idx = sample_index(); + assert_eq!(idx.neighbors(99), &[] as &[u64]); + assert_eq!(idx.degree(99), 0); + } + + #[test] + fn test_empty_graph() { + let idx = CsrIndexBuilder::new().with_num_vertices(0).build(); + assert_eq!(idx.num_vertices(), 0); + assert_eq!(idx.num_edges(), 0); + assert_eq!(idx.neighbors(0), &[] as &[u64]); + } + + #[test] + fn test_isolated_vertices() { + let idx = CsrIndexBuilder::new() + .with_num_vertices(5) + .add_edge(1, 3) + .build(); + + assert_eq!(idx.neighbors(0), &[] as &[u64]); + assert_eq!(idx.neighbors(1), &[3]); + assert_eq!(idx.neighbors(2), &[] as &[u64]); + assert_eq!(idx.neighbors(3), &[] as &[u64]); + assert_eq!(idx.neighbors(4), &[] as &[u64]); + } + + #[test] + fn test_build_from_record_batch() { + let schema = Arc::new(Schema::new(vec![ + Field::new("src_id", DataType::UInt64, false), + Field::new("dst_id", DataType::UInt64, false), + ])); + let batch = RecordBatch::try_new( + schema, + vec![ + Arc::new(UInt64Array::from(vec![0, 0, 1, 2])), + Arc::new(UInt64Array::from(vec![1, 2, 2, 0])), + ], + ) + .unwrap(); + + let idx = CsrIndexBuilder::new() + .add_edges_from_batch(&batch) + .unwrap() + .build(); + + assert_eq!(idx.neighbors(0), &[1, 2]); + assert_eq!(idx.neighbors(1), &[2]); + assert_eq!(idx.neighbors(2), &[0]); + } + + #[test] + fn test_to_record_batch() { + let idx = sample_index(); + let batch = idx.to_record_batch().unwrap(); + + assert_eq!(batch.num_rows(), 4); + assert_eq!(batch.num_columns(), 3); + + let offsets = batch + .column_by_name("offset") + .unwrap() + .as_any() + .downcast_ref::() + .unwrap(); + assert_eq!(offsets.value(0), 0); + assert_eq!(offsets.value(1), 3); + assert_eq!(offsets.value(2), 4); + assert_eq!(offsets.value(3), 5); + + let degrees = batch + .column_by_name("degree") + .unwrap() + .as_any() + .downcast_ref::() + .unwrap(); + assert_eq!(degrees.value(0), 3); + assert_eq!(degrees.value(1), 1); + assert_eq!(degrees.value(2), 1); + assert_eq!(degrees.value(3), 0); + } + + #[test] + fn test_neighbors_to_record_batch() { + let idx = sample_index(); + let batch = idx.neighbors_to_record_batch().unwrap(); + + assert_eq!(batch.num_rows(), 5); + let dst = batch + .column_by_name("dst_id") + .unwrap() + .as_any() + .downcast_ref::() + .unwrap(); + assert_eq!(dst.value(0), 1); + assert_eq!(dst.value(1), 2); + assert_eq!(dst.value(2), 3); + assert_eq!(dst.value(3), 2); + assert_eq!(dst.value(4), 3); + } + + #[test] + fn test_bidirectional_index() { + let edges = vec![(0, 1), (0, 2), (1, 2)]; + let (outgoing, incoming) = build_bidirectional_index(&edges, 3); + + // Outgoing + assert_eq!(outgoing.neighbors(0), &[1, 2]); + assert_eq!(outgoing.neighbors(1), &[2]); + assert_eq!(outgoing.neighbors(2), &[] as &[u64]); + + // Incoming (reversed edges) + assert_eq!(incoming.neighbors(0), &[] as &[u64]); + assert_eq!(incoming.neighbors(1), &[0]); + assert_eq!(incoming.neighbors(2), &[0, 1]); + } + + #[test] + fn test_bfs_traversal() { + // Graph: 0->1, 0->2, 1->3, 2->3, 3->4 + let idx = CsrIndexBuilder::new() + .with_num_vertices(5) + .add_edge(0, 1) + .add_edge(0, 2) + .add_edge(1, 3) + .add_edge(2, 3) + .add_edge(3, 4) + .build(); + + let levels = idx.bfs(0, 3); + assert_eq!(levels.len(), 4); + assert_eq!(levels[0], vec![0]); // start + assert_eq!(levels[1], vec![1, 2]); // 1-hop + assert_eq!(levels[2], vec![3]); // 2-hop (3 reached from both 1 and 2, but visited once) + assert_eq!(levels[3], vec![4]); // 3-hop + } + + #[test] + fn test_bfs_limited_hops() { + let idx = CsrIndexBuilder::new() + .with_num_vertices(4) + .add_edge(0, 1) + .add_edge(1, 2) + .add_edge(2, 3) + .build(); + + let levels = idx.bfs(0, 1); + assert_eq!(levels.len(), 2); + assert_eq!(levels[0], vec![0]); + assert_eq!(levels[1], vec![1]); + } + + #[test] + fn test_bfs_disconnected() { + let idx = CsrIndexBuilder::new() + .with_num_vertices(4) + .add_edge(0, 1) + // 2, 3 are disconnected + .build(); + + let levels = idx.bfs(0, 10); + assert_eq!(levels.len(), 2); + assert_eq!(levels[0], vec![0]); + assert_eq!(levels[1], vec![1]); + } + + #[test] + fn test_bfs_invalid_start() { + let idx = CsrIndexBuilder::new().with_num_vertices(3).build(); + let levels = idx.bfs(99, 5); + assert!(levels.is_empty()); + } + + #[test] + fn test_shortest_path_direct() { + let idx = CsrIndexBuilder::new() + .with_num_vertices(3) + .add_edge(0, 1) + .add_edge(1, 2) + .add_edge(0, 2) + .build(); + + // Direct edge 0->2 is shorter than 0->1->2 + let path = idx.shortest_path(0, 2).unwrap(); + assert_eq!(path, vec![0, 2]); + } + + #[test] + fn test_shortest_path_multi_hop() { + let idx = CsrIndexBuilder::new() + .with_num_vertices(4) + .add_edge(0, 1) + .add_edge(1, 2) + .add_edge(2, 3) + .build(); + + let path = idx.shortest_path(0, 3).unwrap(); + assert_eq!(path, vec![0, 1, 2, 3]); + } + + #[test] + fn test_shortest_path_same_vertex() { + let idx = sample_index(); + let path = idx.shortest_path(2, 2).unwrap(); + assert_eq!(path, vec![2]); + } + + #[test] + fn test_shortest_path_unreachable() { + let idx = CsrIndexBuilder::new() + .with_num_vertices(3) + .add_edge(0, 1) + // no path from 0 to 2 + .build(); + + assert!(idx.shortest_path(0, 2).is_none()); + } + + #[test] + fn test_shortest_path_invalid_vertices() { + let idx = sample_index(); + assert!(idx.shortest_path(99, 0).is_none()); + assert!(idx.shortest_path(0, 99).is_none()); + } + + #[test] + fn test_auto_inferred_num_vertices() { + let idx = CsrIndexBuilder::new().add_edge(0, 5).add_edge(3, 7).build(); + + // Should infer num_vertices = 8 (max ID 7 + 1) + assert_eq!(idx.num_vertices(), 8); + assert_eq!(idx.neighbors(0), &[5]); + assert_eq!(idx.neighbors(3), &[7]); + } + + #[test] + fn test_self_loops() { + let idx = CsrIndexBuilder::new() + .with_num_vertices(3) + .add_edge(0, 0) + .add_edge(1, 1) + .add_edge(0, 1) + .build(); + + assert_eq!(idx.neighbors(0), &[0, 1]); + assert_eq!(idx.neighbors(1), &[1]); + } + + #[test] + fn test_parallel_edges() { + let idx = CsrIndexBuilder::new() + .with_num_vertices(2) + .add_edge(0, 1) + .add_edge(0, 1) + .add_edge(0, 1) + .build(); + + // CSR preserves multi-edges + assert_eq!(idx.neighbors(0), &[1, 1, 1]); + assert_eq!(idx.degree(0), 3); + } + /// FAILS IF: a null `src_id`/`dst_id` slot is read through `value(i)` + /// (which ignores validity) and becomes an edge — here it would be an + /// edge from or to vertex 0, which the valid rows never touch. + #[test] + fn test_null_endpoints_contribute_no_edge() { + let schema = Arc::new(Schema::new(vec![ + Field::new("src_id", DataType::UInt64, true), + Field::new("dst_id", DataType::UInt64, true), + ])); + let batch = RecordBatch::try_new( + schema, + vec![ + Arc::new(UInt64Array::from(vec![Some(1), None, Some(2)])), + Arc::new(UInt64Array::from(vec![Some(2), Some(1), None])), + ], + ) + .unwrap(); + + let idx = CsrIndexBuilder::new() + .add_edges_from_batch(&batch) + .unwrap() + .build(); + + assert_eq!(idx.num_edges(), 1); + assert_eq!(idx.neighbors(1), &[2]); + assert_eq!(idx.neighbors(0), &[] as &[u64]); + assert_eq!(idx.neighbors(2), &[] as &[u64]); + } + + /// FAILS IF: a declared vertex count below an endpoint truncates the + /// offsets while the edge is still stored — the edge from vertex 5 would + /// be counted in `num_edges` but unreachable through `neighbors(5)`. + #[test] + fn test_endpoints_past_declared_range_stay_reachable() { + let idx = CsrIndexBuilder::new() + .with_num_vertices(2) + .add_edge(0, 1) + .add_edge(5, 3) + .build(); + + assert_eq!(idx.num_vertices(), 6); + assert_eq!(idx.neighbors(5), &[3]); + assert_eq!(idx.neighbors(0), &[1]); + let reachable: u64 = (0..idx.num_vertices()).map(|v| idx.degree(v) as u64).sum(); + assert_eq!(reachable, idx.num_edges()); + + let (out, inc) = build_bidirectional_index(&[(0, 1), (5, 3)], 2); + assert_eq!(out.neighbors(5), &[3]); + assert_eq!(inc.neighbors(3), &[5]); + } + + /// FAILS IF: the source sort is unstable — a source's neighbors then come + /// back in an arbitrary order instead of insertion order. The input is + /// large and interleaved on purpose: a short slice is insertion-sorted + /// (stable by accident) and would pass either way. + #[test] + fn test_neighbors_keep_insertion_order_per_source() { + let mut b = CsrIndexBuilder::new(); + for i in 0..600u64 { + // Sources cycle 2, 1, 0 so every source's edges are scattered + // across the input; destinations descend so insertion order is + // not the sorted order either. + b = b.add_edge(2 - (i % 3), 10_000 - i); + } + let idx = b.build(); + for src in 0..3u64 { + let expected: Vec = (0..600u64) + .filter(|i| 2 - (i % 3) == src) + .map(|i| 10_000 - i) + .collect(); + assert_eq!(idx.neighbors(src), expected.as_slice(), "source {src}"); + } + } + + /// FAILS IF: `max endpoint + 1` (or `num_vertices + 1`) is computed + /// unchecked — it wraps in release and panics in debug instead of being + /// reported. Only checks that the error is returned, never allocates. + #[test] + fn test_unaddressable_vertex_ids_are_an_error() { + let max_endpoint = CsrIndexBuilder::new().add_edge(0, u64::MAX).try_build(); + assert!(max_endpoint.is_err()); + let max_source = CsrIndexBuilder::new().add_edge(u64::MAX, 0).try_build(); + assert!(max_source.is_err()); + // A declared count of u64::MAX fits, but its `count + 1` offsets do + // not. + let declared = CsrIndexBuilder::new() + .with_num_vertices(u64::MAX) + .try_build(); + assert!(declared.is_err()); + } + + /// FAILS IF: the `start == end` shortcut runs before the range check. + #[test] + fn test_shortest_path_equal_endpoints_out_of_range() { + let idx = sample_index(); + assert!(idx.shortest_path(99, 99).is_none()); + assert_eq!(idx.shortest_path(2, 2), Some(vec![2])); + } +} diff --git a/crates/lance-graph/src/lib.rs b/crates/lance-graph/src/lib.rs index 7dfd809b1..8f6dfd978 100644 --- a/crates/lance-graph/src/lib.rs +++ b/crates/lance-graph/src/lib.rs @@ -39,6 +39,7 @@ pub mod ast; pub mod cam_pq; pub mod case_insensitive; pub mod config; +pub mod csr_index; pub mod datafusion_planner; pub mod dev_s3_env; pub mod error; @@ -66,6 +67,7 @@ pub mod table_readers; pub const MAX_VARIABLE_LENGTH_HOPS: u32 = 20; pub use config::{GraphConfig, NodeMapping, RelationshipMapping}; +pub use csr_index::{build_bidirectional_index, CsrIndex, CsrIndexBuilder}; pub use error::{GraphError, Result}; pub use lance_graph_catalog::{ DirNamespace, GraphSourceCatalog, InMemoryCatalog, SimpleTableSource, diff --git a/python/python/knowledge_graph/__init__.py b/python/python/knowledge_graph/__init__.py index c1669836c..fbf640232 100644 --- a/python/python/knowledge_graph/__init__.py +++ b/python/python/knowledge_graph/__init__.py @@ -13,6 +13,8 @@ except ImportError: # pragma: no cover - builder is available in normal installs. GraphConfigBuilder = object # type: ignore[assignment] +from lance_graph import DistanceMetric, VectorSearch + from .component import KnowledgeGraphComponent from .config import KnowledgeGraphConfig, build_graph_config_from_mapping from .extraction import ( @@ -66,6 +68,23 @@ def run( ) return query.execute(sources) + def run_with_vector_rerank( + self, + statement: str, + vector_search: "VectorSearch", + *, + datasets: Optional[TableMapping] = None, + ) -> pa.Table: + """Execute a Cypher statement and rerank results by vector similarity.""" + + query = CypherQuery(statement).with_config(self.config) + sources: Dict[str, pa.Table] = dict(self._tables) + if datasets: + sources.update( + {name: _ensure_table(name, table) for name, table in datasets.items()} + ) + return query.execute_with_vector_rerank(sources, vector_search) + def tables(self) -> Dict[str, pa.Table]: """Return a shallow copy of the registered datasets.""" return dict(self._tables) @@ -129,4 +148,6 @@ def build(self) -> KnowledgeGraph: "preview_extraction", "HeuristicExtractor", "LLMExtractor", + "VectorSearch", + "DistanceMetric", ] diff --git a/python/python/knowledge_graph/component.py b/python/python/knowledge_graph/component.py index ecc1460c0..b895c34c4 100644 --- a/python/python/knowledge_graph/component.py +++ b/python/python/knowledge_graph/component.py @@ -2,17 +2,20 @@ from __future__ import annotations -from typing import Any, Dict, List, Optional +import logging +from typing import Any, Dict, List, Literal, Optional import pyarrow as pa import yaml from fastapi import APIRouter, HTTPException -from pydantic import BaseModel +from pydantic import BaseModel, Field from .config import KnowledgeGraphConfig from .service import LanceKnowledgeGraph from .store import LanceGraphStore +LOGGER = logging.getLogger(__name__) + class QueryRequest(BaseModel): query: str @@ -23,6 +26,48 @@ class QueryResponse(BaseModel): row_count: int +class VectorQueryRequest(BaseModel): + """Request body for vector-reranked Cypher queries. + + Supply either ``vector`` (raw floats) or ``query_text`` (auto-embedded via OpenAI). + """ + + query: str = Field(..., description="Cypher statement to execute.") + column: str = Field(..., description="Name of the vector column to search.") + + # Choose one: pass the vector directly or pass the text + # and let the server automatically embed it + vector: Optional[List[float]] = Field( + None, + description="Query vector (float list). Mutually exclusive with query_text.", + ) + query_text: Optional[str] = Field( + None, + min_length=1, + description="Text to embed as query vector. Requires OpenAI API key.", + ) + + metric: Literal["cosine", "l2", "dot"] = Field( + "cosine", description="Distance metric: cosine | l2 | dot." + ) + top_k: int = Field(10, ge=1, le=10000, description="Number of nearest neighbours.") + include_distance: bool = Field( + True, description="Include _distance column in results." + ) + embedding_model: str = Field( + "text-embedding-3-small", + description="OpenAI embedding model (only used when query_text is provided).", + ) + + +class VectorQueryResponse(BaseModel): + rows: List[Dict[str, Any]] + row_count: int + column: str + metric: str + top_k: int + + class DatasetUpsertRequest(BaseModel): records: List[Dict[str, Any]] merge: bool = True @@ -98,6 +143,76 @@ async def get_schema() -> Dict[str, Any]: payload = yaml.safe_load(handle) or {} return {"path": str(schema_path), "schema": payload} + @self.router.post("/query/vector", response_model=VectorQueryResponse) + async def execute_vector_query( + request: VectorQueryRequest, + ) -> VectorQueryResponse: + """Execute a Cypher query with vector similarity reranking. + + Supply ``vector`` (raw floats) or ``query_text`` (auto-embedded). + """ + if request.vector is None and request.query_text is None: + raise HTTPException( + status_code=400, + detail="Either 'vector' or 'query_text' must be provided.", + ) + if request.vector is not None and request.query_text is not None: + raise HTTPException( + status_code=400, + detail="Provide only one of 'vector' or 'query_text', not both.", + ) + + service = self._get_service() + + try: + if request.query_text is not None: + # Text: service internally calls EmbeddingGenerator + result = service.query_by_text( + request.query, + request.query_text, + request.column, + top_k=request.top_k, + metric=request.metric, + include_distance=request.include_distance, + embedding_model=request.embedding_model, + ) + else: + # Vector: Constructing VectorSearch directly + from lance_graph import DistanceMetric, VectorSearch + + _metric_map = { + "cosine": DistanceMetric.Cosine, + "l2": DistanceMetric.L2, + "dot": DistanceMetric.Dot, + } + vs = ( + VectorSearch(request.column) + .query_vector(request.vector) + .metric(_metric_map[request.metric]) + .top_k(request.top_k) + .include_distance(request.include_distance) + ) + result = service.run_with_vector_rerank(request.query, vs) + + except RuntimeError as exc: + # The message can carry embedding-client and engine internals + # (and the caller's query text): log it, return a generic one. + LOGGER.exception("Vector query failed") + raise HTTPException( + status_code=500, detail="Vector query execution failed." + ) from exc + except ValueError as exc: + raise HTTPException(status_code=400, detail=str(exc)) from exc + + rows = result.to_pylist() + return VectorQueryResponse( + rows=rows, + row_count=len(rows), + column=request.column, + metric=request.metric, + top_k=request.top_k, + ) + def close(self) -> None: """Release retained resources.""" self._service = None diff --git a/python/python/knowledge_graph/service.py b/python/python/knowledge_graph/service.py index a0e53f634..3d9678a5e 100644 --- a/python/python/knowledge_graph/service.py +++ b/python/python/knowledge_graph/service.py @@ -4,7 +4,7 @@ from typing import TYPE_CHECKING, Iterable, Mapping, MutableMapping, Optional -from lance_graph import CypherQuery, GraphConfig +from lance_graph import CypherQuery, DistanceMetric, GraphConfig, VectorSearch from .config import KnowledgeGraphConfig, build_default_graph_config from .store import LanceGraphStore @@ -141,6 +141,94 @@ def query( """Alias for :meth:`run` to match the semantic service naming.""" return self.run(statement, datasets=datasets) + def run_with_vector_rerank( + self, + statement: str, + vector_search: "VectorSearch", + *, + datasets: Optional[Mapping[str, "pa.Table"]] = None, + ) -> "pa.Table": + """Execute a Cypher statement and rerank results by vector similarity. + + Parameters + ---------- + statement: + Cypher query string. + vector_search: + A configured ``VectorSearch`` instance (column, vector, metric, top_k). + datasets: + Optional override tables injected on top of persisted datasets. + """ + query = CypherQuery(statement).with_config(self._config) + + referenced_tables = set(query.node_labels()) | set(query.relationship_types()) + base_tables: MutableMapping[str, "pa.Table"] = dict( + self._store.load_tables(referenced_tables) + ) + if datasets: + base_tables.update(datasets) + return query.execute_with_vector_rerank(base_tables, vector_search) + + def query_by_text( + self, + statement: str, + query_text: str, + column: str, + *, + top_k: int = 10, + metric: str = "cosine", + include_distance: bool = True, + embedding_model: str = "text-embedding-3-small", + datasets: Optional[Mapping[str, "pa.Table"]] = None, + ) -> "pa.Table": + """Convenience method: embed ``query_text`` then call run_with_vector_rerank. + + Parameters + ---------- + statement: + Cypher query string. + query_text: + Natural-language text to embed as the query vector. + column: + Name of the vector column in the dataset. + top_k: + Number of nearest neighbours to return. + metric: + Distance metric: "cosine", "l2", or "dot". + include_distance: + Whether to include the ``_distance`` column in results. + embedding_model: + OpenAI embedding model name. + datasets: + Optional override tables. + """ + from .embeddings import EmbeddingGenerator + + _metric_map = { + "cosine": DistanceMetric.Cosine, + "l2": DistanceMetric.L2, + "dot": DistanceMetric.Dot, + } + try: + rust_metric = _metric_map[metric.lower()] + except KeyError as exc: + raise ValueError( + f"Unsupported metric {metric!r}; expected one of {sorted(_metric_map)}" + ) from exc + + vector = EmbeddingGenerator(model=embedding_model).embed_one(query_text) + if vector is None: + raise RuntimeError(f"Failed to generate embedding for text: {query_text!r}") + + vs = ( + VectorSearch(column) + .query_vector(vector) + .metric(rust_metric) + .top_k(top_k) + .include_distance(include_distance) + ) + return self.run_with_vector_rerank(statement, vs, datasets=datasets) + def create_default_service( config: Optional[KnowledgeGraphConfig] = None, diff --git a/python/python/tests/test_vector_query_api.py b/python/python/tests/test_vector_query_api.py new file mode 100644 index 000000000..d756dd9f7 --- /dev/null +++ b/python/python/tests/test_vector_query_api.py @@ -0,0 +1,272 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright The Lance Authors + +"""Functional tests for the ``POST /query/vector`` route. + +The raw-vector path runs end to end: a real Lance-backed store in a temp +directory, the real Cypher engine, the real vector rerank. Only the OpenAI +embedding call on the ``query_text`` path is replaced, because it needs a +network and an API key. +""" + +import pyarrow as pa +import pytest +from fastapi import FastAPI +from fastapi.testclient import TestClient +from knowledge_graph import service as service_module +from knowledge_graph.component import KnowledgeGraphComponent +from knowledge_graph.config import KnowledgeGraphConfig +from knowledge_graph.service import LanceKnowledgeGraph +from knowledge_graph.store import LanceGraphStore +from lance_graph import GraphConfig + +QUERY = "MATCH (d:Document) RETURN d.name, d.embedding" + + +@pytest.fixture +def client(tmp_path): + kg_config = KnowledgeGraphConfig( + storage_path=tmp_path / "storage", + schema_path=tmp_path / "graph.yaml", + ) + store = LanceGraphStore(kg_config) + store.ensure_layout() + # Doc1 and Doc2 point the same way; Doc3 is orthogonal to them; Doc4 is + # anti-parallel to Doc1, so cosine and dot order it LAST while L2 still + # ranks it by raw distance — the metrics disagree, which is what lets a + # test see which one the route actually applied. + store.write_tables( + { + "Document": pa.table( + { + "id": [1, 2, 3, 4], + "name": ["Doc1", "Doc2", "Doc3", "Doc4"], + "embedding": pa.array( + [ + [1.0, 0.0, 0.0], + [4.0, 0.2, 0.0], + [0.0, 1.0, 0.0], + [-1.0, 0.0, 0.0], + ], + type=pa.list_(pa.float32()), + ), + } + ) + } + ) + graph_config = GraphConfig.builder().with_node_label("Document", "id").build() + + component = KnowledgeGraphComponent(kg_config) + component._service = LanceKnowledgeGraph(graph_config, storage=store) + app = FastAPI() + app.include_router(component.router) + test_client = TestClient(app) + test_client.kg_service = component._service + return test_client + + +def _service_of(client): + return client.kg_service + + +def _names(response): + return [row["d.name"] for row in response.json()["rows"]] + + +def test_raw_vector_returns_top_k_nearest(client): + response = client.post( + "/query/vector", + json={ + "query": QUERY, + "column": "d.embedding", + "vector": [1.0, 0.0, 0.0], + "metric": "cosine", + "top_k": 2, + }, + ) + + assert response.status_code == 200, response.text + body = response.json() + assert body["row_count"] == 2 + assert _names(response) == ["Doc1", "Doc2"] + assert body["column"] == "d.embedding" + assert body["metric"] == "cosine" + assert body["top_k"] == 2 + assert "_distance" in body["rows"][0] + + +def test_metric_selects_the_distance_function(client): + """L2 and cosine rank Doc2 differently: it points the same way as the + query (cosine-nearest) but is far away in raw distance (L2).""" + + def top(metric): + response = client.post( + "/query/vector", + json={ + "query": QUERY, + "column": "d.embedding", + "vector": [1.0, 0.0, 0.0], + "metric": metric, + "top_k": 4, + }, + ) + assert response.status_code == 200, response.text + return _names(response) + + cosine = top("cosine") + l2 = top("l2") + assert cosine[:2] == ["Doc1", "Doc2"] + # Doc2 sits at distance ~3.0 from the query, farther than Doc3 (~1.41). + assert l2.index("Doc3") < l2.index("Doc2") + assert cosine != l2 + # Dot product favours the long, aligned Doc2 over the unit-length Doc1. + assert top("dot")[0] == "Doc2" + + +def test_include_distance_false_drops_the_column(client): + response = client.post( + "/query/vector", + json={ + "query": QUERY, + "column": "d.embedding", + "vector": [1.0, 0.0, 0.0], + "top_k": 1, + "include_distance": False, + }, + ) + + assert response.status_code == 200, response.text + assert "_distance" not in response.json()["rows"][0] + + +def test_neither_vector_nor_text_is_rejected(client): + response = client.post( + "/query/vector", json={"query": QUERY, "column": "d.embedding"} + ) + assert response.status_code == 400 + assert "Either 'vector' or 'query_text'" in response.json()["detail"] + + +def test_both_vector_and_text_is_rejected(client): + response = client.post( + "/query/vector", + json={ + "query": QUERY, + "column": "d.embedding", + "vector": [1.0, 0.0, 0.0], + "query_text": "anything", + }, + ) + assert response.status_code == 400 + assert "only one of" in response.json()["detail"] + + +@pytest.mark.parametrize( + "field, value", + [("metric", "manhattan"), ("top_k", 0), ("top_k", 10001)], +) +def test_out_of_contract_fields_fail_validation(client, field, value): + payload = { + "query": QUERY, + "column": "d.embedding", + "vector": [1.0, 0.0, 0.0], + field: value, + } + assert client.post("/query/vector", json=payload).status_code == 422 + + +def test_query_text_is_embedded_then_reranked(client, monkeypatch): + """The text path must embed with the requested model and rank by the + resulting vector — the stub maps the text onto Doc3's direction, so a + route that ignored the embedding would not put Doc3 first.""" + seen = {} + + class FakeEmbeddingGenerator: + def __init__(self, model): + seen["model"] = model + + def embed_one(self, text): + seen["text"] = text + return [0.0, 1.0, 0.0] + + import knowledge_graph.embeddings as embeddings + + monkeypatch.setattr(embeddings, "EmbeddingGenerator", FakeEmbeddingGenerator) + + response = client.post( + "/query/vector", + json={ + "query": QUERY, + "column": "d.embedding", + "query_text": "science", + "top_k": 1, + "embedding_model": "text-embedding-3-large", + }, + ) + + assert response.status_code == 200, response.text + assert _names(response) == ["Doc3"] + assert seen == {"model": "text-embedding-3-large", "text": "science"} + + +def test_failed_embedding_is_a_generic_server_error(client, monkeypatch): + """A failure inside the service is a 500 whose body names no internals: + the raw message quotes the caller's text and the embedding client.""" + + class NoEmbedding: + def __init__(self, model): + pass + + def embed_one(self, text): + return None + + import knowledge_graph.embeddings as embeddings + + monkeypatch.setattr(embeddings, "EmbeddingGenerator", NoEmbedding) + + response = client.post( + "/query/vector", + json={"query": QUERY, "column": "d.embedding", "query_text": "secret text"}, + ) + assert response.status_code == 500 + detail = response.json()["detail"] + assert detail == "Vector query execution failed." + assert "secret text" not in detail + + +def test_empty_query_text_is_a_client_error(client): + response = client.post( + "/query/vector", + json={"query": QUERY, "column": "d.embedding", "query_text": ""}, + ) + assert response.status_code == 422 + + +def test_query_by_text_rejects_an_unknown_metric(client, monkeypatch): + """The service must refuse a metric it cannot honour rather than rank by + cosine without saying so.""" + + class ReachedEmbedding(Exception): + pass + + class Unused: + def __init__(self, model): + raise ReachedEmbedding + + import knowledge_graph.embeddings as embeddings + + monkeypatch.setattr(embeddings, "EmbeddingGenerator", Unused) + kg = _service_of(client) + with pytest.raises(ValueError, match="Unsupported metric 'euclidean'"): + kg.query_by_text(QUERY, "text", "d.embedding", metric="euclidean") + # The three documented names are accepted case-insensitively: "L2" gets + # past the metric check and on to the embedding step. + with pytest.raises(ReachedEmbedding): + kg.query_by_text(QUERY, "text", "d.embedding", metric="L2") + + +def test_service_module_exposes_the_rerank_entry_points(): + # Guards the import the route relies on: the component calls these two + # methods by name, so a rename would otherwise surface only as a 500. + assert hasattr(service_module.LanceKnowledgeGraph, "run_with_vector_rerank") + assert hasattr(service_module.LanceKnowledgeGraph, "query_by_text")