diff --git a/.github/codeql/suppressions.json b/.github/codeql/suppressions.json index ea1f2937a..7f60cea6d 100644 --- a/.github/codeql/suppressions.json +++ b/.github/codeql/suppressions.json @@ -39,6 +39,24 @@ "path": "nodedb/src/control/server/shared/ddl/neutral/timeseries/rewrite.rs", "sink": "^\\s*let partition_dir = ts_base\\.join\\(dir_name\\);", "reason": "dir_name is validated with is_plain_path_component immediately above this join (rejects / \\ : NUL . .. leading/trailing dot-space and control chars); a failing name is skipped with a warning before any path is built, so the join stays one level under ts_base." + }, + { + "rule": "rust/cleartext-logging", + "path": "nodedb-client/src/pg_cell/text.rs", + "sink": "^\\s*Ok\\(v\\) => panic!\\(\"\\{text:\\?\\} as \\{ty\\} must be refused, got \\{v:\\?\\}\"\\),", + "reason": "Test-only: a panic inside #[cfg(test)] mod tests prints a decoded test fixture when a decode is wrongly accepted. The value is a hard-coded literal, never user data, and the panic goes to test output, not a log." + }, + { + "rule": "rust/cleartext-logging", + "path": "nodedb/src/control/server/response_shape/compose/kernel.rs", + "sink": "^\\s*panic!\\(\"the whole row is an object: \\{:\\?\\}\", shaped\\.rows\\[0\\]\\);", + "reason": "Test-only: a panic inside #[cfg(test)] mod tests prints a row built from hard-coded fixture values when the shape assertion fails. No user data reaches it, and the panic goes to test output, not a log." + }, + { + "rule": "rust/cleartext-logging", + "path": "nodedb/src/data/executor/strict_format/coerce.rs", + "sink": "^\\s*panic!\\(\"(\\{input\\}: expected a decimal|expected an instant), got \\{got:\\?\\}\"\\);", + "reason": "Test-only: panics inside #[cfg(test)] mod tests print a coerced fixture value when a coercion returns the wrong variant. The inputs are hard-coded literals, never user data, and the panic goes to test output, not a log." } ] } diff --git a/Cargo.lock b/Cargo.lock index da0d071b5..55f8abdb3 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -4343,6 +4343,7 @@ version = "0.5.0" dependencies = [ "async-trait", "nodedb-types", + "rust_decimal", "rustls-pemfile", "serde", "serde_json", @@ -4363,6 +4364,8 @@ dependencies = [ "nodedb-client", "nodedb-test-support", "nodedb-types", + "serde_json", + "sonic-rs", "tempfile", "tokio", "tokio-postgres", @@ -4464,6 +4467,7 @@ name = "nodedb-columnar" version = "0.5.0" dependencies = [ "crc32c", + "faultbox", "nodedb-codec", "nodedb-mem", "nodedb-types", @@ -4474,6 +4478,7 @@ dependencies = [ "serde_json", "sonic-rs", "thiserror 2.0.20", + "ulid", "uuid", "zerompk", ] diff --git a/docs/graph.md b/docs/graph.md index 08e8e99b9..244bf0fae 100644 --- a/docs/graph.md +++ b/docs/graph.md @@ -74,13 +74,28 @@ GRAPH TRAVERSE FROM 'users:alice' DEPTH 3; GRAPH TRAVERSE FROM 'users:alice' DEPTH 2 LABEL 'follows' DIRECTION out; ``` -Breadth-first search from a start node. Returns discovered nodes at each depth level. +Breadth-first search from a start node. Returns the discovered nodes with their depth, and the crossed edges with their properties: `{"nodes": [{"id", "depth"}], "edges": [{"from", "to", "label", "properties"}]}`. Nodes at the last depth are not expanded. Their edges to nodes already in the result are included. -| Parameter | Default | Description | -| ----------- | ------- | ---------------------- | -| `DEPTH` | 2 | Maximum hop count | -| `LABEL` | (any) | Filter by edge label | -| `DIRECTION` | out | `in`, `out`, or `both` | +| Parameter | Default | Description | +| ------------ | ------- | --------------------------------------------- | +| `DEPTH` | 2 | Maximum hop count | +| `LABEL` | (any) | Follow edges with any listed label | +| `DIRECTION` | out | `in`, `out`, or `both` | +| `EDGE WHERE` | (none) | Cross only edges whose properties match; last | + +### Edge Property Predicate + +```sql +GRAPH TRAVERSE FROM 'users:alice' DEPTH 2 LABEL 'follows' EDGE WHERE since >= 2020 AND active = TRUE; +GRAPH PATH FROM 'a' TO 'z' MAX_DEPTH 6 EDGE WHERE "kind" IN ('road', 'rail') AND NOT (closed = TRUE); +``` + +- `EDGE WHERE` is the last clause. Only `GRAPH TRAVERSE` and `GRAPH PATH` accept it. +- Terms: `=`, `<>`, `!=`, `>`, `>=`, `<`, `<=`, `IN`, `NOT IN`, `IS NULL`, `IS NOT NULL`, joined by `AND`, `OR`, `NOT`. +- Each comparison names one property and one literal: `NULL`, `TRUE`, `FALSE`, a quoted string, an integer or a finite float. +- A missing property matches `IS NULL` and `NOT IN`. It never matches `>`, `>=`, `<`, `<=` or `IN`. +- An edge without properties evaluates as an empty object. +- The predicate runs on the core that stores the edge, before the edge counts against the visit cap. ### Neighbors (1-Hop) @@ -98,7 +113,7 @@ GRAPH PATH FROM 'users:alice' TO 'users:charlie'; GRAPH PATH FROM 'users:alice' TO 'users:charlie' MAX_DEPTH 5 LABEL 'knows'; ``` -Cross-core BFS path finding. Returns an ordered list of node IDs, or empty array if no path exists within `MAX_DEPTH` (default 10). +Cross-core BFS path finding. Returns an ordered list of node IDs, or empty array if no path exists within `MAX_DEPTH` (default 10). `EDGE WHERE` restricts the path to edges whose properties match, tested in each edge's stored direction. --- diff --git a/docs/query-language.md b/docs/query-language.md index bb264bab9..711f09d87 100644 --- a/docs/query-language.md +++ b/docs/query-language.md @@ -857,11 +857,15 @@ GRAPH INSERT EDGE IN 'edges' FROM 'alice' TO 'bob' TYPE 'knows' PROPERTIES { sin -- Traversal GRAPH TRAVERSE FROM 'users:alice' DEPTH 3 LABEL 'follows' DIRECTION out; +-- Traversal over edges whose properties match (EDGE WHERE is the last clause) +GRAPH TRAVERSE FROM 'users:alice' DEPTH 2 LABEL 'follows', 'knows' EDGE WHERE since >= 2020; + -- Neighbors GRAPH NEIGHBORS OF 'users:alice' LABEL 'follows' DIRECTION out; -- Shortest path GRAPH PATH FROM 'users:alice' TO 'users:carol' MAX_DEPTH 5; +GRAPH PATH FROM 'users:alice' TO 'users:carol' MAX_DEPTH 5 EDGE WHERE "weight" > 0.5; -- Pattern matching (Cypher subset) MATCH (u:User)-[follows]->(other:User) diff --git a/fuzz/src/targets/strict_tuple.rs b/fuzz/src/targets/strict_tuple.rs index 433ea9255..304101a13 100644 --- a/fuzz/src/targets/strict_tuple.rs +++ b/fuzz/src/targets/strict_tuple.rs @@ -38,10 +38,7 @@ fn schema_from_descriptor(descriptor: &[u8]) -> StrictSchema { 4 => ColumnType::Bytes, 5 => ColumnType::Timestamp, 6 => ColumnType::Timestamptz, - 7 => ColumnType::Decimal { - precision: 38, - scale: 10, - }, + 7 => ColumnType::Decimal(None), 8 => ColumnType::Uuid, 9 => ColumnType::Vector( descriptor @@ -147,10 +144,7 @@ fn all_supported_schema() -> StrictSchema { ColumnType::Timestamp, ColumnType::Timestamptz, ColumnType::SystemTimestamp, - ColumnType::Decimal { - precision: 38, - scale: 10, - }, + ColumnType::Decimal(None), ColumnType::Uuid, ColumnType::Vector(1), ColumnType::SparseVector, diff --git a/nodedb-client-tests/Cargo.toml b/nodedb-client-tests/Cargo.toml index a43714e29..22cc6923d 100644 --- a/nodedb-client-tests/Cargo.toml +++ b/nodedb-client-tests/Cargo.toml @@ -21,3 +21,5 @@ nodedb-types = { workspace = true } tokio = { workspace = true, features = ["test-util", "macros", "rt-multi-thread"] } tokio-postgres = { workspace = true } tempfile = { workspace = true } +serde_json = { workspace = true } +sonic-rs = { workspace = true } diff --git a/nodedb-client-tests/tests/document_declared_key_round_trip.rs b/nodedb-client-tests/tests/document_declared_key_round_trip.rs new file mode 100644 index 000000000..87c575274 --- /dev/null +++ b/nodedb-client-tests/tests/document_declared_key_round_trip.rs @@ -0,0 +1,230 @@ +// SPDX-License-Identifier: BUSL-1.1 + +//! End-to-end tests for a collection whose declared primary key is not `id`. +//! +//! The key column is the identity column: every write stores the document id +//! under it, and every key read and key restriction names it. Both clients +//! resolve it from the server catalog, so a document put, read, or deleted +//! through either client is the same row SQL sees under that key. + +use std::collections::HashSet; + +use nodedb_client::native::pool::PoolConfig; +use nodedb_client::{Document, NativeClient, NodeDb, NodeDbRemote, SearchResult, Value}; +use nodedb_test_support::pgwire_harness::TestServer; + +async fn remote(server: &TestServer) -> NodeDbRemote { + NodeDbRemote::connect(&format!( + "host=127.0.0.1 port={} user=nodedb dbname=default", + server.pg_port + )) + .await + .expect("pgwire connect to harness must succeed") +} + +fn native(server: &TestServer) -> NativeClient { + NativeClient::new(PoolConfig::new( + format!("127.0.0.1:{}", server.native_port), + nodedb_types::protocol::AuthMethod::Trust { + username: "nodedb".into(), + }, + )) +} + +async fn sql(remote: &NodeDbRemote, statement: &str) { + remote + .execute_sql(statement, &[]) + .await + .unwrap_or_else(|e| panic!("{statement}: {e}")); +} + +fn item(id: &str, name: &str, qty: i64) -> Document { + let mut doc = Document::new(id); + doc.set("name", Value::String(name.into())); + doc.set("qty", Value::Integer(qty)); + doc +} + +/// `written` as it reads back: the declared key is a declared field, so it +/// reads back holding the document id. +fn stored(written: &Document) -> Document { + let mut doc = written.clone(); + doc.set("sku", Value::String(written.id.clone())); + doc +} + +/// Read `id` through both clients and assert each returns `expected`, or +/// nothing when `expected` is `None`. +async fn assert_both_read( + remote: &NodeDbRemote, + native: &NativeClient, + id: &str, + expected: Option<&Document>, +) { + let over_pgwire = remote + .document_get("items", id) + .await + .unwrap_or_else(|e| panic!("remote get {id}: {e}")); + let over_native = native + .document_get("items", id) + .await + .unwrap_or_else(|e| panic!("native get {id}: {e}")); + assert_eq!(over_pgwire.as_ref(), expected, "remote read of {id}"); + assert_eq!(over_native.as_ref(), expected, "native read of {id}"); +} + +#[tokio::test] +async fn documents_round_trip_under_a_declared_key() { + let server = TestServer::start().await; + let remote = remote(&server).await; + let native = native(&server); + sql( + &remote, + "CREATE COLLECTION items (sku STRING PRIMARY KEY, name STRING) \ + WITH (engine='document_schemaless')", + ) + .await; + + let by_remote = item("p1", "pen", 3); + remote + .document_put("items", by_remote.clone()) + .await + .expect("remote put under a declared key"); + let by_native = item("p2", "ink", 7); + native + .document_put("items", by_native.clone()) + .await + .expect("native put under a declared key"); + + assert_both_read(&remote, &native, "p1", Some(&stored(&by_remote))).await; + assert_both_read(&remote, &native, "p2", Some(&stored(&by_native))).await; + + // SQL sees each document under its key column. + let keys = remote + .execute_sql("SELECT sku FROM items WHERE name = 'ink'", &[]) + .await + .expect("SQL reads the native-written row by its key"); + assert_eq!(keys.rows, vec![vec![Value::String("p2".into())]]); + + // A scan on a non-key predicate returns the stored fields only: the key + // column holds the id, and no `id` column appears beside it. + let scanned = remote + .execute_sql("SELECT * FROM items WHERE name = 'ink'", &[]) + .await + .expect("SELECT * scan of a declared-key collection"); + assert!( + !scanned.columns.iter().any(|c| c == "id"), + "a declared-key scan shows no id column: {:?}", + scanned.columns + ); + let sku = scanned + .columns + .iter() + .position(|c| c == "sku") + .expect("the scan shows the key column"); + assert_eq!(scanned.rows.len(), 1); + assert_eq!(scanned.rows[0][sku], Value::String("p2".into())); + let documents = remote + .execute_sql( + "SELECT to_jsonb(*) AS document FROM items WHERE name = 'ink'", + &[], + ) + .await + .expect("to_jsonb(*) scan of a declared-key collection"); + let [row] = documents.rows.as_slice() else { + panic!("one row matches: {:?}", documents.rows); + }; + let Some(Value::String(json)) = row.first() else { + panic!("the document cell is JSON text: {row:?}"); + }; + let fields: serde_json::Value = sonic_rs::from_str(json).expect("the cell is JSON"); + assert_eq!( + fields, + serde_json::json!({"sku": "p2", "name": "ink", "qty": 7}), + "a scanned row is exactly its stored fields" + ); + + // A put replaces the document whole, through either client. + let replacement = item("p1", "pencil", 4); + native + .document_put("items", replacement.clone()) + .await + .expect("native replace of a remote-written document"); + assert_both_read(&remote, &native, "p1", Some(&stored(&replacement))).await; + + // A key field that names another document is refused. + let mut conflicting = item("p3", "cap", 1); + conflicting.set("sku", Value::String("other".into())); + remote + .document_put("items", conflicting.clone()) + .await + .expect_err("remote put whose key field names another id"); + native + .document_put("items", conflicting) + .await + .expect_err("native put whose key field names another id"); + + remote + .document_delete("items", "p1") + .await + .expect("remote delete under a declared key"); + native + .document_delete("items", "p2") + .await + .expect("native delete under a declared key"); + assert_both_read(&remote, &native, "p1", None).await; + assert_both_read(&remote, &native, "p2", None).await; + + server.graceful_shutdown().await; +} + +fn ids(hits: &[SearchResult]) -> Vec<&str> { + hits.iter().map(|h| h.id.as_str()).collect() +} + +#[tokio::test] +async fn vector_search_restricts_allowed_ids_on_a_declared_key() { + let server = TestServer::start().await; + let remote = remote(&server).await; + let native = native(&server); + sql( + &remote, + "CREATE COLLECTION sku_vecs \ + FIELDS (sku TEXT PRIMARY KEY, embedding VECTOR(2)) \ + WITH (engine='vector', m=8, ef_construction=50)", + ) + .await; + for (sku, x) in [ + ("near1", 0.0), + ("near2", 0.1), + ("near3", 0.2), + ("far1", 10.0), + ("far2", 20.0), + ("far3", 30.0), + ] { + sql( + &remote, + &format!("INSERT INTO sku_vecs (sku, embedding) VALUES ('{sku}', ARRAY[{x}, 0.0])"), + ) + .await; + } + + // The nearest vectors lie outside the allowed set. The restriction must + // apply before the top-k cut, so k rows come back from the allowed set. + let allowed: HashSet = ["far1", "far2", "far3"] + .iter() + .map(|s| s.to_string()) + .collect(); + let over_pgwire = remote + .vector_search("sku_vecs", &[0.0, 0.0], 2, None, Some(&allowed)) + .await + .expect("remote vector_search with allowed_ids on a declared key"); + assert_eq!(ids(&over_pgwire), vec!["far1", "far2"], "remote"); + let over_native = native + .vector_search("sku_vecs", &[0.0, 0.0], 2, None, Some(&allowed)) + .await + .expect("native vector_search with allowed_ids on a declared key"); + assert_eq!(ids(&over_native), vec!["far1", "far2"], "native"); + + server.graceful_shutdown().await; +} diff --git a/nodedb-client-tests/tests/document_put_transaction.rs b/nodedb-client-tests/tests/document_put_transaction.rs new file mode 100644 index 000000000..b1f418777 --- /dev/null +++ b/nodedb-client-tests/tests/document_put_transaction.rs @@ -0,0 +1,160 @@ +// SPDX-License-Identifier: BUSL-1.1 + +//! `NodeDbRemote::document_put` inside and outside a caller's transaction. +//! +//! The put is a `DELETE` then an `INSERT`. Inside the caller's block it runs +//! under a savepoint, so it never commits or rolls back that block. A failed +//! put rolls back to its savepoint and leaves the block usable with every +//! earlier write. Outside a block it runs in its own transaction. + +use nodedb_client::{Document, NodeDb, NodeDbRemote, Value}; +use nodedb_test_support::pgwire_harness::TestServer; + +async fn remote(server: &TestServer) -> NodeDbRemote { + let conn_str = format!( + "host=127.0.0.1 port={} user=nodedb dbname=default", + server.pg_port + ); + NodeDbRemote::connect(&conn_str) + .await + .expect("pgwire connect to harness must succeed") +} + +async fn sql(remote: &NodeDbRemote, statement: &str) { + remote + .execute_sql(statement, &[]) + .await + .unwrap_or_else(|e| panic!("{statement}: {e}")); +} + +fn document(id: &str, field: &str, value: &str) -> Document { + let mut doc = Document::new(id); + doc.set(field, Value::String(value.into())); + doc +} + +async fn exists(remote: &NodeDbRemote, collection: &str, id: &str) -> bool { + remote + .document_get(collection, id) + .await + .expect("document_get") + .is_some() +} + +#[tokio::test] +async fn a_put_inside_a_rolled_back_block_rolls_back_with_it() { + let server = TestServer::start().await; + let remote = remote(&server).await; + sql(&remote, "CREATE COLLECTION docs").await; + + sql(&remote, "BEGIN").await; + sql(&remote, "INSERT INTO docs (id, body) VALUES ('early', 'x')").await; + remote + .document_put("docs", document("d1", "body", "in block")) + .await + .expect("put inside the block"); + sql(&remote, "ROLLBACK").await; + + assert!( + !exists(&remote, "docs", "d1").await, + "the put must not commit the caller's block" + ); + assert!( + !exists(&remote, "docs", "early").await, + "the block's earlier write rolls back with it" + ); + + server.graceful_shutdown().await; +} + +#[tokio::test] +async fn a_put_inside_a_committed_block_persists_with_it() { + let server = TestServer::start().await; + let remote = remote(&server).await; + sql(&remote, "CREATE COLLECTION docs").await; + + sql(&remote, "BEGIN").await; + sql(&remote, "INSERT INTO docs (id, body) VALUES ('early', 'x')").await; + remote + .document_put("docs", document("d1", "body", "in block")) + .await + .expect("put inside the block"); + sql(&remote, "COMMIT").await; + + assert!(exists(&remote, "docs", "d1").await); + assert!(exists(&remote, "docs", "early").await); + + server.graceful_shutdown().await; +} + +#[tokio::test] +async fn a_failed_put_inside_a_block_leaves_the_block_usable() { + let server = TestServer::start().await; + let remote = remote(&server).await; + sql(&remote, "CREATE COLLECTION docs").await; + // A strict BIGINT column refuses the put's text value. The refusal comes + // from the INSERT, after the put opened its savepoint. + sql( + &remote, + "CREATE COLLECTION counters (id TEXT PRIMARY KEY, n BIGINT) \ + WITH (engine='document_strict')", + ) + .await; + + sql(&remote, "BEGIN").await; + sql(&remote, "INSERT INTO docs (id, body) VALUES ('early', 'x')").await; + remote + .document_put("counters", document("c1", "n", "not-a-number")) + .await + .expect_err("a text value in a BIGINT column is refused"); + // The put rolled back to its savepoint: the block takes more writes. + sql(&remote, "INSERT INTO docs (id, body) VALUES ('after', 'y')").await; + sql(&remote, "COMMIT").await; + + assert!( + exists(&remote, "docs", "early").await, + "the failed put must not discard the block's earlier write" + ); + assert!(exists(&remote, "docs", "after").await); + + server.graceful_shutdown().await; +} + +#[tokio::test] +async fn a_failed_standalone_put_keeps_the_old_document() { + let server = TestServer::start().await; + let remote = remote(&server).await; + sql(&remote, "CREATE COLLECTION people").await; + sql(&remote, "CREATE UNIQUE INDEX ON people(email)").await; + + remote + .document_put("people", document("a", "email", "a@x")) + .await + .expect("put a"); + remote + .document_put("people", document("b", "email", "b@x")) + .await + .expect("put b"); + + // The replace deletes `b`, then its insert breaks the unique index. + remote + .document_put("people", document("b", "email", "a@x")) + .await + .expect_err("a duplicate unique value is refused"); + + let b = remote + .document_get("people", "b") + .await + .expect("document_get") + .expect("the refused replace must not delete the old document"); + assert_eq!(b.fields.get("email"), Some(&Value::String("b@x".into()))); + + // The connection is back outside any block: a later put commits alone. + remote + .document_put("people", document("c", "email", "c@x")) + .await + .expect("put c"); + assert!(exists(&remote, "people", "c").await); + + server.graceful_shutdown().await; +} diff --git a/nodedb-client-tests/tests/document_round_trip.rs b/nodedb-client-tests/tests/document_round_trip.rs new file mode 100644 index 000000000..b74458e24 --- /dev/null +++ b/nodedb-client-tests/tests/document_round_trip.rs @@ -0,0 +1,305 @@ +// SPDX-License-Identifier: BUSL-1.1 + +//! End-to-end document shape tests for `NodeDbRemote` and `NativeClient`. +//! +//! A document's fields are the row's top-level columns under both clients. +//! SQL filters and text-searches each field by name. A document written by +//! one client reads back with the same fields through the other. + +use nodedb_client::native::pool::PoolConfig; +use nodedb_client::{Document, NativeClient, NodeDb, NodeDbRemote, Value}; +use nodedb_test_support::pgwire_harness::TestServer; +use nodedb_types::text_search::{QueryMode, TextSearchParams}; + +/// All-terms fuzzy search: params the remote client renders as named +/// `text_match` options. +fn pgwire_search_params() -> TextSearchParams { + TextSearchParams { + mode: QueryMode::And, + fuzzy: true, + } +} + +async fn remote(server: &TestServer) -> NodeDbRemote { + let conn_str = format!( + "host=127.0.0.1 port={} user=nodedb dbname=default", + server.pg_port + ); + NodeDbRemote::connect(&conn_str) + .await + .expect("pgwire connect to harness must succeed") +} + +fn native(server: &TestServer) -> NativeClient { + NativeClient::new(PoolConfig::new( + format!("127.0.0.1:{}", server.native_port), + nodedb_types::protocol::AuthMethod::Trust { + username: "nodedb".into(), + }, + )) +} + +async fn seed_collection(remote: &NodeDbRemote) { + remote + .execute_sql("CREATE COLLECTION docs", &[]) + .await + .expect("CREATE COLLECTION docs"); + remote + .execute_sql("CREATE SEARCH INDEX ON docs FIELDS body", &[]) + .await + .expect("CREATE SEARCH INDEX on body"); +} + +fn document(id: &str, fields: &[(&str, &str)]) -> Document { + let mut doc = Document::new(id); + for (name, value) in fields { + doc.set(*name, Value::String((*value).into())); + } + doc +} + +#[tokio::test] +async fn remote_document_round_trips_as_top_level_fields() { + let server = TestServer::start().await; + let remote = remote(&server).await; + seed_collection(&remote).await; + + let written = document( + "d1", + &[("body", "machine learning is everywhere"), ("title", "ml")], + ); + remote + .document_put("docs", written.clone()) + .await + .expect("remote document_put"); + + let read = remote + .document_get("docs", "d1") + .await + .expect("remote document_get") + .expect("the written document exists"); + assert_eq!(read.id, "d1"); + assert_eq!( + read.fields, written.fields, + "fields must round-trip unchanged" + ); + + let filtered = remote + .execute_sql( + "SELECT body FROM docs WHERE body = 'machine learning is everywhere'", + &[], + ) + .await + .expect("filter on the body field"); + assert_eq!( + filtered.rows, + vec![vec![Value::String("machine learning is everywhere".into())]], + "the body field must be a top-level column SQL can filter and project" + ); + + let hits = remote + .text_search( + "docs", + "body", + "machine learning", + 10, + pgwire_search_params(), + None, + ) + .await + .expect("text_search on the body field"); + assert_eq!( + hits.iter().map(|h| h.id.as_str()).collect::>(), + vec!["d1"], + "the body field must be text-searchable under the document id" + ); + + server.graceful_shutdown().await; +} + +#[tokio::test] +async fn remote_document_put_replaces_the_whole_field_set() { + let server = TestServer::start().await; + let remote = remote(&server).await; + seed_collection(&remote).await; + + remote + .document_put("docs", document("d1", &[("body", "one"), ("extra", "x")])) + .await + .expect("first put"); + let replacement = document("d1", &[("body", "two")]); + remote + .document_put("docs", replacement.clone()) + .await + .expect("second put of the same id"); + + let read = remote + .document_get("docs", "d1") + .await + .expect("document_get") + .expect("the document exists"); + assert_eq!( + read.fields, replacement.fields, + "a put replaces every field; `extra` must be gone" + ); + + server.graceful_shutdown().await; +} + +#[tokio::test] +async fn documents_cross_between_remote_and_native_unchanged() { + let server = TestServer::start().await; + let remote = remote(&server).await; + let native = native(&server); + seed_collection(&remote).await; + + let by_remote = document("from-remote", &[("body", "written over pgwire")]); + remote + .document_put("docs", by_remote.clone()) + .await + .expect("remote put"); + let seen_by_native = native + .document_get("docs", "from-remote") + .await + .expect("native get") + .expect("native sees the remote-written document"); + assert_eq!(seen_by_native.id, "from-remote"); + assert_eq!(seen_by_native.fields, by_remote.fields); + + let by_native = document( + "from-native", + &[("body", "written over the native protocol")], + ); + native + .document_put("docs", by_native.clone()) + .await + .expect("native put"); + let seen_by_remote = remote + .document_get("docs", "from-native") + .await + .expect("remote get") + .expect("remote sees the native-written document"); + assert_eq!(seen_by_remote.id, "from-native"); + assert_eq!(seen_by_remote.fields, by_native.fields); + + let ids = remote + .execute_sql( + "SELECT id FROM docs WHERE body = 'written over the native protocol'", + &[], + ) + .await + .expect("SQL reads the native-written row"); + assert_eq!( + ids.rows, + vec![vec![Value::String("from-native".into())]], + "SQL must see a native-written document under its own id" + ); + + server.graceful_shutdown().await; +} + +#[tokio::test] +async fn native_document_round_trips_typed_fields_arrays_and_nulls() { + let server = TestServer::start().await; + let remote = remote(&server).await; + let native = native(&server); + seed_collection(&remote).await; + + let mut written = Document::new("typed"); + written.set("count", Value::Integer(42)); + written.set("ratio", Value::Float(1.5)); + written.set("active", Value::Bool(true)); + written.set("name", Value::String("n".into())); + written.set( + "tags", + Value::Array(vec![Value::String("a".into()), Value::Integer(2)]), + ); + written.set("missing", Value::Null); + native + .document_put("docs", written.clone()) + .await + .expect("native put of typed fields"); + + let read = native + .document_get("docs", "typed") + .await + .expect("native get") + .expect("the typed document exists"); + assert_eq!(read.id, "typed"); + assert_eq!( + read.fields, written.fields, + "the native client returns every field with its own type" + ); + + server.graceful_shutdown().await; +} + +/// Read `id` through both clients and assert each returns `expected`. +async fn assert_both_read( + remote: &NodeDbRemote, + native: &NativeClient, + id: &str, + expected: &Document, +) { + let over_pgwire = remote + .document_get("docs", id) + .await + .expect("remote get") + .expect("the document exists"); + let over_native = native + .document_get("docs", id) + .await + .expect("native get") + .expect("the document exists"); + assert_eq!(over_pgwire.id, id); + assert_eq!(over_native.id, id); + assert_eq!( + over_pgwire.fields, expected.fields, + "the pgwire client reads every field with its stored type" + ); + assert_eq!( + over_native.fields, expected.fields, + "the native client reads every field with its stored type" + ); +} + +#[tokio::test] +async fn remote_and_native_read_the_same_typed_document() { + let server = TestServer::start().await; + let remote = remote(&server).await; + let native = native(&server); + seed_collection(&remote).await; + + // Every kind a pgwire put binds. + let mut by_remote = Document::new("n1"); + by_remote.set("count", Value::Integer(5)); + by_remote.set("ratio", Value::Float(1.5)); + by_remote.set("whole", Value::Float(2.0)); + by_remote.set("active", Value::Bool(true)); + by_remote.set("name", Value::String("n".into())); + by_remote.set( + "tags", + Value::Array(vec![Value::String("a".into()), Value::String("b".into())]), + ); + by_remote.set("missing", Value::Null); + remote + .document_put("docs", by_remote.clone()) + .await + .expect("remote put of typed fields"); + assert_both_read(&remote, &native, "n1", &by_remote).await; + + // A nested object, which only the native put stores. + let mut inner = std::collections::HashMap::new(); + inner.insert("city".to_string(), Value::String("KL".into())); + inner.insert("zip".to_string(), Value::Integer(50000)); + let mut by_native = Document::new("n2"); + by_native.set("address", Value::Object(inner)); + by_native.set("score", Value::Float(0.25)); + native + .document_put("docs", by_native.clone()) + .await + .expect("native put of a nested object"); + assert_both_read(&remote, &native, "n2", &by_native).await; + + server.graceful_shutdown().await; +} diff --git a/nodedb-client-tests/tests/graph_edge_filter_round_trip.rs b/nodedb-client-tests/tests/graph_edge_filter_round_trip.rs new file mode 100644 index 000000000..5c1288730 --- /dev/null +++ b/nodedb-client-tests/tests/graph_edge_filter_round_trip.rs @@ -0,0 +1,248 @@ +// SPDX-License-Identifier: BUSL-1.1 + +//! End-to-end `graph_traverse` and `graph_shortest_path` with an +//! `EdgeFilter`, over pgwire (`NodeDbRemote`) and the native protocol +//! (`NativeClient`). +//! +//! Both clients send `GRAPH TRAVERSE` / `GRAPH PATH` SQL with every label +//! and the property filters as `EDGE WHERE`, and decode one result shape. +//! The fixture gives each filter one edge it admits and one it rejects, so +//! a dropped label, a dropped predicate, or a dropped property is visible. +//! +//! Fixture, in one collection: +//! - `hub -road{score 9, owner O'Reilly}-> m1 -rail{score 9}-> dst` +//! - `hub -road{score 1}-> cheap -road{score 1}-> dst` +//! - `hub -other{score 9}-> out` + +use std::collections::BTreeSet; + +use nodedb_client::native::pool::PoolConfig; +use nodedb_client::{Document, NativeClient, NodeDb, NodeDbRemote, NodeId, Value}; +use nodedb_test_support::pgwire_harness::TestServer; +use nodedb_types::filter::{EdgeFilter, MetadataFilter}; +use nodedb_types::graph::Direction; + +fn node(id: &str) -> NodeId { + NodeId::try_new(id).expect("fixture node id") +} + +fn properties(score: i64, owner: Option<&str>) -> Document { + let mut doc = Document::new("edge"); + doc.set("score", Value::Integer(score)); + if let Some(owner) = owner { + doc.set("owner", Value::String(owner.to_string())); + } + doc +} + +fn filter(labels: &[&str], property_filters: Vec) -> EdgeFilter { + EdgeFilter { + labels: labels.iter().map(|label| label.to_string()).collect(), + property_filters, + } +} + +fn score_above(n: i64) -> Vec { + vec![MetadataFilter::Gt { + field: "score".into(), + value: Value::Integer(n), + }] +} + +async fn seed(db: &D, collection: &str) { + db.execute_sql(&format!("CREATE COLLECTION {collection}"), &[]) + .await + .expect("create collection"); + for (src, dst, label, score, owner) in [ + ("hub", "m1", "road", 9, Some("O'Reilly")), + ("m1", "dst", "rail", 9, None), + ("hub", "cheap", "road", 1, None), + ("cheap", "dst", "road", 1, None), + ("hub", "out", "other", 9, None), + ] { + db.graph_insert_edge( + collection, + &node(src), + &node(dst), + label, + Some(properties(score, owner)), + ) + .await + .unwrap_or_else(|e| panic!("seed {src}-{label}->{dst}: {e}")); + } +} + +async fn node_set( + db: &D, + collection: &str, + start: &str, + depth: u8, + direction: Direction, + edge_filter: &EdgeFilter, +) -> BTreeSet { + db.graph_traverse( + collection, + &node(start), + depth, + direction, + Some(edge_filter), + ) + .await + .expect("traversal completes") + .nodes + .into_iter() + .map(|n| n.id.as_str().to_string()) + .collect() +} + +fn set(ids: &[&str]) -> BTreeSet { + ids.iter().map(|id| id.to_string()).collect() +} + +async fn exercise(db: &D, collection: &str) { + seed(db, collection).await; + + // Every listed label is followed. `other` is not listed. + let road_rail = filter(&["road", "rail"], Vec::new()); + assert_eq!( + node_set(db, collection, "hub", 2, Direction::Out, &road_rail).await, + set(&["cheap", "dst", "hub", "m1"]) + ); + + // Each crossed edge carries its properties, string escaping included. + let sg = db + .graph_traverse( + collection, + &node("hub"), + 1, + Direction::Out, + Some(&road_rail), + ) + .await + .expect("traversal completes"); + let to_m1 = sg + .edges + .iter() + .find(|edge| edge.to.as_str() == "m1") + .expect("hub -road-> m1 is crossed"); + assert_eq!(to_m1.label, "road"); + assert_eq!(to_m1.properties.get("score"), Some(&Value::Integer(9))); + assert_eq!( + to_m1.properties.get("owner"), + Some(&Value::String("O'Reilly".into())) + ); + let to_cheap = sg + .edges + .iter() + .find(|edge| edge.to.as_str() == "cheap") + .expect("hub -road-> cheap is crossed"); + assert_eq!(to_cheap.properties.get("score"), Some(&Value::Integer(1))); + assert_eq!(to_cheap.properties.get("owner"), None); + + // The predicate keeps only edges whose properties match. + let reilly = filter(&[], vec![MetadataFilter::eq("owner", "O'Reilly")]); + assert_eq!( + node_set(db, collection, "hub", 1, Direction::Out, &reilly).await, + set(&["hub", "m1"]) + ); + let scored = filter(&["road", "rail"], score_above(5)); + assert_eq!( + node_set(db, collection, "hub", 2, Direction::Out, &scored).await, + set(&["dst", "hub", "m1"]) + ); + + // Incoming: from `dst`, only `m1 -rail{9}->` passes `score > 5`. + let inward = db + .graph_traverse( + collection, + &node("dst"), + 1, + Direction::In, + Some(&filter(&[], score_above(5))), + ) + .await + .expect("traversal completes"); + let inward_nodes: BTreeSet<&str> = inward.nodes.iter().map(|n| n.id.as_str()).collect(); + assert_eq!(inward_nodes, BTreeSet::from(["dst", "m1"])); + let inward_edges: Vec<(&str, &str, &str)> = inward + .edges + .iter() + .map(|e| (e.from.as_str(), e.label.as_str(), e.to.as_str())) + .collect(); + assert_eq!(inward_edges, vec![("m1", "rail", "dst")]); + + // Both ways from `m1` over `road` only reaches `hub`. + assert_eq!( + node_set( + db, + collection, + "m1", + 1, + Direction::Both, + &filter(&["road"], Vec::new()) + ) + .await, + set(&["hub", "m1"]) + ); + + // A path that needs two labels, under the predicate. + let path = db + .graph_shortest_path(collection, &node("hub"), &node("dst"), 4, Some(&scored)) + .await + .expect("path completes") + .expect("hub -road-> m1 -rail-> dst passes"); + let ids: Vec<&str> = path.iter().map(NodeId::as_str).collect(); + assert_eq!(ids, vec!["hub", "m1", "dst"]); + + // `road` alone under the predicate reaches only `m1`. + let road_only = filter(&["road"], score_above(5)); + assert_eq!( + db.graph_shortest_path(collection, &node("hub"), &node("dst"), 4, Some(&road_only)) + .await + .expect("path completes"), + None + ); + + // `road` alone without the predicate takes the cheap route. + let cheap = db + .graph_shortest_path( + collection, + &node("hub"), + &node("dst"), + 4, + Some(&filter(&["road"], Vec::new())), + ) + .await + .expect("path completes") + .expect("hub -road-> cheap -road-> dst"); + let ids: Vec<&str> = cheap.iter().map(NodeId::as_str).collect(); + assert_eq!(ids, vec!["hub", "cheap", "dst"]); +} + +#[tokio::test] +async fn remote_edge_filter_round_trip() { + let server = TestServer::start().await; + let conn_str = format!( + "host=127.0.0.1 port={} user=nodedb dbname=default", + server.pg_port + ); + let remote = NodeDbRemote::connect(&conn_str) + .await + .expect("pgwire connect to harness must succeed"); + exercise(&remote, "ef_remote").await; + server.graceful_shutdown().await; +} + +#[tokio::test] +async fn native_edge_filter_round_trip() { + let server = TestServer::start().await; + let pool = PoolConfig::new( + format!("127.0.0.1:{}", server.native_port), + nodedb_types::protocol::AuthMethod::Trust { + username: "nodedb".into(), + }, + ); + let native = NativeClient::new(pool); + exercise(&native, "ef_native").await; + server.graceful_shutdown().await; +} diff --git a/nodedb-client-tests/tests/graph_traverse_remote_round_trip.rs b/nodedb-client-tests/tests/graph_traverse_remote_round_trip.rs index 724521651..bbb5611b9 100644 --- a/nodedb-client-tests/tests/graph_traverse_remote_round_trip.rs +++ b/nodedb-client-tests/tests/graph_traverse_remote_round_trip.rs @@ -9,6 +9,9 @@ //! An empty subgraph is indistinguishable from "the wire short-circuits //! before the server's traversal runs" — the silent-fake pattern this //! test guards against. +//! +//! The test also checks In/Out/Both direction, edge orientation, reciprocal +//! edges, a self-loop and discovery depths. use nodedb_client::{NodeDb, NodeDbRemote, NodeId}; use nodedb_test_support::pgwire_harness::TestServer; @@ -48,19 +51,14 @@ async fn graph_traverse_returns_inserted_subgraph() { .expect("seed edge a->s"); let sg = remote - .graph_traverse("smoke_g", &a, 2, None) + .graph_traverse("smoke_g", &a, 2, nodedb_types::graph::Direction::Out, None) .await .expect("graph_traverse must complete against a populated graph"); - // Spec: with edges {a→b, b→c, a→s} present, a depth-2 traversal - // from `a` reaches at least {a, b, c, s}. Returning an empty - // subgraph reproduces the original bug: the wire returns success - // with `nodes=[], edges=[]` indistinguishable from "the seed has - // no out-edges within depth N". + // A depth-2 outgoing traversal reaches both direct and two-hop neighbors. assert!( !sg.nodes.is_empty(), - "depth-2 traversal from seed must surface reachable nodes, got empty subgraph \ - (regression: graph_traverse short-circuits after same-session graph_insert_edge)" + "depth-2 traversal from seed must surface reachable nodes, got empty subgraph" ); let node_ids: std::collections::HashSet<&str> = @@ -78,14 +76,127 @@ async fn graph_traverse_returns_inserted_subgraph() { "depth-2 traversal must reach direct neighbor sess; nodes={node_ids:?}" ); - // Edge regression guard: a populated traversal must also carry - // the edges it crossed. The original bug returned an object with - // both `nodes=[]` AND `edges=[]`; asserting on edges separately - // catches a half-fix that surfaces nodes but drops edges. + // A populated traversal includes the edges it crossed. assert!( !sg.edges.is_empty(), - "depth-2 traversal must surface the edges it crossed; got empty edges \ - (regression: traverse drops traversed edges from the wire response)" + "depth-2 traversal must surface the edges it crossed; got empty edges" + ); + + for (direction, includes_a, includes_c) in [ + (nodedb_types::graph::Direction::In, true, false), + (nodedb_types::graph::Direction::Out, false, true), + (nodedb_types::graph::Direction::Both, true, true), + ] { + let traversal = remote + .graph_traverse("smoke_g", &b, 1, direction, None) + .await + .expect("depth-1 traversal must complete"); + let nodes: std::collections::HashSet<&str> = traversal + .nodes + .iter() + .map(|node| node.id.as_str()) + .collect(); + assert_eq!( + nodes.contains(a.as_str()), + includes_a, + "direction={direction}, nodes={nodes:?}" + ); + assert_eq!( + nodes.contains(c.as_str()), + includes_c, + "direction={direction}, nodes={nodes:?}" + ); + let edges: std::collections::HashSet<(&str, &str, &str)> = traversal + .edges + .iter() + .map(|edge| (edge.from.as_str(), edge.label.as_str(), edge.to.as_str())) + .collect(); + assert_eq!( + edges.contains(&(a.as_str(), "next", b.as_str())), + includes_a + ); + assert_eq!( + edges.contains(&(b.as_str(), "next", c.as_str())), + includes_c + ); + assert_eq!( + edges.len(), + usize::from(includes_a) + usize::from(includes_c) + ); + } + + remote + .graph_insert_edge("smoke_g", &b, &a, "next", None) + .await + .expect("reciprocal edge b->a"); + remote + .graph_insert_edge("smoke_g", &b, &b, "self", None) + .await + .expect("self loop b->b"); + let reciprocal = remote + .graph_traverse("smoke_g", &b, 2, nodedb_types::graph::Direction::Both, None) + .await + .expect("bidirectional traversal with reciprocal edges"); + let edges: std::collections::HashSet<(&str, &str, &str)> = reciprocal + .edges + .iter() + .map(|edge| (edge.from.as_str(), edge.label.as_str(), edge.to.as_str())) + .collect(); + assert_eq!( + edges.len(), + reciprocal.edges.len(), + "physical edges appear once" + ); + assert_eq!(edges.len(), 5); + for edge in [ + (a.as_str(), "next", b.as_str()), + (b.as_str(), "next", a.as_str()), + (b.as_str(), "next", c.as_str()), + (a.as_str(), "in_session", s.as_str()), + (b.as_str(), "self", b.as_str()), + ] { + assert!(edges.contains(&edge), "missing physical edge {edge:?}"); + } + for node in &reciprocal.nodes { + let expected_depth = match node.id.as_str() { + "chunk_b" => 0, + "chunk_a" | "chunk_c" => 1, + "sess" => 2, + other => panic!("unexpected node {other}"), + }; + assert_eq!(node.depth, expected_depth); + } + + // A label set follows every listed label and no other. + let x = NodeId::try_new("outside").expect("fixture"); + remote + .graph_insert_edge("smoke_g", &a, &x, "other", None) + .await + .expect("seed edge a->outside"); + let labelled = remote + .graph_traverse( + "smoke_g", + &a, + 1, + nodedb_types::graph::Direction::Out, + Some(&nodedb_types::filter::EdgeFilter::labels([ + "next", + "in_session", + ])), + ) + .await + .expect("label-set traversal must complete"); + let nodes: std::collections::BTreeSet<&str> = + labelled.nodes.iter().map(|node| node.id.as_str()).collect(); + assert_eq!( + nodes, + std::collections::BTreeSet::from(["chunk_a", "chunk_b", "sess"]), + "'other' is not listed" + ); + assert!( + labelled.edges.iter().all(|edge| edge.label != "other"), + "no 'other' edge is crossed: {:?}", + labelled.edges ); server.graceful_shutdown().await; diff --git a/nodedb-client-tests/tests/text_search_default_round_trip.rs b/nodedb-client-tests/tests/text_search_default_round_trip.rs index 71ae874d4..b3f65bb87 100644 --- a/nodedb-client-tests/tests/text_search_default_round_trip.rs +++ b/nodedb-client-tests/tests/text_search_default_round_trip.rs @@ -8,9 +8,12 @@ //! against — a fake "no matches" answer is indistinguishable from a //! real one and lets callers proceed as if FTS were working. -use nodedb_client::{NodeDb, NodeDbRemote}; +use std::collections::HashSet; + +use nodedb_client::native::pool::PoolConfig; +use nodedb_client::{NativeClient, NodeDb, NodeDbRemote}; use nodedb_test_support::pgwire_harness::TestServer; -use nodedb_types::text_search::TextSearchParams; +use nodedb_types::text_search::{QueryMode, TextSearchParams}; #[tokio::test] async fn text_search_returns_real_matches() { @@ -58,7 +61,10 @@ async fn text_search_returns_real_matches() { "body", "machine learning", 10, - TextSearchParams::default(), + TextSearchParams { + mode: QueryMode::And, + fuzzy: true, + }, None, ) .await @@ -68,5 +74,185 @@ async fn text_search_returns_real_matches() { "text_search must return real BM25-ranked matches; got empty" ); + // A second document holds one of the two terms. `And` keeps only the + // document holding both; `Or` returns both. + let mut doc = nodedb_client::Document::new("d2"); + doc.set( + "body", + nodedb_client::Value::String("machine shop tools".into()), + ); + remote + .document_put("docs", doc) + .await + .expect("seed second document"); + let ids = |hits: &[nodedb_client::SearchResult]| -> Vec { + let mut ids: Vec = hits.iter().map(|h| h.id.clone()).collect(); + ids.sort(); + ids + }; + let all_terms = remote + .text_search( + "docs", + "body", + "machine learning", + 10, + TextSearchParams { + mode: QueryMode::And, + fuzzy: false, + }, + None, + ) + .await + .expect("And-mode text_search"); + assert_eq!(ids(&all_terms), vec!["d1".to_string()]); + let any_term = remote + .text_search( + "docs", + "body", + "machine learning", + 10, + TextSearchParams::default(), + None, + ) + .await + .expect("Or-mode text_search"); + assert_eq!(ids(&any_term), vec!["d1".to_string(), "d2".to_string()]); + + // A typo matches only through the fuzzy fallback. + let exact = remote + .text_search( + "docs", + "body", + "machime", + 10, + TextSearchParams::default(), + None, + ) + .await + .expect("non-fuzzy text_search"); + assert!(exact.is_empty(), "no exact match for a typo; got {exact:?}"); + let fuzzy = remote + .text_search( + "docs", + "body", + "machime", + 10, + TextSearchParams { + mode: QueryMode::Or, + fuzzy: true, + }, + None, + ) + .await + .expect("fuzzy text_search"); + assert_eq!(ids(&fuzzy), vec!["d1".to_string(), "d2".to_string()]); + + // `allowed_ids` restricts the candidates on the server. + let only_other: std::collections::HashSet = + std::iter::once("other".to_string()).collect(); + let restricted = remote + .text_search( + "docs", + "body", + "machine learning", + 10, + TextSearchParams { + mode: QueryMode::And, + fuzzy: true, + }, + Some(&only_other), + ) + .await + .expect("text_search with allowed_ids"); + assert!( + restricted.is_empty(), + "d1 is not in allowed_ids, so no hit may return; got {restricted:?}" + ); + + server.graceful_shutdown().await; +} + +/// The hits as `(id, score)` pairs, in rank order. +fn ranked(hits: &[nodedb_client::SearchResult]) -> Vec<(String, f32)> { + hits.iter().map(|h| (h.id.clone(), h.distance)).collect() +} + +/// The native client sends the statement the pgwire client sends, so the +/// two return the same ids with the same scores for every search. +#[tokio::test] +async fn native_text_search_matches_remote() { + let server = TestServer::start().await; + let remote = NodeDbRemote::connect(&format!( + "host=127.0.0.1 port={} user=nodedb dbname=default", + server.pg_port + )) + .await + .expect("pgwire connect to harness must succeed"); + let native = NativeClient::new(PoolConfig::new( + format!("127.0.0.1:{}", server.native_port), + nodedb_types::protocol::AuthMethod::Trust { + username: "nodedb".into(), + }, + )); + + remote + .execute_sql("CREATE COLLECTION docs", &[]) + .await + .expect("CREATE COLLECTION docs"); + remote + .execute_sql("CREATE SEARCH INDEX ON docs FIELDS body", &[]) + .await + .expect("CREATE SEARCH INDEX on body field"); + for (id, body) in [ + ("d1", "machine learning is everywhere"), + ("d2", "machine shop tools"), + ("d3", "learning to cook"), + ] { + let mut doc = nodedb_client::Document::new(id); + doc.set("body", nodedb_client::Value::String(body.into())); + native + .document_put("docs", doc) + .await + .unwrap_or_else(|e| panic!("seed {id}: {e}")); + } + + let only_d1: HashSet = std::iter::once("d1".to_string()).collect(); + let cases: [(&str, TextSearchParams, Option<&HashSet>); 4] = [ + ("any term", TextSearchParams::default(), None), + ( + "all terms", + TextSearchParams { + mode: QueryMode::And, + fuzzy: false, + }, + None, + ), + ( + "fuzzy", + TextSearchParams { + mode: QueryMode::Or, + fuzzy: true, + }, + None, + ), + ("allowed ids", TextSearchParams::default(), Some(&only_d1)), + ]; + for (case, params, allowed) in cases { + let over_pgwire = remote + .text_search("docs", "body", "machine learning", 10, params, allowed) + .await + .unwrap_or_else(|e| panic!("{case}: remote text_search: {e}")); + let over_native = native + .text_search("docs", "body", "machine learning", 10, params, allowed) + .await + .unwrap_or_else(|e| panic!("{case}: native text_search: {e}")); + assert!(!over_native.is_empty(), "{case}: native found no hit"); + assert_eq!( + ranked(&over_native), + ranked(&over_pgwire), + "{case}: both clients return the same ids and scores" + ); + } + server.graceful_shutdown().await; } diff --git a/nodedb-client-tests/tests/typed_cells_round_trip.rs b/nodedb-client-tests/tests/typed_cells_round_trip.rs new file mode 100644 index 000000000..dd80d8945 --- /dev/null +++ b/nodedb-client-tests/tests/typed_cells_round_trip.rs @@ -0,0 +1,125 @@ +// SPDX-License-Identifier: BUSL-1.1 + +//! End-to-end cell encoding for `NodeDbRemote`. +//! +//! A strict document holds a `BYTEA`, a `VECTOR`, a `TIMESTAMP`, a +//! `DECIMAL` and a `UUID` field. Read over the extended protocol, where +//! `tokio_postgres` asks for binary results, every field comes back as its +//! own typed `Value`, equal to what was written. Read over the simple +//! protocol, every cell is its PostgreSQL text form. + +use nodedb_client::{NodeDb, NodeDbRemote, Value}; +use nodedb_test_support::pgwire_harness::TestServer; +use nodedb_types::NdbDateTime; + +const UUID: &str = "550e8400-e29b-41d4-a716-446655440000"; + +async fn remote(server: &TestServer) -> NodeDbRemote { + let conn_str = format!( + "host=127.0.0.1 port={} user=nodedb dbname=default", + server.pg_port + ); + NodeDbRemote::connect(&conn_str) + .await + .expect("pgwire connect to harness must succeed") +} + +/// Create the strict collection and write row `r1` with every typed field +/// set and row `r2` with every typed field NULL. +async fn seed(remote: &NodeDbRemote) { + remote + .execute_sql( + "CREATE COLLECTION typed_cells (id TEXT PRIMARY KEY, payload BYTEA, \ + embedding VECTOR(3), at TIMESTAMP, price DECIMAL(10,2), uid UUID) \ + WITH (engine='document_strict')", + &[], + ) + .await + .expect("CREATE COLLECTION typed_cells"); + // A string written to a BYTEA column is read as base64: `AQL/` is the + // bytes 0x01 0x02 0xff. + remote + .execute_sql( + &format!( + "INSERT INTO typed_cells (id, payload, embedding, at, price, uid) VALUES \ + ('r1', 'AQL/', '[0.5, 1.25, -2]', '2024-01-02T03:04:05Z', 12.50, '{UUID}')" + ), + &[], + ) + .await + .expect("INSERT the typed row"); + remote + .execute_sql("INSERT INTO typed_cells (id) VALUES ('r2')", &[]) + .await + .expect("INSERT the NULL row"); +} + +const SELECT_TYPED: &str = + "SELECT payload, embedding, at, price, uid FROM typed_cells WHERE id = $1"; + +#[tokio::test] +async fn remote_reads_bytes_vector_timestamp_decimal_and_uuid_typed() { + let server = TestServer::start().await; + let remote = remote(&server).await; + seed(&remote).await; + + let result = remote + .execute_sql(SELECT_TYPED, &[Value::String("r1".into())]) + .await + .expect("extended-protocol read of the typed row"); + assert_eq!( + result.columns, + vec!["payload", "embedding", "at", "price", "uid"] + ); + let at = NdbDateTime::parse("2024-01-02T03:04:05Z").expect("valid instant"); + assert_eq!( + result.rows, + vec![vec![ + Value::Bytes(vec![0x01, 0x02, 0xff]), + Value::Array(vec![ + Value::Float(0.5), + Value::Float(1.25), + Value::Float(-2.0) + ]), + Value::NaiveDateTime(at), + Value::Decimal("12.50".parse().expect("decimal literal")), + Value::Uuid(UUID.into()), + ]], + "every field reads back as the typed value written" + ); + + let nulls = remote + .execute_sql(SELECT_TYPED, &[Value::String("r2".into())]) + .await + .expect("extended-protocol read of the NULL row"); + assert_eq!(nulls.rows, vec![vec![Value::Null; 5]]); + + server.graceful_shutdown().await; +} + +#[tokio::test] +async fn simple_protocol_cells_are_postgres_text() { + let server = TestServer::start().await; + let remote = remote(&server).await; + seed(&remote).await; + + let result = remote + .execute_sql( + "SELECT payload, embedding, price, uid FROM typed_cells WHERE id = 'r1'", + &[], + ) + .await + .expect("simple-protocol read of the typed row"); + assert_eq!( + result.rows, + vec![vec![ + Value::String("\\x0102ff".into()), + Value::String("{0.5,1.25,-2}".into()), + Value::String("12.50".into()), + Value::String(UUID.into()), + ]], + "bytea is \\x hex, a vector is a {{...}} literal, numeric keeps its scale" + ); + + server.graceful_shutdown().await; +} diff --git a/nodedb-client-tests/tests/vector_search_allowed_ids_round_trip.rs b/nodedb-client-tests/tests/vector_search_allowed_ids_round_trip.rs new file mode 100644 index 000000000..42ae1b2c8 --- /dev/null +++ b/nodedb-client-tests/tests/vector_search_allowed_ids_round_trip.rs @@ -0,0 +1,142 @@ +// SPDX-License-Identifier: BUSL-1.1 + +//! End-to-end test that `NodeDb::vector_search` ranks only within +//! `allowed_ids` on both Origin clients. +//! +//! The nearest vectors to the query lie outside the allowed set. A +//! restriction applied after a top-k cut would return nothing (or fewer +//! than `k` rows). The restriction must apply before ranking: `k` rows come +//! back, every one from the allowed set, nearest first. + +use std::collections::HashSet; + +use nodedb_client::native::pool::PoolConfig; +use nodedb_client::{NativeClient, NodeDb, NodeDbRemote, SearchResult}; +use nodedb_test_support::pgwire_harness::TestServer; + +/// Seed `allowed_ids_vecs`: three near rows the query prefers and three far +/// rows, increasingly distant from the query `[0, 0]`. +async fn seed(remote: &NodeDbRemote) { + remote + .execute_sql( + "CREATE COLLECTION allowed_ids_vecs \ + FIELDS (id TEXT, embedding VECTOR(2)) \ + WITH (engine='vector', m=8, ef_construction=50)", + &[], + ) + .await + .expect("CREATE COLLECTION allowed_ids_vecs"); + for (id, x) in [ + ("near1", 0.0), + ("near2", 0.1), + ("near3", 0.2), + ("far1", 10.0), + ("far2", 20.0), + ("far3", 30.0), + ] { + remote + .execute_sql( + &format!( + "INSERT INTO allowed_ids_vecs (id, embedding) VALUES ('{id}', ARRAY[{x}, 0.0])" + ), + &[], + ) + .await + .unwrap_or_else(|e| panic!("seed {id}: {e}")); + } +} + +fn far_ids() -> HashSet { + ["far1", "far2", "far3"] + .iter() + .map(|s| s.to_string()) + .collect() +} + +fn ids(hits: &[SearchResult]) -> Vec<&str> { + hits.iter().map(|h| h.id.as_str()).collect() +} + +/// `k = 2` over the far rows returns the two nearest far rows, nearest first. +fn assert_ranked_within_allowed(client: &str, hits: &[SearchResult]) { + assert_eq!( + ids(hits), + vec!["far1", "far2"], + "{client}: k rows from the allowed set, nearest first" + ); +} + +#[tokio::test] +async fn vector_search_ranks_only_within_allowed_ids() { + let server = TestServer::start().await; + let remote = NodeDbRemote::connect(&format!( + "host=127.0.0.1 port={} user=nodedb dbname=default", + server.pg_port + )) + .await + .expect("pgwire connect to harness must succeed"); + seed(&remote).await; + let native = NativeClient::new(PoolConfig::new( + format!("127.0.0.1:{}", server.native_port), + nodedb_types::protocol::AuthMethod::Trust { + username: "nodedb".into(), + }, + )); + let allowed = far_ids(); + + // Unrestricted, the near rows win: the restriction below is what moves + // the result into the allowed set. + let open = remote + .vector_search("allowed_ids_vecs", &[0.0, 0.0], 2, None, None) + .await + .expect("remote unrestricted vector_search"); + assert_eq!(ids(&open), vec!["near1", "near2"]); + let open = native + .vector_search("allowed_ids_vecs", &[0.0, 0.0], 2, None, None) + .await + .expect("native unrestricted vector_search"); + assert_eq!(ids(&open), vec!["near1", "near2"], "native unrestricted"); + + let hits = remote + .vector_search("allowed_ids_vecs", &[0.0, 0.0], 2, None, Some(&allowed)) + .await + .expect("remote vector_search with allowed_ids"); + assert_ranked_within_allowed("remote", &hits); + + let hits = native + .vector_search("allowed_ids_vecs", &[0.0, 0.0], 2, None, Some(&allowed)) + .await + .expect("native vector_search with allowed_ids"); + assert_ranked_within_allowed("native", &hits); + + // An empty allowed set admits no row on either client. + let none: HashSet = HashSet::new(); + for (client, hits) in [ + ( + "remote", + remote + .vector_search("allowed_ids_vecs", &[0.0, 0.0], 2, None, Some(&none)) + .await + .expect("remote empty allowed set"), + ), + ( + "native", + native + .vector_search("allowed_ids_vecs", &[0.0, 0.0], 2, None, Some(&none)) + .await + .expect("native empty allowed set"), + ), + ] { + assert!(hits.is_empty(), "{client}: got {hits:?}"); + } + + // An allowed id that names no row restricts to the ids that do. + let with_ghost: HashSet = ["far3", "ghost"].iter().map(|s| s.to_string()).collect(); + let hits = native + .vector_search("allowed_ids_vecs", &[0.0, 0.0], 2, None, Some(&with_ghost)) + .await + .expect("native vector_search with an unbound id"); + assert_eq!(ids(&hits), vec!["far3"]); + + server.graceful_shutdown().await; +} diff --git a/nodedb-client-tests/tests/vector_search_field_round_trip.rs b/nodedb-client-tests/tests/vector_search_field_round_trip.rs index c49336524..b0d3e68aa 100644 --- a/nodedb-client-tests/tests/vector_search_field_round_trip.rs +++ b/nodedb-client-tests/tests/vector_search_field_round_trip.rs @@ -27,13 +27,8 @@ async fn vector_search_field_must_not_silently_delegate_to_unfielded() { // Provision a vector-primary collection with an explicit // `body_embedding` column and seed it so the search has real data - // to return. `primary='vector'` is the server contract that wires - // the named-field HNSW index — `vector_distance(body_embedding, - // ARRAY[...])` resolves through it. Plain `engine='vector'` - // without `primary='vector'` would store the row only in the - // document body and the named-field HNSW index would never get - // populated, masking a broken trait default behind an empty result - // set. + // to return. `primary='vector'` wires the named-field HNSW index — + // `vector_distance(body_embedding, ARRAY[...])` resolves through it. remote .execute_sql( "CREATE COLLECTION embeddings_multi \ diff --git a/nodedb-client-tests/tests/vector_search_metadata_filter_window.rs b/nodedb-client-tests/tests/vector_search_metadata_filter_window.rs new file mode 100644 index 000000000..134921ece --- /dev/null +++ b/nodedb-client-tests/tests/vector_search_metadata_filter_window.rs @@ -0,0 +1,140 @@ +// SPDX-License-Identifier: BUSL-1.1 + +//! End-to-end test that `NodeDb::vector_search` applies a metadata filter +//! before the top-k cut, identically on both Origin clients. +//! +//! The nearest rows to the query all fail the filter, and they outnumber +//! the search's initial candidate window. A filter applied to a fixed +//! window would return fewer than `k` rows. The candidates must narrow to +//! the matching rows before the cut: `k` rows come back, every one +//! matching the filter, nearest first, the same rows on both clients. + +use nodedb_client::native::pool::PoolConfig; +use nodedb_client::{MetadataFilter, NativeClient, NodeDb, NodeDbRemote, SearchResult, Value}; +use nodedb_test_support::pgwire_harness::TestServer; + +/// Rows near the query that fail the filter. More than the initial window +/// of a `k = 2` filtered search. +const NEAR_ROWS: usize = 40; + +/// Seed `meta_filter_vecs`: `NEAR_ROWS` rows close to the query `[0, 0]` +/// in category `skip`, then three rows far from it in category `keep`. +async fn seed(remote: &NodeDbRemote) { + remote + .execute_sql( + "CREATE COLLECTION meta_filter_vecs \ + FIELDS (id TEXT, embedding VECTOR(2), category TEXT) \ + WITH (engine='vector', m=8, ef_construction=50)", + &[], + ) + .await + .expect("CREATE COLLECTION meta_filter_vecs"); + let near = (0..NEAR_ROWS).map(|i| (format!("near{i}"), i as f64 * 0.01, "skip")); + let far = [("far1", 10.0), ("far2", 20.0), ("far3", 30.0)] + .into_iter() + .map(|(id, x)| (id.to_string(), x, "keep")); + for (id, x, category) in near.chain(far) { + remote + .execute_sql( + &format!( + "INSERT INTO meta_filter_vecs (id, embedding, category) \ + VALUES ('{id}', ARRAY[{x}, 0.0], '{category}')" + ), + &[], + ) + .await + .unwrap_or_else(|e| panic!("seed {id}: {e}")); + } +} + +fn ids(hits: &[SearchResult]) -> Vec<&str> { + hits.iter().map(|h| h.id.as_str()).collect() +} + +#[tokio::test] +async fn vector_search_metadata_filter_narrows_candidates_before_top_k() { + let server = TestServer::start().await; + let remote = NodeDbRemote::connect(&format!( + "host=127.0.0.1 port={} user=nodedb dbname=default", + server.pg_port + )) + .await + .expect("pgwire connect to harness must succeed"); + seed(&remote).await; + let native = NativeClient::new(PoolConfig::new( + format!("127.0.0.1:{}", server.native_port), + nodedb_types::protocol::AuthMethod::Trust { + username: "nodedb".into(), + }, + )); + + // Unfiltered, the near rows win: the filter below is what moves the + // result to the far rows. + let open = remote + .vector_search("meta_filter_vecs", &[0.0, 0.0], 2, None, None) + .await + .expect("remote unfiltered vector_search"); + assert_eq!(ids(&open), vec!["near0", "near1"]); + + let keep = MetadataFilter::eq("category", Value::String("keep".into())); + let remote_hits = remote + .vector_search("meta_filter_vecs", &[0.0, 0.0], 2, Some(&keep), None) + .await + .expect("remote vector_search with a metadata filter"); + let native_hits = native + .vector_search("meta_filter_vecs", &[0.0, 0.0], 2, Some(&keep), None) + .await + .expect("native vector_search with a metadata filter"); + assert_eq!( + ids(&remote_hits), + vec!["far1", "far2"], + "remote: k matching rows, nearest first" + ); + assert_eq!( + ids(&native_hits), + ids(&remote_hits), + "native and remote return the same rows" + ); + + // A compound filter matches the same rows on both clients. + let far_two = MetadataFilter::and(vec![ + keep.clone(), + MetadataFilter::Not(Box::new(MetadataFilter::eq( + "id", + Value::String("far1".into()), + ))), + ]); + let remote_hits = remote + .vector_search("meta_filter_vecs", &[0.0, 0.0], 2, Some(&far_two), None) + .await + .expect("remote vector_search with a compound filter"); + let native_hits = native + .vector_search("meta_filter_vecs", &[0.0, 0.0], 2, Some(&far_two), None) + .await + .expect("native vector_search with a compound filter"); + assert_eq!(ids(&remote_hits), vec!["far2", "far3"]); + assert_eq!(ids(&native_hits), ids(&remote_hits)); + + // A filter no row matches returns no row on either client. + let none = MetadataFilter::eq("category", Value::String("absent".into())); + for (client, hits) in [ + ( + "remote", + remote + .vector_search("meta_filter_vecs", &[0.0, 0.0], 2, Some(&none), None) + .await + .expect("remote vector_search matching nothing"), + ), + ( + "native", + native + .vector_search("meta_filter_vecs", &[0.0, 0.0], 2, Some(&none), None) + .await + .expect("native vector_search matching nothing"), + ), + ] { + assert!(hits.is_empty(), "{client}: got {hits:?}"); + } + + server.graceful_shutdown().await; +} diff --git a/nodedb-client/Cargo.toml b/nodedb-client/Cargo.toml index e036aca22..90dfb3f9a 100644 --- a/nodedb-client/Cargo.toml +++ b/nodedb-client/Cargo.toml @@ -14,7 +14,7 @@ categories = ["database", "api-bindings"] [features] default = [] -remote = ["tokio-postgres", "tokio", "serde_json"] +remote = ["tokio-postgres", "tokio", "serde_json", "rust_decimal"] native = ["tokio", "zerompk", "serde_json", "tokio-rustls", "rustls-pemfile"] [dependencies] @@ -28,6 +28,7 @@ tracing = { workspace = true } tokio-postgres = { workspace = true, optional = true, features = ["with-serde_json-1"] } tokio = { workspace = true, optional = true, features = ["rt", "sync"] } serde_json = { workspace = true, optional = true } +rust_decimal = { workspace = true, optional = true } sonic-rs = { workspace = true } # Native protocol client dependencies (feature-gated) @@ -37,3 +38,4 @@ rustls-pemfile = { workspace = true, optional = true } [dev-dependencies] tokio = { workspace = true, features = ["rt", "macros"] } +rust_decimal = { workspace = true } diff --git a/nodedb-client/src/document_identity.rs b/nodedb-client/src/document_identity.rs new file mode 100644 index 000000000..5a4f78063 --- /dev/null +++ b/nodedb-client/src/document_identity.rs @@ -0,0 +1,166 @@ +// SPDX-License-Identifier: Apache-2.0 + +//! The identity column of a stored schemaless document. +//! +//! A collection's identity column is its declared primary key, else `id`. +//! The server catalog owns that resolution. Every write stores the document +//! id under the identity column: +//! - A native put sends the id beside the body. The server writes it under +//! the identity column and refuses a body that names another id there. +//! - A pgwire put names the identity column in its `INSERT`. The client +//! reads the column from `DESCRIBE`. +//! +//! A read returns the implicit `id` cell as `Document::id`, never as one of +//! `Document::fields`. A declared key column is a declared field: it reads +//! back as a field that holds the document id. + +use nodedb_types::DEFAULT_IDENTITY_COLUMN; +use nodedb_types::error::{NodeDbError, NodeDbResult}; +use nodedb_types::value::Value; + +/// `DESCRIBE` output column naming each field. +const DESCRIBE_FIELD_COLUMN: &str = "field"; +/// `DESCRIBE` output column marking the key field. +const DESCRIBE_KEY_COLUMN: &str = "primary_key"; + +/// Whether a stored cell is the document's implicit identity cell. +/// +/// An `id` cell that holds any other value is a field the writer set, and +/// stays a field. +pub(crate) fn is_identity_cell(name: &str, value: &Value, document_id: &str) -> bool { + name == DEFAULT_IDENTITY_COLUMN && matches!(value, Value::String(s) if s == document_id) +} + +/// The identity column named by a `DESCRIBE ` result. +/// +/// The result holds one row per field, with a `primary_key` cell marking the +/// key. `is_key` decodes that cell in the transport's own shape. Exactly one +/// row must be the key: none or several is an error naming the collection. +pub(crate) fn identity_column_from_describe( + collection: &str, + columns: &[String], + rows: &[Vec], + is_key: fn(&Value) -> Option, +) -> NodeDbResult { + let index = |name: &str| { + columns.iter().position(|c| c == name).ok_or_else(|| { + describe_error( + collection, + format!("result has no '{name}' column; columns are {columns:?}"), + ) + }) + }; + let field_idx = index(DESCRIBE_FIELD_COLUMN)?; + let key_idx = index(DESCRIBE_KEY_COLUMN)?; + let mut keys = Vec::new(); + for row in rows { + let cell = row.get(key_idx).ok_or_else(|| { + describe_error(collection, format!("row {row:?} has no primary_key cell")) + })?; + let marked = is_key(cell).ok_or_else(|| { + describe_error( + collection, + format!("primary_key cell {cell:?} is not a bool"), + ) + })?; + if !marked { + continue; + } + match row.get(field_idx) { + Some(Value::String(field)) => keys.push(field.clone()), + other => { + return Err(describe_error( + collection, + format!("key row names no field: {other:?}"), + )); + } + } + } + match keys.as_slice() { + [key] => Ok(key.clone()), + [] => Err(describe_error(collection, "no field is the primary key")), + several => Err(describe_error( + collection, + format!("several fields are the primary key: {several:?}"), + )), + } +} + +fn describe_error(collection: &str, detail: impl std::fmt::Display) -> NodeDbError { + NodeDbError::serialization( + "describe", + format!("identity column of '{collection}': {detail}"), + ) +} + +#[cfg(test)] +mod tests { + use super::*; + + fn text_bool(cell: &Value) -> Option { + match cell { + Value::String(s) if s == "t" => Some(true), + Value::String(s) if s == "f" => Some(false), + _ => None, + } + } + + fn describe(rows: &[(&str, &str)]) -> (Vec, Vec>) { + let columns = ["field", "type", "nullable", "primary_key"] + .iter() + .map(|c| c.to_string()) + .collect(); + let rows = rows + .iter() + .map(|(field, key)| { + vec![ + Value::String((*field).into()), + Value::String("TEXT".into()), + Value::String("true".into()), + Value::String((*key).into()), + ] + }) + .collect(); + (columns, rows) + } + + #[test] + fn only_the_matching_id_cell_is_the_identity() { + assert!(is_identity_cell("id", &Value::String("d1".into()), "d1")); + assert!(!is_identity_cell( + "id", + &Value::String("other".into()), + "d1" + )); + assert!(!is_identity_cell("id", &Value::Integer(1), "d1")); + assert!(!is_identity_cell("body", &Value::String("d1".into()), "d1")); + } + + #[test] + fn the_describe_key_row_names_the_identity_column() { + let (columns, rows) = describe(&[("id", "f"), ("sku", "t"), ("name", "f")]); + let column = + identity_column_from_describe("c", &columns, &rows, text_bool).expect("one key row"); + assert_eq!(column, "sku"); + } + + #[test] + fn a_describe_result_without_one_key_row_is_an_error() { + let (columns, rows) = describe(&[("id", "f"), ("name", "f")]); + let err = + identity_column_from_describe("c", &columns, &rows, text_bool).expect_err("no key row"); + assert!(err.to_string().contains("'c'"), "{err}"); + + let (columns, rows) = describe(&[("a", "t"), ("b", "t")]); + assert!(identity_column_from_describe("c", &columns, &rows, text_bool).is_err()); + + let (columns, rows) = describe(&[("a", "maybe")]); + let err = + identity_column_from_describe("c", &columns, &rows, text_bool).expect_err("not a bool"); + assert!(err.to_string().contains("not a bool"), "{err}"); + + let err = identity_column_from_describe("c", &columns[..3], &rows, text_bool) + .expect_err("no primary_key column"); + assert!(err.to_string().contains("'primary_key'"), "{err}"); + } +} diff --git a/nodedb-client/src/graph_dsl/decode.rs b/nodedb-client/src/graph_dsl/decode.rs new file mode 100644 index 000000000..994d46e4c --- /dev/null +++ b/nodedb-client/src/graph_dsl/decode.rs @@ -0,0 +1,292 @@ +// SPDX-License-Identifier: Apache-2.0 + +//! Strict decoders for the `GRAPH TRAVERSE` and `GRAPH PATH` results. Both +//! clients use them. +//! +//! The server answers either statement with one `result` column holding one +//! JSON text cell: +//! +//! - `GRAPH TRAVERSE`: `{"nodes": [{"id", "depth"}], "edges": [{"from", +//! "to", "label"[, "properties"]}]}`. +//! - `GRAPH PATH`: `["src", …, "dst"]`, or `[]` when no path exists. +//! +//! Any other shape is an error naming what is wrong. A dropped row would +//! return a partial subgraph as a complete one. + +use std::collections::HashMap; + +use nodedb_types::error::{NodeDbError, NodeDbResult}; +use nodedb_types::id::{EdgeId, NodeId}; +use nodedb_types::result::{SubGraph, SubGraphEdge, SubGraphNode}; +use nodedb_types::value::Value; +use serde_json::{Map, Value as Json}; + +/// Decode a `GRAPH TRAVERSE` result. +pub(crate) fn decode_traverse_result( + columns: &[String], + rows: &[Vec], +) -> NodeDbResult { + let json = result_json(columns, rows, "GRAPH TRAVERSE")?; + let Json::Object(mut top) = json else { + return Err(malformed("GRAPH TRAVERSE result is not a JSON object")); + }; + let nodes = take_array(&mut top, "nodes")? + .into_iter() + .enumerate() + .map(|(index, node)| decode_node(index, node)) + .collect::>>()?; + let edges = take_array(&mut top, "edges")? + .into_iter() + .enumerate() + .map(|(index, edge)| decode_edge(index, edge)) + .collect::>>()?; + Ok(SubGraph { nodes, edges }) +} + +/// Decode a `GRAPH PATH` result. `None` when no path exists. +pub(crate) fn decode_path_result( + columns: &[String], + rows: &[Vec], +) -> NodeDbResult>> { + let json = result_json(columns, rows, "GRAPH PATH")?; + let Json::Array(items) = json else { + return Err(malformed("GRAPH PATH result is not a JSON array")); + }; + if items.is_empty() { + return Ok(None); + } + items + .into_iter() + .enumerate() + .map(|(index, item)| match item { + Json::String(id) => node_id(id, &format!("GRAPH PATH node {index}")), + other => Err(malformed(format!( + "GRAPH PATH node {index} is not a string: {other}" + ))), + }) + .collect::>>() + .map(Some) +} + +/// The JSON text of the one `result` cell. +fn result_json(columns: &[String], rows: &[Vec], statement: &str) -> NodeDbResult { + if columns.len() != 1 || columns[0] != "result" { + return Err(malformed(format!( + "{statement} answered columns {columns:?}, expected [\"result\"]" + ))); + } + let [row] = rows else { + return Err(malformed(format!( + "{statement} answered {} rows, expected 1", + rows.len() + ))); + }; + let [Value::String(text)] = row.as_slice() else { + return Err(malformed(format!( + "{statement} result row is not one text cell: {row:?}" + ))); + }; + sonic_rs::from_str(text).map_err(|e| malformed(format!("{statement} result is not JSON: {e}"))) +} + +fn take_array(top: &mut Map, key: &str) -> NodeDbResult> { + match top.remove(key) { + Some(Json::Array(items)) => Ok(items), + Some(other) => Err(malformed(format!( + "GRAPH TRAVERSE `{key}` is not an array: {other}" + ))), + None => Err(malformed(format!("GRAPH TRAVERSE result lacks `{key}`"))), + } +} + +fn decode_node(index: usize, node: Json) -> NodeDbResult { + let what = format!("GRAPH TRAVERSE node {index}"); + let Json::Object(mut fields) = node else { + return Err(malformed(format!("{what} is not an object"))); + }; + let id = node_id(take_string(&mut fields, "id", &what)?, &what)?; + let depth = match fields.remove("depth") { + Some(Json::Number(n)) => n + .as_u64() + .and_then(|d| u8::try_from(d).ok()) + .ok_or_else(|| malformed(format!("{what} depth {n} is not an integer in 0..=255")))?, + Some(other) => { + return Err(malformed(format!("{what} depth is not a number: {other}"))); + } + None => return Err(malformed(format!("{what} lacks `depth`"))), + }; + Ok(SubGraphNode { + id, + depth, + properties: HashMap::new(), + }) +} + +fn decode_edge(index: usize, edge: Json) -> NodeDbResult { + let what = format!("GRAPH TRAVERSE edge {index}"); + let Json::Object(mut fields) = edge else { + return Err(malformed(format!("{what} is not an object"))); + }; + let from = node_id(take_string(&mut fields, "from", &what)?, &what)?; + let to = node_id(take_string(&mut fields, "to", &what)?, &what)?; + let label = take_string(&mut fields, "label", &what)?; + let properties = match fields.remove("properties") { + None => HashMap::new(), + Some(Json::Object(properties)) => properties + .into_iter() + .map(|(name, value)| (name, Value::from(value))) + .collect(), + Some(other) => { + return Err(malformed(format!( + "{what} properties are not an object: {other}" + ))); + } + }; + let id = EdgeId::try_first(from.clone(), to.clone(), label.clone()).map_err(|e| { + malformed(format!( + "{what} label '{label}' is not a valid edge label: {e}" + )) + })?; + Ok(SubGraphEdge { + id, + from, + to, + label, + properties, + }) +} + +fn take_string(fields: &mut Map, key: &str, what: &str) -> NodeDbResult { + match fields.remove(key) { + Some(Json::String(s)) => Ok(s), + Some(other) => Err(malformed(format!( + "{what} `{key}` is not a string: {other}" + ))), + None => Err(malformed(format!("{what} lacks `{key}`"))), + } +} + +fn node_id(id: String, what: &str) -> NodeDbResult { + NodeId::try_new(id.clone()) + .map_err(|e| malformed(format!("{what} id '{id}' is not a valid node id: {e}"))) +} + +fn malformed(detail: impl std::fmt::Display) -> NodeDbError { + NodeDbError::serialization("json", detail) +} + +#[cfg(test)] +mod tests { + use super::*; + + fn result(json: &str) -> (Vec, Vec>) { + ( + vec!["result".to_string()], + vec![vec![Value::String(json.to_string())]], + ) + } + + fn traverse(json: &str) -> NodeDbResult { + let (columns, rows) = result(json); + decode_traverse_result(&columns, &rows) + } + + fn path(json: &str) -> NodeDbResult>> { + let (columns, rows) = result(json); + decode_path_result(&columns, &rows) + } + + #[test] + fn a_well_formed_traversal_decodes_with_properties() { + let sg = traverse( + r#"{"nodes":[{"id":"a","depth":0},{"id":"b","depth":1}], + "edges":[{"from":"a","to":"b","label":"KNOWS","properties":{"since":2020,"w":0.5}}, + {"from":"b","to":"a","label":"KNOWS"}]}"#, + ) + .expect("decodes"); + assert_eq!(sg.nodes.len(), 2); + assert_eq!(sg.nodes[1].id.as_str(), "b"); + assert_eq!(sg.nodes[1].depth, 1); + assert_eq!(sg.edges.len(), 2); + assert_eq!(sg.edges[0].label, "KNOWS"); + assert_eq!(sg.edges[0].from.as_str(), "a"); + assert_eq!( + sg.edges[0].properties.get("since"), + Some(&Value::Integer(2020)) + ); + assert_eq!(sg.edges[0].properties.get("w"), Some(&Value::Float(0.5))); + assert!(sg.edges[1].properties.is_empty()); + } + + #[test] + fn an_absent_start_decodes_as_the_empty_subgraph() { + let sg = traverse(r#"{"nodes":[],"edges":[]}"#).expect("decodes"); + assert!(sg.nodes.is_empty()); + assert!(sg.edges.is_empty()); + } + + #[test] + fn every_malformed_traversal_shape_is_an_error() { + for json in [ + r#"[]"#, + r#"{"edges":[]}"#, + r#"{"nodes":[]}"#, + r#"{"nodes":{},"edges":[]}"#, + r#"{"nodes":["a"],"edges":[]}"#, + r#"{"nodes":[{"depth":0}],"edges":[]}"#, + r#"{"nodes":[{"id":7,"depth":0}],"edges":[]}"#, + r#"{"nodes":[{"id":"","depth":0}],"edges":[]}"#, + r#"{"nodes":[{"id":"a"}],"edges":[]}"#, + r#"{"nodes":[{"id":"a","depth":1.5}],"edges":[]}"#, + r#"{"nodes":[{"id":"a","depth":-1}],"edges":[]}"#, + r#"{"nodes":[{"id":"a","depth":256}],"edges":[]}"#, + r#"{"nodes":[],"edges":[{"to":"b","label":"L"}]}"#, + r#"{"nodes":[],"edges":[{"from":"a","label":"L"}]}"#, + r#"{"nodes":[],"edges":[{"from":"a","to":"b"}]}"#, + r#"{"nodes":[],"edges":[{"from":"a","to":"b","label":3}]}"#, + r#"{"nodes":[],"edges":[{"from":"a","to":"b","label":""}]}"#, + r#"{"nodes":[],"edges":[{"from":"a","to":"b","label":"L","properties":[1]}]}"#, + "not json", + ] { + assert!(traverse(json).is_err(), "{json} must be refused"); + } + } + + #[test] + fn a_malformed_row_names_its_index() { + let error = traverse(r#"{"nodes":[{"id":"a","depth":0},{"id":"b"}],"edges":[]}"#) + .expect_err("node 1 lacks depth"); + assert!(error.to_string().contains("node 1"), "{error}"); + let error = traverse( + r#"{"nodes":[],"edges":[{"from":"a","to":"b","label":"L"},{"from":"a","to":"b"}]}"#, + ) + .expect_err("edge 1 lacks a label"); + assert!(error.to_string().contains("edge 1"), "{error}"); + } + + #[test] + fn the_result_column_shape_is_checked() { + let rows = vec![vec![Value::String("{}".into())]]; + assert!(decode_traverse_result(&["node_id".to_string()], &rows).is_err()); + assert!(decode_traverse_result(&["result".to_string()], &[]).is_err()); + assert!( + decode_traverse_result(&["result".to_string()], &[vec![Value::Integer(1)]]).is_err() + ); + assert!(decode_traverse_result(&[], &[]).is_err()); + } + + #[test] + fn a_path_decodes_and_an_empty_path_is_none() { + let found = path(r#"["a","b","c"]"#).expect("decodes").expect("a path"); + let ids: Vec<&str> = found.iter().map(NodeId::as_str).collect(); + assert_eq!(ids, vec!["a", "b", "c"]); + assert_eq!(path("[]").expect("decodes"), None); + } + + #[test] + fn a_malformed_path_is_an_error() { + for json in [r#"{"path":[]}"#, r#"["a",1]"#, r#"["a",""]"#, "nope"] { + assert!(path(json).is_err(), "{json} must be refused"); + } + } +} diff --git a/nodedb-client/src/graph_dsl/edge_predicate.rs b/nodedb-client/src/graph_dsl/edge_predicate.rs new file mode 100644 index 000000000..792b25ceb --- /dev/null +++ b/nodedb-client/src/graph_dsl/edge_predicate.rs @@ -0,0 +1,299 @@ +// SPDX-License-Identifier: Apache-2.0 + +//! Render `EdgeFilter::property_filters` as the `EDGE WHERE` clause of +//! `GRAPH TRAVERSE` / `GRAPH PATH`. +//! +//! Every property name renders quoted, so its case and any reserved word +//! survive. Every term renders parenthesised, so operator precedence never +//! depends on the server's grammar. + +use nodedb_types::error::{NodeDbError, NodeDbResult}; +use nodedb_types::filter::MetadataFilter; +use nodedb_types::value::Value; + +use crate::sql_escape::{quote_identifier, quote_string_literal}; + +/// ` EDGE WHERE ` for `filters`, AND-ed. Empty filters render nothing. +pub(crate) fn edge_where_clause(filters: &[MetadataFilter]) -> NodeDbResult { + if filters.is_empty() { + return Ok(String::new()); + } + let mut out = String::from(" EDGE WHERE "); + for (i, filter) in filters.iter().enumerate() { + if i > 0 { + out.push_str(" AND "); + } + out.push('('); + render(filter, &mut out)?; + out.push(')'); + } + Ok(out) +} + +fn render(filter: &MetadataFilter, out: &mut String) -> NodeDbResult<()> { + match filter { + MetadataFilter::Eq { field, value } if value.is_null() => { + null_test(field, "IS NULL", out); + Ok(()) + } + MetadataFilter::Ne { field, value } if value.is_null() => { + null_test(field, "IS NOT NULL", out); + Ok(()) + } + MetadataFilter::Eq { field, value } => compare(field, "=", value, out), + MetadataFilter::Ne { field, value } => compare(field, "<>", value, out), + MetadataFilter::Gt { field, value } => compare(field, ">", value, out), + MetadataFilter::Gte { field, value } => compare(field, ">=", value, out), + MetadataFilter::Lt { field, value } => compare(field, "<", value, out), + MetadataFilter::Lte { field, value } => compare(field, "<=", value, out), + // An empty IN admits nothing. An empty NOT IN admits everything. + MetadataFilter::In { values, .. } if values.is_empty() => { + out.push_str("FALSE"); + Ok(()) + } + MetadataFilter::NotIn { values, .. } if values.is_empty() => { + out.push_str("TRUE"); + Ok(()) + } + MetadataFilter::In { field, values } => list(field, "IN", values, out), + MetadataFilter::NotIn { field, values } => list(field, "NOT IN", values, out), + MetadataFilter::And(children) => group(children, " AND ", "TRUE", out), + MetadataFilter::Or(children) => group(children, " OR ", "FALSE", out), + MetadataFilter::Not(inner) => { + out.push_str("NOT ("); + render(inner, out)?; + out.push(')'); + Ok(()) + } + // `MetadataFilter` is `#[non_exhaustive]`: a variant added later has + // no rendering until one is written here. + other => Err(NodeDbError::bad_request(format!( + "edge filter {other:?} has no GRAPH SQL form: use Eq, Ne, Gt, Gte, Lt, Lte, In, \ + NotIn, And, Or or Not" + ))), + } +} + +fn null_test(field: &str, test: &str, out: &mut String) { + out.push_str("e_identifier(field)); + out.push(' '); + out.push_str(test); +} + +fn compare(field: &str, op: &str, value: &Value, out: &mut String) -> NodeDbResult<()> { + out.push_str("e_identifier(field)); + out.push(' '); + out.push_str(op); + out.push(' '); + literal(value, out) +} + +fn list(field: &str, op: &str, values: &[Value], out: &mut String) -> NodeDbResult<()> { + out.push_str("e_identifier(field)); + out.push(' '); + out.push_str(op); + out.push_str(" ("); + for (i, value) in values.iter().enumerate() { + if i > 0 { + out.push_str(", "); + } + literal(value, out)?; + } + out.push(')'); + Ok(()) +} + +fn group( + children: &[MetadataFilter], + joiner: &str, + empty: &str, + out: &mut String, +) -> NodeDbResult<()> { + if children.is_empty() { + out.push_str(empty); + return Ok(()); + } + for (i, child) in children.iter().enumerate() { + if i > 0 { + out.push_str(joiner); + } + out.push('('); + render(child, out)?; + out.push(')'); + } + Ok(()) +} + +fn literal(value: &Value, out: &mut String) -> NodeDbResult<()> { + match value { + Value::Null => out.push_str("NULL"), + Value::Bool(b) => out.push_str(if *b { "TRUE" } else { "FALSE" }), + Value::Integer(i) => out.push_str(&i.to_string()), + // `{:?}` keeps a `.` or an exponent, so the server reads a float back. + Value::Float(f) if f.is_finite() => out.push_str(&format!("{f:?}")), + Value::String(s) | Value::Uuid(s) | Value::Ulid(s) | Value::Regex(s) => { + out.push_str("e_string_literal(s)) + } + Value::DateTime(dt) | Value::NaiveDateTime(dt) => { + out.push_str("e_string_literal(&dt.to_iso8601())) + } + Value::Decimal(d) => out.push_str(&d.to_string()), + other => { + return Err(NodeDbError::bad_request(format!( + "edge filter value {other:?} has no GRAPH SQL literal: use null, bool, a finite \ + number, string or timestamp" + ))); + } + } + Ok(()) +} + +#[cfg(test)] +mod tests { + use super::*; + + fn clause(filters: Vec) -> String { + edge_where_clause(&filters).expect("renders") + } + + fn cmp(field: &str, value: Value, ctor: fn(String, Value) -> MetadataFilter) -> MetadataFilter { + ctor(field.to_string(), value) + } + + #[test] + fn no_filters_render_nothing() { + assert_eq!(clause(Vec::new()), ""); + } + + #[test] + fn every_comparison_renders_with_a_quoted_name() { + let cases: [(MetadataFilter, &str); 6] = [ + ( + cmp("score", Value::Integer(5), |field, value| { + MetadataFilter::Eq { field, value } + }), + r#" EDGE WHERE ("score" = 5)"#, + ), + ( + cmp("score", Value::Integer(5), |field, value| { + MetadataFilter::Ne { field, value } + }), + r#" EDGE WHERE ("score" <> 5)"#, + ), + ( + cmp("score", Value::Float(2.5), |field, value| { + MetadataFilter::Gt { field, value } + }), + r#" EDGE WHERE ("score" > 2.5)"#, + ), + ( + cmp("score", Value::Float(5.0), |field, value| { + MetadataFilter::Gte { field, value } + }), + r#" EDGE WHERE ("score" >= 5.0)"#, + ), + ( + cmp("score", Value::Integer(-3), |field, value| { + MetadataFilter::Lt { field, value } + }), + r#" EDGE WHERE ("score" < -3)"#, + ), + ( + cmp("Score", Value::Float(1e-7), |field, value| { + MetadataFilter::Lte { field, value } + }), + r#" EDGE WHERE ("Score" <= 1e-7)"#, + ), + ]; + for (filter, expected) in cases { + assert_eq!(clause(vec![filter]), expected); + } + } + + #[test] + fn null_comparisons_render_null_tests() { + assert_eq!( + clause(vec![ + MetadataFilter::eq("gone", Value::Null), + MetadataFilter::Ne { + field: "here".into(), + value: Value::Null, + }, + ]), + r#" EDGE WHERE ("gone" IS NULL) AND ("here" IS NOT NULL)"# + ); + } + + #[test] + fn lists_and_empty_lists() { + assert_eq!( + clause(vec![MetadataFilter::In { + field: "kind".into(), + values: vec![Value::from("road"), Value::from("rail")], + }]), + r#" EDGE WHERE ("kind" IN ('road', 'rail'))"# + ); + assert_eq!( + clause(vec![MetadataFilter::NotIn { + field: "tag".into(), + values: vec![Value::Integer(1), Value::Bool(true)], + }]), + r#" EDGE WHERE ("tag" NOT IN (1, TRUE))"# + ); + assert_eq!( + clause(vec![MetadataFilter::In { + field: "kind".into(), + values: Vec::new(), + }]), + " EDGE WHERE (FALSE)" + ); + assert_eq!( + clause(vec![MetadataFilter::NotIn { + field: "kind".into(), + values: Vec::new(), + }]), + " EDGE WHERE (TRUE)" + ); + } + + #[test] + fn groups_nest_and_empties_render_constants() { + assert_eq!( + clause(vec![MetadataFilter::Not(Box::new(MetadataFilter::Or( + vec![ + MetadataFilter::eq("a", 1i64), + MetadataFilter::And(vec![ + MetadataFilter::eq("b", 2i64), + MetadataFilter::eq("c", 3i64), + ]), + ] + )))]), + r#" EDGE WHERE (NOT (("a" = 1) OR (("b" = 2) AND ("c" = 3))))"# + ); + assert_eq!( + clause(vec![MetadataFilter::And(Vec::new())]), + " EDGE WHERE (TRUE)" + ); + assert_eq!( + clause(vec![MetadataFilter::Or(Vec::new())]), + " EDGE WHERE (FALSE)" + ); + } + + #[test] + fn names_and_strings_escape_their_quotes() { + assert_eq!( + clause(vec![MetadataFilter::eq("we\"ird", "O'Reilly")]), + r#" EDGE WHERE ("we""ird" = 'O''Reilly')"# + ); + } + + #[test] + fn bytes_and_non_finite_floats_are_refused() { + assert!(edge_where_clause(&[MetadataFilter::eq("b", Value::Bytes(vec![1]))]).is_err()); + assert!(edge_where_clause(&[MetadataFilter::eq("f", Value::Float(f64::NAN))]).is_err()); + assert!( + edge_where_clause(&[MetadataFilter::eq("f", Value::Float(f64::INFINITY))]).is_err() + ); + } +} diff --git a/nodedb-client/src/graph_dsl/mod.rs b/nodedb-client/src/graph_dsl/mod.rs new file mode 100644 index 000000000..3535685ea --- /dev/null +++ b/nodedb-client/src/graph_dsl/mod.rs @@ -0,0 +1,14 @@ +// SPDX-License-Identifier: Apache-2.0 + +//! Shared graph-DSL builders and result decoders used by both the native and +//! the remote clients: one implementation per concern, no per-transport +//! duplicates. + +pub(crate) mod decode; +pub(crate) mod edge_predicate; +pub(crate) mod pagerank; +pub(crate) mod traverse; + +pub(crate) use decode::{decode_path_result, decode_traverse_result}; +pub(crate) use pagerank::{build_pagerank_sql, parse_pagerank_rows}; +pub(crate) use traverse::{build_graph_path_sql, build_graph_traverse_sql}; diff --git a/nodedb-client/src/graph_dsl.rs b/nodedb-client/src/graph_dsl/pagerank.rs similarity index 93% rename from nodedb-client/src/graph_dsl.rs rename to nodedb-client/src/graph_dsl/pagerank.rs index a70a7a9a5..54feb6fbb 100644 --- a/nodedb-client/src/graph_dsl.rs +++ b/nodedb-client/src/graph_dsl/pagerank.rs @@ -1,10 +1,6 @@ // SPDX-License-Identifier: Apache-2.0 -//! Shared graph-DSL builders and result parsers used by both the native and -//! the remote clients — one implementation per concern, no per-transport -//! duplicates. Both transports ultimately speak the same `GRAPH ALGO …` DSL -//! and receive the same `(columns, rows)` shape, so the SQL construction and -//! row decoding live here once. +//! `GRAPH ALGO PAGERANK` statement builder and result decoder. use std::collections::HashMap; diff --git a/nodedb-client/src/graph_dsl/traverse.rs b/nodedb-client/src/graph_dsl/traverse.rs new file mode 100644 index 000000000..cb7713118 --- /dev/null +++ b/nodedb-client/src/graph_dsl/traverse.rs @@ -0,0 +1,151 @@ +// SPDX-License-Identifier: Apache-2.0 + +//! `GRAPH TRAVERSE` and `GRAPH PATH` statement builders. Both clients send +//! these statements: the remote client over pgwire, the native client over +//! its SQL query path. + +use nodedb_types::error::NodeDbResult; +use nodedb_types::filter::EdgeFilter; +use nodedb_types::graph::Direction; +use nodedb_types::id::NodeId; + +use super::edge_predicate::edge_where_clause; +use crate::sql_escape::quote_string_literal; + +/// `GRAPH TRAVERSE IN '' FROM '' DEPTH DIRECTION +/// [LABEL '', …] [EDGE WHERE …]`. +pub(crate) fn build_graph_traverse_sql( + collection: &str, + start: &NodeId, + depth: u8, + direction: Direction, + edge_filter: Option<&EdgeFilter>, +) -> NodeDbResult { + Ok(format!( + "GRAPH TRAVERSE IN {} FROM {} DEPTH {depth} DIRECTION {}{}{}", + quote_string_literal(collection), + quote_string_literal(start.as_str()), + direction.as_str(), + label_clause(edge_filter), + filter_clause(edge_filter)?, + )) +} + +/// `GRAPH PATH IN '' FROM '' TO '' MAX_DEPTH +/// [LABEL '', …] [EDGE WHERE …]`. +pub(crate) fn build_graph_path_sql( + collection: &str, + from: &NodeId, + to: &NodeId, + max_depth: u8, + edge_filter: Option<&EdgeFilter>, +) -> NodeDbResult { + Ok(format!( + "GRAPH PATH IN {} FROM {} TO {} MAX_DEPTH {max_depth}{}{}", + quote_string_literal(collection), + quote_string_literal(from.as_str()), + quote_string_literal(to.as_str()), + label_clause(edge_filter), + filter_clause(edge_filter)?, + )) +} + +/// ` LABEL 'a', 'b'` for every label of `edge_filter`. No labels render +/// nothing, which keeps every edge. +fn label_clause(edge_filter: Option<&EdgeFilter>) -> String { + let labels = edge_filter.map_or(&[][..], |filter| filter.labels.as_slice()); + if labels.is_empty() { + return String::new(); + } + let quoted: Vec = labels + .iter() + .map(|label| quote_string_literal(label)) + .collect(); + format!(" LABEL {}", quoted.join(", ")) +} + +/// ` EDGE WHERE …` for the property filters of `edge_filter`. +fn filter_clause(edge_filter: Option<&EdgeFilter>) -> NodeDbResult { + edge_where_clause(edge_filter.map_or(&[][..], |filter| filter.property_filters.as_slice())) +} + +#[cfg(test)] +mod tests { + use super::*; + use nodedb_types::filter::MetadataFilter; + use nodedb_types::value::Value; + + fn node(id: &str) -> NodeId { + NodeId::try_new(id).expect("valid node id") + } + + #[test] + fn traverse_renders_direction_depth_and_escaped_literals() { + for direction in [Direction::Out, Direction::In, Direction::Both] { + let sql = build_graph_traverse_sql("it's", &node("seed'one"), 3, direction, None) + .expect("builds"); + assert_eq!( + sql, + format!( + "GRAPH TRAVERSE IN 'it''s' FROM 'seed''one' DEPTH 3 DIRECTION {}", + direction.as_str() + ) + ); + } + } + + #[test] + fn every_label_renders_escaped() { + let filter = EdgeFilter::labels(["next", "it's"]); + let sql = build_graph_traverse_sql("g", &node("a"), 1, Direction::Out, Some(&filter)) + .expect("builds"); + assert_eq!( + sql, + "GRAPH TRAVERSE IN 'g' FROM 'a' DEPTH 1 DIRECTION out LABEL 'next', 'it''s'" + ); + } + + #[test] + fn labels_precede_the_edge_predicate() { + let filter = EdgeFilter { + labels: vec!["road".into()], + property_filters: vec![ + MetadataFilter::Gt { + field: "score".into(), + value: Value::Integer(5), + }, + MetadataFilter::eq("owner", "O'Reilly"), + ], + }; + let sql = build_graph_traverse_sql("g", &node("a"), 2, Direction::Both, Some(&filter)) + .expect("builds"); + assert_eq!( + sql, + r#"GRAPH TRAVERSE IN 'g' FROM 'a' DEPTH 2 DIRECTION both LABEL 'road' EDGE WHERE ("score" > 5) AND ("owner" = 'O''Reilly')"# + ); + let path = + build_graph_path_sql("g", &node("a"), &node("z"), 6, Some(&filter)).expect("builds"); + assert_eq!( + path, + r#"GRAPH PATH IN 'g' FROM 'a' TO 'z' MAX_DEPTH 6 LABEL 'road' EDGE WHERE ("score" > 5) AND ("owner" = 'O''Reilly')"# + ); + } + + #[test] + fn path_without_a_filter_has_no_clauses() { + let sql = build_graph_path_sql("g", &node("a"), &node("b'c"), 4, Some(&EdgeFilter::all())) + .expect("builds"); + assert_eq!(sql, "GRAPH PATH IN 'g' FROM 'a' TO 'b''c' MAX_DEPTH 4"); + } + + #[test] + fn an_unrenderable_filter_value_is_an_error() { + let filter = EdgeFilter { + labels: Vec::new(), + property_filters: vec![MetadataFilter::eq("b", Value::Bytes(vec![0]))], + }; + assert!( + build_graph_traverse_sql("g", &node("a"), 1, Direction::Out, Some(&filter)).is_err() + ); + } +} diff --git a/nodedb-client/src/lib.rs b/nodedb-client/src/lib.rs index 36c7ca97f..68b6cd610 100644 --- a/nodedb-client/src/lib.rs +++ b/nodedb-client/src/lib.rs @@ -25,12 +25,24 @@ mod row_decode; #[cfg(any(feature = "native", feature = "remote"))] mod sql_escape; -/// Shared graph-DSL builders and result parsers. Used by both clients so the -/// `GRAPH ALGO …` SQL construction and row decoding exist once, not per -/// transport. +/// Shared graph-DSL builders and result decoders. Used by both clients so the +/// `GRAPH ALGO`, `GRAPH TRAVERSE` and `GRAPH PATH` SQL construction and result +/// decoding exist once, not per transport. #[cfg(any(feature = "native", feature = "remote"))] mod graph_dsl; +/// The identity column of a stored schemaless document, shared by both +/// clients so a document reads back with the same fields whichever wrote it. +#[cfg(any(feature = "native", feature = "remote"))] +mod document_identity; + +/// Shared search SQL builders: the text-search statement and the allowed-id +/// key conjunct. Used by both clients, so both send the same statement. +#[cfg(any(feature = "native", feature = "remote"))] +mod search_sql; + +#[cfg(feature = "remote")] +mod pg_cell; #[cfg(feature = "remote")] pub mod remote; #[cfg(feature = "remote")] diff --git a/nodedb-client/src/native/client/dispatch.rs b/nodedb-client/src/native/client/dispatch.rs index e13d5cdd6..a5509d29f 100644 --- a/nodedb-client/src/native/client/dispatch.rs +++ b/nodedb-client/src/native/client/dispatch.rs @@ -15,6 +15,7 @@ use nodedb_types::graph::GraphStats; use nodedb_types::id::{EdgeId, NodeId}; use nodedb_types::protocol::Limits; use nodedb_types::result::{QueryResult, SearchResult, SubGraph}; +use nodedb_types::text_search::TextSearchParams; use nodedb_types::value::Value; use crate::traits::NodeDb; @@ -58,9 +59,10 @@ impl NodeDb for NativeClient { query: &[f32], k: usize, filter: Option<&MetadataFilter>, - _allowed_ids: Option<&std::collections::HashSet>, + allowed_ids: Option<&std::collections::HashSet>, ) -> NodeDbResult> { - self.vector_search_impl(collection, query, k, filter).await + self.vector_search_impl(collection, query, k, filter, allowed_ids) + .await } async fn vector_insert( @@ -83,9 +85,10 @@ impl NodeDb for NativeClient { collection: &str, start: &NodeId, depth: u8, + direction: nodedb_types::graph::Direction, edge_filter: Option<&EdgeFilter>, ) -> NodeDbResult { - self.graph_traverse_impl(collection, start, depth, edge_filter) + self.graph_traverse_impl(collection, start, depth, direction, edge_filter) .await } @@ -129,6 +132,18 @@ impl NodeDb for NativeClient { .await } + async fn graph_shortest_path( + &self, + collection: &str, + from: &NodeId, + to: &NodeId, + max_depth: u8, + edge_filter: Option<&EdgeFilter>, + ) -> NodeDbResult>> { + self.graph_shortest_path_impl(collection, from, to, max_depth, edge_filter) + .await + } + async fn list_insert( &self, collection: &str, @@ -176,6 +191,19 @@ impl NodeDb for NativeClient { self.document_delete_impl(collection, id).await } + async fn text_search( + &self, + collection: &str, + field: &str, + query: &str, + top_k: usize, + params: TextSearchParams, + allowed_ids: Option<&std::collections::HashSet>, + ) -> NodeDbResult> { + self.text_search_impl(collection, field, query, top_k, params, allowed_ids) + .await + } + async fn execute_sql(&self, query: &str, params: &[Value]) -> NodeDbResult { self.execute_sql_impl(query, params).await } diff --git a/nodedb-client/src/native/client/document.rs b/nodedb-client/src/native/client/document.rs index 9ff9ab455..d160ec8bf 100644 --- a/nodedb-client/src/native/client/document.rs +++ b/nodedb-client/src/native/client/document.rs @@ -7,6 +7,7 @@ use nodedb_types::error::{NodeDbError, NodeDbResult}; use nodedb_types::protocol::{NativeResponse, OpCode, TextFields}; use super::core::NativeClient; +use crate::document_identity::is_identity_cell; use crate::native::connection::check_error; impl NativeClient { @@ -30,6 +31,10 @@ impl NativeClient { point_get_response_to_document(collection, id, resp) } + /// Replace the document's whole field set, creating it when absent. + /// + /// The server owns the identity column: it writes `doc.id` under the + /// collection's key and refuses a field there that names another id. pub(super) async fn document_put_impl( &self, collection: &str, @@ -128,7 +133,9 @@ fn point_get_response_to_document( let mut doc = Document::new(id); for (name, value) in columns.into_iter().zip(row) { - doc.set(name, value); + if !is_identity_cell(&name, &value, id) { + doc.set(name, value); + } } Ok(Some(doc)) } @@ -174,6 +181,20 @@ mod tests { assert_eq!(doc.get("active"), Some(&Value::Bool(true))); } + #[test] + fn the_stored_identity_cell_is_the_document_id_not_a_field() { + let resp = hit( + &["id", "name"], + vec![Value::String("u-1".into()), Value::String("alice".into())], + ); + let doc = point_get_response_to_document("users", "u-1", resp) + .expect("a well-formed hit must parse") + .expect("a hit must yield a document"); + assert_eq!(doc.id, "u-1"); + assert_eq!(doc.fields.len(), 1); + assert_eq!(doc.get("name"), Some(&Value::String("alice".into()))); + } + #[test] fn hit_preserves_nested_structure() { let mut inner = std::collections::HashMap::new(); diff --git a/nodedb-client/src/native/client/graph.rs b/nodedb-client/src/native/client/graph.rs index 8b91d6571..17cf5e02e 100644 --- a/nodedb-client/src/native/client/graph.rs +++ b/nodedb-client/src/native/client/graph.rs @@ -10,7 +10,6 @@ use nodedb_types::id::{EdgeId, NodeId}; use nodedb_types::protocol::{OpCode, TextFields}; use nodedb_types::result::SubGraph; -use super::super::response_parse::parse_subgraph_response; use super::core::NativeClient; use crate::native::connection::check_error; use crate::sql_escape::quote_string_literal; @@ -21,25 +20,35 @@ impl NativeClient { collection: &str, start: &NodeId, depth: u8, + direction: nodedb_types::graph::Direction, edge_filter: Option<&EdgeFilter>, ) -> NodeDbResult { - let mut conn = self.pool.acquire().await?; - let resp = conn - .send( - OpCode::GraphHop, - TextFields { - collection: Some(collection.to_string()), - start_node: Some(start.as_str().to_string()), - depth: Some(depth as u32), - edge_label: edge_filter.and_then(|f| f.labels.first().cloned()), - ..Default::default() - }, - ) - .await?; - // An error frame carries no rows; parsed unchecked it would read as - // "the traversal found nothing". - check_error(&resp)?; - parse_subgraph_response(&resp) + // `GRAPH TRAVERSE` answers with nodes, their depths and the crossed + // edges with their properties, the same result the remote client + // decodes. `OpCode::GraphHop` answers only the reachable node set. + let sql = crate::graph_dsl::build_graph_traverse_sql( + collection, + start, + depth, + direction, + edge_filter, + )?; + let result = self.query(&sql).await?; + crate::graph_dsl::decode_traverse_result(&result.columns, &result.rows) + } + + pub(super) async fn graph_shortest_path_impl( + &self, + collection: &str, + from: &NodeId, + to: &NodeId, + max_depth: u8, + edge_filter: Option<&EdgeFilter>, + ) -> NodeDbResult>> { + let sql = + crate::graph_dsl::build_graph_path_sql(collection, from, to, max_depth, edge_filter)?; + let result = self.query(&sql).await?; + crate::graph_dsl::decode_path_result(&result.columns, &result.rows) } pub(super) async fn graph_insert_edge_impl( diff --git a/nodedb-client/src/native/client/identity.rs b/nodedb-client/src/native/client/identity.rs new file mode 100644 index 000000000..d326b7a09 --- /dev/null +++ b/nodedb-client/src/native/client/identity.rs @@ -0,0 +1,45 @@ +// SPDX-License-Identifier: Apache-2.0 + +//! The identity column of a collection, read from the server catalog. +//! +//! A search restricted to allowed ids names the key column in its SQL: the +//! declared primary key, else `id`. The client reads it from `DESCRIBE` on +//! every call that needs it, so a dropped and recreated collection never +//! answers with a stale key. + +use nodedb_types::error::NodeDbResult; +use nodedb_types::value::Value; + +use crate::document_identity::identity_column_from_describe; +use crate::sql_escape::quote_identifier; + +use super::core::NativeClient; + +impl NativeClient { + /// The identity column of `collection`. + pub(super) async fn identity_column(&self, collection: &str) -> NodeDbResult { + let sql = format!("DESCRIBE {}", quote_identifier(collection)); + let result = self.query(&sql).await?; + identity_column_from_describe(collection, &result.columns, &result.rows, typed_bool) + } +} + +/// A `bool` cell as the native protocol carries it: a typed `Bool`. +fn typed_bool(cell: &Value) -> Option { + match cell { + Value::Bool(b) => Some(*b), + _ => None, + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn only_typed_bools_decode() { + assert_eq!(typed_bool(&Value::Bool(true)), Some(true)); + assert_eq!(typed_bool(&Value::Bool(false)), Some(false)); + assert_eq!(typed_bool(&Value::String("t".into())), None); + } +} diff --git a/nodedb-client/src/native/client/mod.rs b/nodedb-client/src/native/client/mod.rs index eb6d6d9e3..e232e9dfa 100644 --- a/nodedb-client/src/native/client/mod.rs +++ b/nodedb-client/src/native/client/mod.rs @@ -5,7 +5,9 @@ mod crdt_list; mod dispatch; mod document; mod graph; +mod identity; mod sql_lifecycle; +mod text_search; mod vector; pub use core::NativeClient; diff --git a/nodedb-client/src/native/client/text_search.rs b/nodedb-client/src/native/client/text_search.rs new file mode 100644 index 000000000..9a50abda6 --- /dev/null +++ b/nodedb-client/src/native/client/text_search.rs @@ -0,0 +1,48 @@ +// SPDX-License-Identifier: Apache-2.0 + +//! Native full-text search. +//! +//! The statement comes from the shared builder in `search_sql`, the one the +//! pgwire client sends, and the hits decode through the same row decoder. +//! Both clients return the same ids and scores for one search. + +use std::collections::HashSet; + +use nodedb_types::error::NodeDbResult; +use nodedb_types::result::SearchResult; +use nodedb_types::text_search::TextSearchParams; + +use crate::row_decode::decode_search_hits; +use crate::search_sql::{TextSearchRequest, text_hit_source, text_search_sql}; + +use super::core::NativeClient; + +impl NativeClient { + pub(super) async fn text_search_impl( + &self, + collection: &str, + field: &str, + query: &str, + top_k: usize, + params: TextSearchParams, + allowed_ids: Option<&HashSet>, + ) -> NodeDbResult> { + if allowed_ids.is_some_and(HashSet::is_empty) || top_k == 0 { + return Ok(Vec::new()); + } + let key = match allowed_ids { + Some(_) => Some(self.identity_column(collection).await?), + None => None, + }; + let sql = text_search_sql(&TextSearchRequest { + collection, + field, + query, + top_k, + params: ¶ms, + allowed: key.as_deref().zip(allowed_ids), + }); + let result = self.query(&sql).await?; + decode_search_hits(&text_hit_source(collection), &result.columns, &result.rows) + } +} diff --git a/nodedb-client/src/native/client/vector.rs b/nodedb-client/src/native/client/vector.rs index 5dbbc315c..4760ea268 100644 --- a/nodedb-client/src/native/client/vector.rs +++ b/nodedb-client/src/native/client/vector.rs @@ -2,15 +2,18 @@ //! Vector operation implementations for `NativeClient`. +use std::collections::HashSet; + use nodedb_types::document::Document; use nodedb_types::error::{NodeDbError, NodeDbResult}; use nodedb_types::filter::MetadataFilter; use nodedb_types::protocol::{OpCode, TextFields}; use nodedb_types::result::SearchResult; -use super::super::response_parse::parse_search_results; use super::core::NativeClient; use crate::native::connection::check_error; +use crate::row_decode::search_hit::DISTANCE_COLUMN; +use crate::row_decode::{HitSource, decode_search_hits}; use crate::sql_escape::{quote_identifier, quote_string_literal}; impl NativeClient { @@ -20,14 +23,27 @@ impl NativeClient { query: &[f32], k: usize, filter: Option<&MetadataFilter>, + allowed_ids: Option<&HashSet>, ) -> NodeDbResult> { - let request = build_vector_search_request(collection, query, k, filter)?; + // An empty allowed set admits no candidate. + if allowed_ids.is_some_and(HashSet::is_empty) { + return Ok(Vec::new()); + } + let request = build_vector_search_request(collection, query, k, filter, allowed_ids)?; let mut conn = self.pool.acquire().await?; let resp = conn.send(OpCode::VectorSearch, request).await?; // An error frame carries no rows; parsed unchecked it would read as // "the search matched nothing". check_error(&resp)?; - parse_search_results(&resp) + decode_search_hits( + &HitSource { + op: "vector_search", + collection, + score_column: DISTANCE_COLUMN, + }, + resp.columns.as_deref().unwrap_or_default(), + resp.rows.as_deref().unwrap_or_default(), + ) } pub(super) async fn vector_insert_impl( @@ -77,39 +93,41 @@ impl NativeClient { /// Build the `TextFields` payload for an `OpCode::VectorSearch` request. /// -/// The native protocol reserves wire byte 68 for the optional -/// `TextFields::filters: Option>` field. When the trait caller -/// passes a non-`None` `MetadataFilter`, the predicate is serialized -/// here so it travels alongside the SQL/DSL request rather than being -/// dropped at the client. +/// A non-`None` `MetadataFilter` travels as MessagePack in +/// `TextFields::filters`. The server plans it as the `WHERE` a pgwire +/// client renders from the same filter, so it narrows the candidates +/// before the top-k cut. /// -/// Wire-format note: the inline doc on `TextFields::filters` calls for -/// MessagePack. Until the server-side decoder is wired (the dispatch -/// path currently constructs plans with `filters: Vec::new()`), the -/// client serializes via sonic_rs JSON. The server-side fix will switch -/// both sides to a single agreed encoding; for now the bytes are -/// observable as non-empty, which is what the trait contract requires. +/// `allowed_ids` travel sorted in `TextFields::allowed_ids`. The server +/// lowers them to the candidate bitmap the index search honors. pub(super) fn build_vector_search_request( collection: &str, query: &[f32], k: usize, filter: Option<&MetadataFilter>, + allowed_ids: Option<&HashSet>, ) -> NodeDbResult { // Serialization failure here must surface to the caller. Dropping // the filter and sending the request anyway would send the query // to the server without the caller's predicate — exactly the // silent-drop pattern this client guards against. let filters_bytes = match filter { - Some(f) => Some(sonic_rs::to_vec(f).map_err(|e| { - NodeDbError::serialization("json", format!("vector_search metadata filter: {e}")) + Some(f) => Some(zerompk::to_msgpack_vec(f).map_err(|e| { + NodeDbError::serialization("msgpack", format!("vector_search metadata filter: {e}")) })?), None => None, }; + let allowed_ids = allowed_ids.map(|ids| { + let mut ids: Vec = ids.iter().cloned().collect(); + ids.sort(); + ids + }); Ok(TextFields { collection: Some(collection.to_string()), query_vector: Some(query.to_vec()), top_k: Some(k as u32), filters: filters_bytes, + allowed_ids, ..Default::default() }) } @@ -132,8 +150,9 @@ mod tests { #[test] fn vector_search_request_without_filter_omits_filter_bytes() { - let req = build_vector_search_request("docs", &[0.1, 0.2], 5, None) + let req = build_vector_search_request("docs", &[0.1, 0.2], 5, None, None) .expect("no-filter request must build"); + assert!(req.allowed_ids.is_none(), "no restriction sends no ids"); assert_eq!(req.collection.as_deref(), Some("docs")); assert_eq!(req.query_vector.as_deref(), Some(&[0.1f32, 0.2][..])); assert_eq!(req.top_k, Some(5)); @@ -146,12 +165,22 @@ mod tests { #[test] fn vector_search_request_serializes_metadata_filter() { let filter = MetadataFilter::eq("category", Value::String("ai".into())); - let req = build_vector_search_request("docs", &[0.1], 3, Some(&filter)) + let req = build_vector_search_request("docs", &[0.1], 3, Some(&filter), None) .expect("derived-Serialize MetadataFilter must encode"); let bytes = req.filters.expect("non-None filter must produce bytes"); - assert!( - !bytes.is_empty(), - "serialized filter bytes must not be empty" + let decoded: MetadataFilter = + zerompk::from_msgpack(&bytes).expect("filter bytes are a MessagePack MetadataFilter"); + assert_eq!(decoded, filter); + } + + #[test] + fn vector_search_request_carries_allowed_ids_sorted() { + let ids: HashSet = ["b", "a"].iter().map(|s| s.to_string()).collect(); + let req = build_vector_search_request("docs", &[0.1], 3, None, Some(&ids)) + .expect("allowed ids must encode"); + assert_eq!( + req.allowed_ids, + Some(vec!["a".to_string(), "b".to_string()]) ); } diff --git a/nodedb-client/src/native/mod.rs b/nodedb-client/src/native/mod.rs index d01bb4bf4..74d4b243c 100644 --- a/nodedb-client/src/native/mod.rs +++ b/nodedb-client/src/native/mod.rs @@ -4,6 +4,5 @@ pub mod builder; pub mod client; pub mod connection; pub mod pool; -pub(crate) mod response_parse; pub use client::NativeClient; diff --git a/nodedb-client/src/native/response_parse.rs b/nodedb-client/src/native/response_parse.rs deleted file mode 100644 index b2bc894bc..000000000 --- a/nodedb-client/src/native/response_parse.rs +++ /dev/null @@ -1,140 +0,0 @@ -// SPDX-License-Identifier: Apache-2.0 - -//! Response parsing helpers for native protocol results. - -use std::collections::HashMap; - -use nodedb_types::error::NodeDbResult; -use nodedb_types::id::{EdgeId, NodeId}; -use nodedb_types::result::{SearchResult, SubGraph, SubGraphEdge, SubGraphNode}; - -/// Parse search results from a native response. -pub(crate) fn parse_search_results( - resp: &nodedb_types::protocol::NativeResponse, -) -> NodeDbResult> { - let rows = match &resp.rows { - Some(r) => r, - None => return Ok(Vec::new()), - }; - - let mut results = Vec::new(); - for row in rows { - if let Some(text) = row.first().and_then(|v| v.as_str()) { - if let Ok(items) = sonic_rs::from_str::>(text) { - for item in items { - if let Some(sr) = parse_single_search_result(&item) { - results.push(sr); - } - } - } else if let Ok(item) = sonic_rs::from_str::(text) - && let Some(sr) = parse_single_search_result(&item) - { - results.push(sr); - } - } - } - Ok(results) -} - -fn parse_single_search_result(v: &serde_json::Value) -> Option { - let id = v.get("id")?.as_str()?.to_string(); - let distance = v.get("distance")?.as_f64()? as f32; - Some(SearchResult { - id, - node_id: None, - distance, - metadata: HashMap::new(), - }) -} - -/// Parse a graph traversal response into a SubGraph. -pub(crate) fn parse_subgraph_response( - resp: &nodedb_types::protocol::NativeResponse, -) -> NodeDbResult { - let rows = match &resp.rows { - Some(r) => r, - None => return Ok(SubGraph::empty()), - }; - - let mut nodes = Vec::new(); - let mut edges = Vec::new(); - - for row in rows { - let text = match row.first().and_then(|v| v.as_str()) { - Some(t) => t, - None => continue, - }; - - if let Ok(val) = sonic_rs::from_str::(text) { - if let Some(obj) = val.as_object() { - if let Some(ns) = obj.get("nodes").and_then(|v| v.as_array()) { - for n in ns { - if let Some(id) = n.get("id").and_then(|v| v.as_str()) { - let depth = n.get("depth").and_then(|v| v.as_u64()).unwrap_or(0) as u8; - nodes.push(SubGraphNode { - id: NodeId::from_validated(id.to_owned()), - depth, - properties: HashMap::new(), - }); - } - } - } - if let Some(es) = obj.get("edges").and_then(|v| v.as_array()) { - for e in es { - let from = e - .get("from") - .or_else(|| e.get("src")) - .and_then(|v| v.as_str()) - .unwrap_or(""); - let to = e - .get("to") - .or_else(|| e.get("dst")) - .and_then(|v| v.as_str()) - .unwrap_or(""); - let label = e.get("label").and_then(|v| v.as_str()).unwrap_or(""); - edges.push(SubGraphEdge { - id: EdgeId::try_first( - NodeId::from_validated(from.to_owned()), - NodeId::from_validated(to.to_owned()), - label, - ) - .expect("server wire label already validated"), - from: NodeId::from_validated(from.to_owned()), - to: NodeId::from_validated(to.to_owned()), - label: label.to_string(), - properties: HashMap::new(), - }); - } - } - } - if let Some(arr) = val.as_array() { - for item in arr { - if let Some(id) = item.get("id").and_then(|v| v.as_str()) { - let depth = item.get("depth").and_then(|v| v.as_u64()).unwrap_or(0) as u8; - nodes.push(SubGraphNode { - id: NodeId::from_validated(id.to_owned()), - depth, - properties: HashMap::new(), - }); - } - } - } - } - } - - Ok(SubGraph { nodes, edges }) -} - -#[cfg(test)] -mod tests { - use super::*; - - #[test] - fn parse_search_result_from_json() { - let v = serde_json::json!({"id": "vec-1", "distance": 0.123}); - let sr = - parse_single_search_result(&v).expect("failed to parse search result from valid JSON"); - assert_eq!(sr.id, "vec-1"); - assert!((sr.distance - 0.123).abs() < 0.001); - } -} diff --git a/nodedb-client/src/pg_cell/array.rs b/nodedb-client/src/pg_cell/array.rs new file mode 100644 index 000000000..0ebc1a128 --- /dev/null +++ b/nodedb-client/src/pg_cell/array.rs @@ -0,0 +1,238 @@ +// SPDX-License-Identifier: Apache-2.0 + +//! Decoder for the PostgreSQL text form of an array: `{1.5,NaN,NULL}`. +//! +//! The server sends array columns in text format. Each dimension is a +//! `{...}` list of comma-separated elements. An element is a nested list, a +//! double-quoted string with backslash escapes, the unquoted word `NULL`, or +//! unquoted text. Every element decodes through the text decoder of the +//! element type, so a `float4[]` element reads `NaN`, `Infinity` and +//! `-Infinity` the way a `float4` cell does. + +use nodedb_types::Value; +use tokio_postgres::types::Type; + +use super::error::DecodeReason; +use super::text; + +/// Decode `text` as an array literal whose elements have PG type `element`. +/// A multi-dimensional array decodes as nested `Value::Array`s. +pub(super) fn array(text: &str, element: &Type) -> Result { + let mut parser = Parser { + bytes: text.as_bytes(), + pos: 0, + index: 0, + element, + }; + let value = parser.list()?; + parser.skip_whitespace(); + if parser.pos != parser.bytes.len() { + return Err(DecodeReason::ArraySyntax("text after the closing brace")); + } + Ok(value) +} + +struct Parser<'a> { + bytes: &'a [u8], + pos: usize, + /// Position of the next element in flattened order, for errors. + index: usize, + element: &'a Type, +} + +impl Parser<'_> { + fn peek(&self) -> Option { + self.bytes.get(self.pos).copied() + } + + fn skip_whitespace(&mut self) { + while self.peek().is_some_and(|b| b.is_ascii_whitespace()) { + self.pos += 1; + } + } + + /// One `{...}` list, nested lists included. + fn list(&mut self) -> Result { + self.skip_whitespace(); + if self.peek() != Some(b'{') { + return Err(DecodeReason::ArraySyntax("a list must start with '{'")); + } + self.pos += 1; + let mut items = Vec::new(); + self.skip_whitespace(); + if self.peek() == Some(b'}') { + self.pos += 1; + return Ok(Value::Array(items)); + } + loop { + self.skip_whitespace(); + let item = match self.peek() { + Some(b'{') => self.list()?, + Some(b'"') => { + let quoted = self.quoted()?; + self.decode(quoted)? + } + _ => self.unquoted()?, + }; + items.push(item); + self.skip_whitespace(); + match self.peek() { + Some(b',') => self.pos += 1, + Some(b'}') => { + self.pos += 1; + return Ok(Value::Array(items)); + } + Some(_) => return Err(DecodeReason::ArraySyntax("expected ',' or '}'")), + None => return Err(DecodeReason::ArraySyntax("missing closing brace")), + } + } + } + + /// A double-quoted element. A backslash takes the next byte literally. + fn quoted(&mut self) -> Result { + // Skip the opening quote. + self.pos += 1; + let mut out = Vec::new(); + loop { + match self.peek() { + None => return Err(DecodeReason::ArraySyntax("unterminated quoted element")), + Some(b'"') => { + self.pos += 1; + break; + } + Some(b'\\') => { + let escaped = self + .bytes + .get(self.pos + 1) + .copied() + .ok_or(DecodeReason::ArraySyntax("dangling backslash"))?; + out.push(escaped); + self.pos += 2; + } + Some(byte) => { + out.push(byte); + self.pos += 1; + } + } + } + // Only ASCII quotes and backslashes are removed from UTF-8 input, so + // the rest stays UTF-8. + String::from_utf8(out).map_err(|_| DecodeReason::NotUtf8) + } + + /// An unquoted element: the text up to the next `,` or `}`, with + /// surrounding whitespace removed. The word `NULL` in any case is SQL + /// NULL. + fn unquoted(&mut self) -> Result { + let start = self.pos; + while self + .peek() + .is_some_and(|b| !matches!(b, b',' | b'}' | b'{' | b'"')) + { + self.pos += 1; + } + let raw = self.bytes.get(start..self.pos).unwrap_or_default(); + let word = std::str::from_utf8(raw) + .map_err(|_| DecodeReason::NotUtf8)? + .trim(); + if word.is_empty() { + return Err(DecodeReason::ArraySyntax("empty element")); + } + if word.eq_ignore_ascii_case("NULL") { + self.index += 1; + return Ok(Value::Null); + } + self.decode(word.to_owned()) + } + + /// Decode one element's text as the element type. + fn decode(&mut self, element_text: String) -> Result { + let index = self.index; + self.index += 1; + text::scalar(self.element, &element_text).map_err(|source| DecodeReason::Element { + index, + text: element_text, + element_type: self.element.name().to_owned(), + source: Box::new(source), + }) + } +} + +#[cfg(test)] +mod tests { + use super::*; + + fn floats(values: &[f64]) -> Value { + Value::Array(values.iter().copied().map(Value::Float).collect()) + } + + #[test] + fn decodes_float_arrays_including_non_finite() { + assert_eq!( + array("{1.5,-2,Infinity,-Infinity}", &Type::FLOAT8), + Ok(floats(&[1.5, -2.0, f64::INFINITY, f64::NEG_INFINITY])) + ); + let Ok(Value::Array(items)) = array("{NaN, 0.1}", &Type::FLOAT4) else { + panic!("a float4 array decodes as an Array"); + }; + assert!(matches!(items.first(), Some(Value::Float(f)) if f.is_nan())); + assert_eq!(items.get(1), Some(&Value::Float(f64::from(0.1f32)))); + assert_eq!(array("{}", &Type::FLOAT8), Ok(Value::Array(Vec::new()))); + } + + #[test] + fn decodes_nulls_quotes_and_nesting() { + assert_eq!( + array(r#"{"a,b",NULL,"say \"hi\"",null}"#, &Type::TEXT), + Ok(Value::Array(vec![ + Value::String("a,b".into()), + Value::Null, + Value::String("say \"hi\"".into()), + Value::Null, + ])) + ); + assert_eq!( + array(r#"{"NULL"}"#, &Type::TEXT), + Ok(Value::Array(vec![Value::String("NULL".into())])), + "a quoted NULL is the string" + ); + assert_eq!( + array("{{1,2},{3,4}}", &Type::INT4), + Ok(Value::Array(vec![ + Value::Array(vec![Value::Integer(1), Value::Integer(2)]), + Value::Array(vec![Value::Integer(3), Value::Integer(4)]), + ])) + ); + } + + #[test] + fn names_the_element_that_does_not_decode() { + assert_eq!( + array("{1.5,abc}", &Type::FLOAT8), + Err(DecodeReason::Element { + index: 1, + text: "abc".into(), + element_type: "float8".into(), + source: Box::new(DecodeReason::Invalid { expected: "float8" }), + }) + ); + } + + #[test] + fn refuses_malformed_literals() { + for text in [ + "[1.5,2]", + "1.5,2", + "{1.5,2", + "{1.5,,2}", + "{1.5} x", + "{\"open}", + "{1.5 2.5}", + ] { + assert!( + array(text, &Type::FLOAT8).is_err(), + "{text:?} must be refused" + ); + } + } +} diff --git a/nodedb-client/src/pg_cell/binary.rs b/nodedb-client/src/pg_cell/binary.rs new file mode 100644 index 000000000..749d37118 --- /dev/null +++ b/nodedb-client/src/pg_cell/binary.rs @@ -0,0 +1,173 @@ +// SPDX-License-Identifier: Apache-2.0 + +//! Binary-format decoders for the scalar types the server sends in binary. +//! +//! `tokio_postgres` requests the binary result format for every column. The +//! server honours it for `bool`, `int2`, `int4`, `int8`, `float4`, `float8`, +//! `timestamp` and `timestamptz`, so these decoders read the PostgreSQL +//! binary wire form of exactly those types. + +use nodedb_types::{NdbDateTime, Value}; + +use super::error::DecodeReason; + +/// Microseconds from the Unix epoch to the PostgreSQL epoch +/// (2000-01-01 00:00:00 UTC), which binary timestamps count from. +const PG_EPOCH_OFFSET_MICROS: i64 = 946_684_800_000_000; + +/// The `N` bytes of a fixed-width binary scalar. +fn fixed(raw: &[u8]) -> Result<[u8; N], DecodeReason> { + raw.try_into().map_err(|_| DecodeReason::Width { + expected: N, + found: raw.len(), + }) +} + +/// Binary `bool`: one byte, 0 or 1. +pub(super) fn boolean(raw: &[u8]) -> Result { + match fixed::<1>(raw)? { + [0] => Ok(Value::Bool(false)), + [1] => Ok(Value::Bool(true)), + _ => Err(DecodeReason::BoolByte), + } +} + +/// Binary `int2`: a big-endian `i16`. +pub(super) fn int2(raw: &[u8]) -> Result { + Ok(Value::Integer(i64::from(i16::from_be_bytes(fixed(raw)?)))) +} + +/// Binary `int4`: a big-endian `i32`. +pub(super) fn int4(raw: &[u8]) -> Result { + Ok(Value::Integer(i64::from(i32::from_be_bytes(fixed(raw)?)))) +} + +/// Binary `int8`: a big-endian `i64`. +pub(super) fn int8(raw: &[u8]) -> Result { + Ok(Value::Integer(i64::from_be_bytes(fixed(raw)?))) +} + +/// Binary `float4`: a big-endian IEEE-754 `f32`, widened exactly to `f64`. +/// `NaN` and the infinities keep their value. +pub(super) fn float4(raw: &[u8]) -> Result { + Ok(Value::Float(f64::from(f32::from_be_bytes(fixed(raw)?)))) +} + +/// Binary `float8`: a big-endian IEEE-754 `f64`. +pub(super) fn float8(raw: &[u8]) -> Result { + Ok(Value::Float(f64::from_be_bytes(fixed(raw)?))) +} + +/// The instant a binary `timestamp`/`timestamptz` holds: an `i64` of +/// microseconds since the PostgreSQL epoch. +fn instant(raw: &[u8]) -> Result { + let pg_micros = i64::from_be_bytes(fixed(raw)?); + pg_micros + .checked_add(PG_EPOCH_OFFSET_MICROS) + .map(NdbDateTime::from_micros) + .ok_or(DecodeReason::OutOfRange { what: "timestamp" }) +} + +/// Binary `timestamp`: a naive instant. +pub(super) fn timestamp(raw: &[u8]) -> Result { + instant(raw).map(Value::NaiveDateTime) +} + +/// Binary `timestamptz`: a UTC instant. +pub(super) fn timestamptz(raw: &[u8]) -> Result { + instant(raw).map(Value::DateTime) +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn decodes_bool() { + assert_eq!(boolean(&[1]), Ok(Value::Bool(true))); + assert_eq!(boolean(&[0]), Ok(Value::Bool(false))); + assert_eq!(boolean(&[2]), Err(DecodeReason::BoolByte)); + assert_eq!( + boolean(b"t"), + Err(DecodeReason::BoolByte), + "text `t` under a binary bool is refused" + ); + assert_eq!( + boolean(&[]), + Err(DecodeReason::Width { + expected: 1, + found: 0 + }) + ); + } + + #[test] + fn decodes_integers() { + assert_eq!(int2(&(-7i16).to_be_bytes()), Ok(Value::Integer(-7))); + assert_eq!( + int4(&i32::MAX.to_be_bytes()), + Ok(Value::Integer(i64::from(i32::MAX))) + ); + assert_eq!(int8(&i64::MIN.to_be_bytes()), Ok(Value::Integer(i64::MIN))); + assert_eq!( + int4(b"42"), + Err(DecodeReason::Width { + expected: 4, + found: 2 + }) + ); + assert_eq!( + int8(&[0; 4]), + Err(DecodeReason::Width { + expected: 8, + found: 4 + }) + ); + } + + #[test] + fn decodes_floats_including_non_finite() { + assert_eq!(float4(&1.5f32.to_be_bytes()), Ok(Value::Float(1.5))); + assert_eq!(float8(&(-2.25f64).to_be_bytes()), Ok(Value::Float(-2.25))); + assert_eq!( + float8(&f64::INFINITY.to_be_bytes()), + Ok(Value::Float(f64::INFINITY)) + ); + assert_eq!( + float4(&f32::NEG_INFINITY.to_be_bytes()), + Ok(Value::Float(f64::NEG_INFINITY)) + ); + let Ok(Value::Float(nan)) = float8(&f64::NAN.to_be_bytes()) else { + panic!("binary NaN decodes as a float"); + }; + assert!(nan.is_nan()); + assert_eq!( + float4(&[0; 8]), + Err(DecodeReason::Width { + expected: 4, + found: 8 + }) + ); + } + + #[test] + fn decodes_timestamps_from_the_postgres_epoch() { + // 2020-03-05T10:00:00Z. + let unix_micros = 1_583_402_400_000_000i64; + let raw = (unix_micros - PG_EPOCH_OFFSET_MICROS).to_be_bytes(); + let at = NdbDateTime::from_micros(unix_micros); + assert_eq!(timestamp(&raw), Ok(Value::NaiveDateTime(at))); + assert_eq!(timestamptz(&raw), Ok(Value::DateTime(at))); + assert_eq!( + timestamptz(&i64::MAX.to_be_bytes()), + Err(DecodeReason::OutOfRange { what: "timestamp" }) + ); + assert_eq!( + timestamp(b"2020"), + Err(DecodeReason::Width { + expected: 8, + found: 4 + }) + ); + } +} diff --git a/nodedb-client/src/pg_cell/decode.rs b/nodedb-client/src/pg_cell/decode.rs new file mode 100644 index 000000000..483d76732 --- /dev/null +++ b/nodedb-client/src/pg_cell/decode.rs @@ -0,0 +1,146 @@ +// SPDX-License-Identifier: Apache-2.0 + +//! Pick the decoder for one result cell from its PG column type. + +use nodedb_types::Value; +use nodedb_types::error::NodeDbResult; +use tokio_postgres::types::{Kind, Type}; + +use super::error::{DecodeReason, cell_error}; +use super::{array, binary, text}; + +/// Decode one result cell of column `column` with PG type `ty`. `raw` is +/// `None` for SQL NULL. +/// +/// `tokio_postgres` requests the binary result format for every column. The +/// server sends `bool`, `int2`, `int4`, `int8`, `float4`, `float8`, +/// `timestamp` and `timestamptz` in binary. It sends every other type as its +/// PostgreSQL text form. A cell that does not decode is an error naming the +/// column, the PG type and the cell text, never a NULL. +pub(crate) fn decode_cell(column: &str, ty: &Type, raw: Option<&[u8]>) -> NodeDbResult { + let Some(raw) = raw else { + return Ok(Value::Null); + }; + decode_bytes(ty, raw).map_err(|reason| cell_error(column, ty, raw, &reason)) +} + +fn decode_bytes(ty: &Type, raw: &[u8]) -> Result { + match *ty { + Type::BOOL => return binary::boolean(raw), + Type::INT2 => return binary::int2(raw), + Type::INT4 => return binary::int4(raw), + Type::INT8 => return binary::int8(raw), + Type::FLOAT4 => return binary::float4(raw), + Type::FLOAT8 => return binary::float8(raw), + Type::TIMESTAMP => return binary::timestamp(raw), + Type::TIMESTAMPTZ => return binary::timestamptz(raw), + _ => {} + } + let cell_text = std::str::from_utf8(raw).map_err(|_| DecodeReason::NotUtf8)?; + match ty.kind() { + Kind::Array(element) => array::array(cell_text, element), + _ => text::scalar(ty, cell_text), + } +} + +#[cfg(test)] +mod tests { + use nodedb_types::NdbDateTime; + use rust_decimal::Decimal; + + use super::*; + + fn decoded(ty: &Type, raw: &[u8]) -> Value { + decode_cell("c", ty, Some(raw)).unwrap_or_else(|e| panic!("{ty}: {e}")) + } + + fn error_message(ty: &Type, raw: &[u8]) -> String { + match decode_cell("price", ty, Some(raw)) { + Ok(v) => panic!("{ty} must refuse {raw:?}, got {v:?}"), + Err(e) => e.to_string(), + } + } + + #[test] + fn null_is_null_for_every_type() { + for ty in [Type::INT8, Type::NUMERIC, Type::BYTEA, Type::FLOAT4_ARRAY] { + assert_eq!(decode_cell("c", &ty, None).ok(), Some(Value::Null)); + } + } + + #[test] + fn binary_scalars_decode_from_binary() { + assert_eq!(decoded(&Type::BOOL, &[1]), Value::Bool(true)); + assert_eq!(decoded(&Type::INT2, &5i16.to_be_bytes()), Value::Integer(5)); + assert_eq!(decoded(&Type::INT4, &5i32.to_be_bytes()), Value::Integer(5)); + assert_eq!(decoded(&Type::INT8, &5i64.to_be_bytes()), Value::Integer(5)); + assert_eq!( + decoded(&Type::FLOAT4, &0.5f32.to_be_bytes()), + Value::Float(0.5) + ); + assert!( + matches!(decoded(&Type::FLOAT8, &f64::NAN.to_be_bytes()), Value::Float(f) if f.is_nan()) + ); + assert_eq!( + decoded(&Type::TIMESTAMPTZ, &0i64.to_be_bytes()), + Value::DateTime(NdbDateTime::from_micros(946_684_800_000_000)) + ); + assert_eq!( + decoded(&Type::TIMESTAMP, &0i64.to_be_bytes()), + Value::NaiveDateTime(NdbDateTime::from_micros(946_684_800_000_000)) + ); + } + + #[test] + fn other_types_decode_from_text() { + assert_eq!(decoded(&Type::TEXT, b"abc"), Value::String("abc".into())); + assert_eq!(decoded(&Type::VARCHAR, b"abc"), Value::String("abc".into())); + assert_eq!( + decoded(&Type::NUMERIC, b"1.25"), + Value::Decimal(Decimal::new(125, 2)) + ); + assert_eq!(decoded(&Type::BYTEA, b"\\x0102"), Value::Bytes(vec![1, 2])); + assert_eq!( + decoded(&Type::UUID, b"550e8400-e29b-41d4-a716-446655440000"), + Value::Uuid("550e8400-e29b-41d4-a716-446655440000".into()) + ); + assert_eq!( + decoded(&Type::JSONB, b"[true]"), + Value::Array(vec![Value::Bool(true)]) + ); + assert_eq!( + decoded(&Type::FLOAT8_ARRAY, b"{1,-Infinity}"), + Value::Array(vec![Value::Float(1.0), Value::Float(f64::NEG_INFINITY)]) + ); + } + + #[test] + fn errors_name_column_type_and_text() { + let message = error_message(&Type::NUMERIC, b"12,5"); + assert!(message.contains("\"price\""), "{message}"); + assert!(message.contains("numeric"), "{message}"); + assert!(message.contains("\"12,5\""), "{message}"); + + let message = error_message(&Type::INT8, b"42"); + assert!(message.contains("int8"), "{message}"); + assert!(message.contains("needs 8 bytes, found 2"), "{message}"); + + let message = error_message(&Type::BYTEA, b"AQI"); + assert!(message.contains("bytea"), "{message}"); + assert!(message.contains("\"AQI\""), "{message}"); + + let message = error_message(&Type::TEXT, &[0xff, 0xfe]); + assert!(message.contains("\\xfffe"), "{message}"); + assert!(message.contains("not UTF-8"), "{message}"); + + let message = error_message(&Type::FLOAT4_ARRAY, b"[0.5,1]"); + assert!(message.contains("_float4"), "{message}"); + assert!(message.contains("array literal syntax"), "{message}"); + + let message = error_message(&Type::FLOAT8_ARRAY, b"{1,x}"); + assert!(message.contains("array element 1"), "{message}"); + + let message = error_message(&Type::POINT, b"(1,2)"); + assert!(message.contains("no decoder"), "{message}"); + } +} diff --git a/nodedb-client/src/pg_cell/error.rs b/nodedb-client/src/pg_cell/error.rs new file mode 100644 index 000000000..927084347 --- /dev/null +++ b/nodedb-client/src/pg_cell/error.rs @@ -0,0 +1,103 @@ +// SPDX-License-Identifier: Apache-2.0 + +//! The error for a cell the client cannot decode. + +use nodedb_types::error::NodeDbError; +use tokio_postgres::types::Type; + +/// Why the bytes of one cell do not decode. +#[derive(Debug, Clone, thiserror::Error, PartialEq)] +pub(super) enum DecodeReason { + /// A binary scalar with the wrong byte count. + #[error("binary form needs {expected} bytes, found {found}")] + Width { expected: usize, found: usize }, + /// A binary `bool` byte other than 0 or 1. + #[error("binary bool byte must be 0 or 1")] + BoolByte, + /// Text bytes that are not UTF-8. + #[error("text is not UTF-8")] + NotUtf8, + /// Text that is not the PostgreSQL text form of `expected`. + #[error("text is not a valid {expected}")] + Invalid { expected: &'static str }, + /// A value outside the range `Value` can hold. + #[error("{what} is out of range")] + OutOfRange { what: &'static str }, + /// A PG type the server never sends and the client cannot decode. + #[error("the client has no decoder for this type")] + Unsupported, + /// An array literal that breaks the PostgreSQL array syntax. + #[error("array literal syntax: {0}")] + ArraySyntax(&'static str), + /// One array element that does not decode as the element type. + #[error("array element {index} ({text:?}) of type {element_type}: {source}")] + Element { + index: usize, + text: String, + element_type: String, + source: Box, + }, +} + +/// The error for a cell of column `column` and PG type `ty` whose bytes +/// `raw` do not decode. It names the column, the PG type, the cell text and +/// the reason. Bytes that are not printable UTF-8 show as `\x` hex. +pub(super) fn cell_error( + column: &str, + ty: &Type, + raw: &[u8], + reason: &DecodeReason, +) -> NodeDbError { + NodeDbError::serialization( + "pgwire", + format!( + "column \"{column}\" of type {}: cannot decode {}: {reason}", + ty.name(), + cell_text(raw) + ), + ) +} + +/// The cell as the error shows it: quoted UTF-8 text, or `\x` hex for bytes +/// that are not printable UTF-8. +fn cell_text(raw: &[u8]) -> String { + match std::str::from_utf8(raw) { + Ok(text) if !text.chars().any(char::is_control) => format!("{text:?}"), + _ => { + let mut hex = String::with_capacity(2 + raw.len() * 2); + hex.push_str("\\x"); + for byte in raw { + hex.push_str(&format!("{byte:02x}")); + } + hex + } + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn names_column_type_text_and_reason() { + let reason = DecodeReason::Invalid { + expected: "numeric", + }; + let message = cell_error("price", &Type::NUMERIC, b"abc", &reason).to_string(); + assert!(message.contains("\"price\""), "{message}"); + assert!(message.contains("type numeric"), "{message}"); + assert!(message.contains("\"abc\""), "{message}"); + assert!(message.contains("not a valid numeric"), "{message}"); + } + + #[test] + fn shows_binary_bytes_as_hex() { + let reason = DecodeReason::Width { + expected: 4, + found: 3, + }; + let message = cell_error("n", &Type::INT4, &[0, 1, 0xff], &reason).to_string(); + assert!(message.contains("\\x0001ff"), "{message}"); + assert!(message.contains("needs 4 bytes, found 3"), "{message}"); + } +} diff --git a/nodedb-client/src/pg_cell/mod.rs b/nodedb-client/src/pg_cell/mod.rs new file mode 100644 index 000000000..5bfdd3706 --- /dev/null +++ b/nodedb-client/src/pg_cell/mod.rs @@ -0,0 +1,13 @@ +// SPDX-License-Identifier: Apache-2.0 + +//! Decode one pgwire result cell into a `nodedb_types::Value`. + +mod array; +mod binary; +mod decode; +mod error; +mod raw; +mod text; + +pub(crate) use decode::decode_cell; +pub(crate) use raw::RawCell; diff --git a/nodedb-client/src/pg_cell/raw.rs b/nodedb-client/src/pg_cell/raw.rs new file mode 100644 index 000000000..a0bedbcce --- /dev/null +++ b/nodedb-client/src/pg_cell/raw.rs @@ -0,0 +1,22 @@ +// SPDX-License-Identifier: Apache-2.0 + +//! The raw bytes of one result cell, read from a `tokio_postgres::Row`. + +use std::error::Error; + +use tokio_postgres::types::{FromSql, Type}; + +/// The undecoded bytes of a non-NULL cell. It accepts every column type, so +/// `Row::try_get::<_, Option>` hands the bytes to [`super::decode_cell`], +/// which decodes them by the column type. +pub(crate) struct RawCell<'a>(pub(crate) &'a [u8]); + +impl<'a> FromSql<'a> for RawCell<'a> { + fn from_sql(_ty: &Type, raw: &'a [u8]) -> Result> { + Ok(RawCell(raw)) + } + + fn accepts(_ty: &Type) -> bool { + true + } +} diff --git a/nodedb-client/src/pg_cell/text.rs b/nodedb-client/src/pg_cell/text.rs new file mode 100644 index 000000000..58327f299 --- /dev/null +++ b/nodedb-client/src/pg_cell/text.rs @@ -0,0 +1,365 @@ +// SPDX-License-Identifier: Apache-2.0 + +//! Text-format decoders: the PostgreSQL text form of each scalar type. +//! +//! The server sends `bytea`, `json`, `jsonb` and the array types in text +//! format. It sends every type outside the binary scalar set (`numeric`, +//! `uuid`, `interval`, `text`, `varchar`, `name`) as its text bytes. Array +//! elements are text too, so every scalar type has a text decoder here. + +use nodedb_types::{NdbDateTime, NdbDuration, Value}; +use rust_decimal::Decimal; +use tokio_postgres::types::Type; + +use super::error::DecodeReason; + +/// Decode `text` as the text form of the scalar PG type `ty`. +pub(super) fn scalar(ty: &Type, text: &str) -> Result { + match *ty { + Type::BOOL => boolean(text), + Type::INT2 => integer::(text, "int2"), + Type::INT4 => integer::(text, "int4"), + Type::INT8 => integer::(text, "int8"), + Type::FLOAT4 => float4(text).map(|f| Value::Float(f64::from(f))), + Type::FLOAT8 => float8(text).map(Value::Float), + Type::NUMERIC => numeric(text), + Type::TEXT | Type::VARCHAR | Type::NAME | Type::BPCHAR => { + Ok(Value::String(text.to_owned())) + } + Type::BYTEA => bytea(text), + Type::UUID => uuid(text), + Type::JSON | Type::JSONB => json(text), + Type::TIMESTAMP => instant(text).map(Value::NaiveDateTime), + Type::TIMESTAMPTZ => instant(text).map(Value::DateTime), + Type::INTERVAL => { + NdbDuration::parse(text) + .map(Value::Duration) + .ok_or(DecodeReason::Invalid { + expected: "interval", + }) + } + _ => Err(DecodeReason::Unsupported), + } +} + +/// `t` or `f`, as PostgreSQL renders a `bool`. +fn boolean(text: &str) -> Result { + match text { + "t" => Ok(Value::Bool(true)), + "f" => Ok(Value::Bool(false)), + _ => Err(DecodeReason::Invalid { expected: "bool" }), + } +} + +/// A decimal integer that fits `T`. +fn integer(text: &str, expected: &'static str) -> Result +where + T: std::str::FromStr + Into, +{ + text.parse::() + .map(|n| Value::Integer(n.into())) + .map_err(|_| DecodeReason::Invalid { expected }) +} + +/// A `float8` in PostgreSQL text: a finite number, `NaN`, `Infinity` or +/// `-Infinity`. A finite text that overflows `f64` is refused. +fn float8(text: &str) -> Result { + match text { + "NaN" => Ok(f64::NAN), + "Infinity" => Ok(f64::INFINITY), + "-Infinity" => Ok(f64::NEG_INFINITY), + _ => text + .parse::() + .ok() + .filter(|f| f.is_finite()) + .ok_or(DecodeReason::Invalid { expected: "float8" }), + } +} + +/// A `float4` in PostgreSQL text. It parses as `f32`, so the value is the +/// one the server held. +fn float4(text: &str) -> Result { + match text { + "NaN" => Ok(f32::NAN), + "Infinity" => Ok(f32::INFINITY), + "-Infinity" => Ok(f32::NEG_INFINITY), + _ => text + .parse::() + .ok() + .filter(|f| f.is_finite()) + .ok_or(DecodeReason::Invalid { expected: "float4" }), + } +} + +/// A `numeric`. A finite number is an exact `Value::Decimal`: digits past +/// the `Decimal` precision are refused, never rounded. An unsigned integer +/// above `i64::MAX` goes through [`Value::from_u64`], the decimal a native +/// `uint64` decodes to. `Decimal` holds no `NaN` or infinity, so those +/// become the `Value::Float` of the same value. +fn numeric(text: &str) -> Result { + match text { + "NaN" => return Ok(Value::Float(f64::NAN)), + "Infinity" => return Ok(Value::Float(f64::INFINITY)), + "-Infinity" => return Ok(Value::Float(f64::NEG_INFINITY)), + _ => {} + } + if let Ok(unsigned) = text.parse::() + && i64::try_from(unsigned).is_err() + { + return Ok(Value::from_u64(unsigned)); + } + let parsed = if text.contains(['e', 'E']) { + Decimal::from_scientific(text) + } else { + Decimal::from_str_exact(text) + }; + parsed + .map(Value::Decimal) + .map_err(|_| DecodeReason::Invalid { + expected: "numeric", + }) +} + +/// A `bytea` in PostgreSQL hex form: `\x` then two hex digits per byte. +fn bytea(text: &str) -> Result { + let invalid = DecodeReason::Invalid { + expected: "bytea in \\x hex form", + }; + let Some(hex) = text.strip_prefix("\\x") else { + return Err(invalid); + }; + if hex.len() % 2 != 0 { + return Err(invalid); + } + hex.as_bytes() + .as_chunks::<2>() + .0 + .iter() + .map(|&[high, low]| match (hex_digit(high), hex_digit(low)) { + (Some(high), Some(low)) => Ok((high << 4) | low), + _ => Err(invalid.clone()), + }) + .collect::, _>>() + .map(Value::Bytes) +} + +fn hex_digit(byte: u8) -> Option { + char::from(byte) + .to_digit(16) + .and_then(|d| u8::try_from(d).ok()) +} + +/// A `uuid` column cell. The server sends a UUID as 36-character hyphenated +/// hex and a ULID, which shares the `uuid` OID, as 26 Crockford base32 +/// characters. +fn uuid(text: &str) -> Result { + if is_hyphenated_uuid(text) { + Ok(Value::Uuid(text.to_ascii_lowercase())) + } else if is_ulid(text) { + Ok(Value::Ulid(text.to_ascii_uppercase())) + } else { + Err(DecodeReason::Invalid { + expected: "uuid or ulid", + }) + } +} + +fn is_hyphenated_uuid(text: &str) -> bool { + text.len() == 36 + && text.bytes().enumerate().all(|(i, b)| match i { + 8 | 13 | 18 | 23 => b == b'-', + _ => b.is_ascii_hexdigit(), + }) +} + +/// 26 Crockford base32 characters whose first is at most `7`, so the +/// 128-bit value does not overflow. +fn is_ulid(text: &str) -> bool { + text.len() == 26 + && text + .bytes() + .next() + .is_some_and(|b| (b'0'..=b'7').contains(&b)) + && text.bytes().all(|b| { + let upper = b.to_ascii_uppercase(); + upper.is_ascii_digit() + || (upper.is_ascii_uppercase() && !matches!(upper, b'I' | b'L' | b'O' | b'U')) + }) +} + +/// A `json` or `jsonb` document, through `json_to_value`. +fn json(text: &str) -> Result { + let invalid = DecodeReason::Invalid { expected: "json" }; + let parsed: serde_json::Value = sonic_rs::from_str(text).map_err(|_| invalid.clone())?; + crate::remote_parse::json_to_value(&parsed).map_err(|_| invalid) +} + +/// An ISO-8601 instant, as the server renders a timestamp in text. +fn instant(text: &str) -> Result { + NdbDateTime::parse(text).ok_or(DecodeReason::Invalid { + expected: "timestamp", + }) +} + +#[cfg(test)] +mod tests { + use super::*; + + fn ok(ty: &Type, text: &str) -> Value { + scalar(ty, text).unwrap_or_else(|e| panic!("{text:?} as {ty}: {e}")) + } + + fn refused(ty: &Type, text: &str) -> DecodeReason { + match scalar(ty, text) { + Ok(v) => panic!("{text:?} as {ty} must be refused, got {v:?}"), + Err(e) => e, + } + } + + #[test] + fn decodes_bool_and_integers() { + assert_eq!(ok(&Type::BOOL, "t"), Value::Bool(true)); + assert_eq!(ok(&Type::BOOL, "f"), Value::Bool(false)); + refused(&Type::BOOL, "true"); + assert_eq!(ok(&Type::INT2, "-32768"), Value::Integer(-32768)); + assert_eq!(ok(&Type::INT4, "7"), Value::Integer(7)); + assert_eq!( + ok(&Type::INT8, "-9223372036854775808"), + Value::Integer(i64::MIN) + ); + refused(&Type::INT2, "32768"); + refused(&Type::INT8, "1.5"); + } + + #[test] + fn decodes_float_text_including_non_finite() { + assert_eq!(ok(&Type::FLOAT8, "2.5"), Value::Float(2.5)); + assert_eq!(ok(&Type::FLOAT8, "1e+20"), Value::Float(1e20)); + assert_eq!(ok(&Type::FLOAT8, "Infinity"), Value::Float(f64::INFINITY)); + assert_eq!( + ok(&Type::FLOAT8, "-Infinity"), + Value::Float(f64::NEG_INFINITY) + ); + assert!(matches!(ok(&Type::FLOAT8, "NaN"), Value::Float(f) if f.is_nan())); + assert_eq!(ok(&Type::FLOAT4, "0.1"), Value::Float(f64::from(0.1f32))); + assert_eq!( + ok(&Type::FLOAT4, "-Infinity"), + Value::Float(f64::NEG_INFINITY) + ); + assert!(matches!(ok(&Type::FLOAT4, "NaN"), Value::Float(f) if f.is_nan())); + for text in ["inf", "nan", "1e400", "abc", ""] { + refused(&Type::FLOAT8, text); + } + refused(&Type::FLOAT4, "1e39"); + } + + #[test] + fn decodes_numeric_exactly() { + assert_eq!( + ok(&Type::NUMERIC, "12.50"), + Value::Decimal(Decimal::new(1250, 2)) + ); + assert_eq!(ok(&Type::NUMERIC, "-3"), Value::Decimal(Decimal::from(-3))); + assert_eq!( + ok(&Type::NUMERIC, "18446744073709551615"), + Value::from_u64(u64::MAX) + ); + assert_eq!( + ok(&Type::NUMERIC, "9223372036854775807"), + Value::Decimal(Decimal::from(i64::MAX)) + ); + assert_eq!( + ok(&Type::NUMERIC, "1e3"), + Value::Decimal(Decimal::from(1000)) + ); + assert_eq!(ok(&Type::NUMERIC, "Infinity"), Value::Float(f64::INFINITY)); + assert!(matches!(ok(&Type::NUMERIC, "NaN"), Value::Float(f) if f.is_nan())); + // More digits than `Decimal` holds: refused, never rounded. + refused(&Type::NUMERIC, "0.12345678901234567890123456789012"); + refused(&Type::NUMERIC, "12,5"); + } + + #[test] + fn decodes_strings() { + for ty in [Type::TEXT, Type::VARCHAR, Type::NAME, Type::BPCHAR] { + assert_eq!(ok(&ty, "héllo"), Value::String("héllo".into())); + } + } + + #[test] + fn decodes_bytea_hex() { + assert_eq!( + ok(&Type::BYTEA, "\\x00ff7A"), + Value::Bytes(vec![0, 255, 0x7a]) + ); + assert_eq!(ok(&Type::BYTEA, "\\x"), Value::Bytes(Vec::new())); + for text in ["00ff", "\\x0", "\\xzz", "AP8"] { + refused(&Type::BYTEA, text); + } + } + + #[test] + fn decodes_uuid_and_ulid() { + assert_eq!( + ok(&Type::UUID, "550E8400-E29B-41D4-A716-446655440000"), + Value::Uuid("550e8400-e29b-41d4-a716-446655440000".into()) + ); + assert_eq!( + ok(&Type::UUID, "01arz3ndektsv4rrffq69g5fav"), + Value::Ulid("01ARZ3NDEKTSV4RRFFQ69G5FAV".into()) + ); + for text in [ + "550e8400e29b41d4a716446655440000", + "550e8400-e29b-41d4-a716-44665544000g", + "81ARZ3NDEKTSV4RRFFQ69G5FAV", + "01ARZ3NDEKTSV4RRFFQ69G5FAI", + ] { + refused(&Type::UUID, text); + } + } + + #[test] + fn decodes_json_through_json_to_value() { + let Value::Object(map) = ok( + &Type::JSONB, + r#"{"a":[1,2.5,null],"big":18446744073709551615}"#, + ) else { + panic!("a JSON object decodes as an Object"); + }; + assert_eq!( + map.get("a"), + Some(&Value::Array(vec![ + Value::Integer(1), + Value::Float(2.5), + Value::Null + ])) + ); + assert_eq!(map.get("big"), Some(&Value::from_u64(u64::MAX))); + assert_eq!(ok(&Type::JSON, "\"x\""), Value::String("x".into())); + refused(&Type::JSON, "{bad"); + } + + #[test] + fn decodes_timestamps_and_interval() { + let at = NdbDateTime::from_micros(1_583_402_400_000_000); + assert_eq!( + ok(&Type::TIMESTAMP, "2020-03-05T10:00:00.000000Z"), + Value::NaiveDateTime(at) + ); + assert_eq!( + ok(&Type::TIMESTAMPTZ, "2020-03-05 10:00:00"), + Value::DateTime(at) + ); + refused(&Type::TIMESTAMPTZ, "yesterday"); + assert_eq!( + ok(&Type::INTERVAL, "1h30m"), + Value::Duration(NdbDuration::from_micros(5_400_000_000)) + ); + refused(&Type::INTERVAL, "01:30:00"); + } + + #[test] + fn refuses_a_type_without_a_decoder() { + assert_eq!(refused(&Type::POINT, "(1,2)"), DecodeReason::Unsupported); + } +} diff --git a/nodedb-client/src/remote/client/core.rs b/nodedb-client/src/remote/client/core.rs index 04380d2e0..ad579cc29 100644 --- a/nodedb-client/src/remote/client/core.rs +++ b/nodedb-client/src/remote/client/core.rs @@ -84,9 +84,8 @@ impl NodeDbRemote { let mut result_rows = Vec::with_capacity(rows.len()); for row in &rows { let mut vals = Vec::with_capacity(columns.len()); - for (i, col) in row.columns().iter().enumerate() { - let val = pg_value_to_value(row, i, col.type_()); - vals.push(val); + for (idx, column) in row.columns().iter().enumerate() { + vals.push(pg_value_to_value(row, idx, column)?); } result_rows.push(vals); } diff --git a/nodedb-client/src/remote/client/dispatch.rs b/nodedb-client/src/remote/client/dispatch.rs index 6027250dd..7bd355d80 100644 --- a/nodedb-client/src/remote/client/dispatch.rs +++ b/nodedb-client/src/remote/client/dispatch.rs @@ -32,9 +32,10 @@ impl NodeDb for NodeDbRemote { query: &[f32], k: usize, filter: Option<&MetadataFilter>, - _allowed_ids: Option<&std::collections::HashSet>, + allowed_ids: Option<&std::collections::HashSet>, ) -> NodeDbResult> { - self.vector_search_impl(collection, query, k, filter).await + self.vector_search_impl(collection, query, k, filter, allowed_ids) + .await } async fn vector_insert_field( @@ -81,9 +82,10 @@ impl NodeDb for NodeDbRemote { collection: &str, start: &NodeId, depth: u8, + direction: nodedb_types::graph::Direction, edge_filter: Option<&EdgeFilter>, ) -> NodeDbResult { - self.graph_traverse_impl(collection, start, depth, edge_filter) + self.graph_traverse_impl(collection, start, depth, direction, edge_filter) .await } @@ -230,9 +232,9 @@ impl NodeDb for NodeDbRemote { query: &str, top_k: usize, params: TextSearchParams, - _allowed_ids: Option<&std::collections::HashSet>, + allowed_ids: Option<&std::collections::HashSet>, ) -> NodeDbResult> { - self.text_search_impl(collection, field, query, top_k, params) + self.text_search_impl(collection, field, query, top_k, params, allowed_ids) .await } diff --git a/nodedb-client/src/remote/client/document.rs b/nodedb-client/src/remote/client/document.rs index c037cd9d4..2c931eb95 100644 --- a/nodedb-client/src/remote/client/document.rs +++ b/nodedb-client/src/remote/client/document.rs @@ -1,97 +1,46 @@ // SPDX-License-Identifier: Apache-2.0 //! Document operation implementations for `NodeDbRemote`. - -use std::collections::HashMap; +//! +//! A document's fields are the row's top-level columns, the shape the native +//! client and Lite store. The collection's identity column holds the +//! document id. use nodedb_types::document::Document; use nodedb_types::error::{NodeDbError, NodeDbResult}; use nodedb_types::value::Value; +use crate::document_identity::is_identity_cell; use crate::remote_parse::json_to_value; -use crate::sql_escape::quote_identifier; +use crate::sql_escape::{quote_identifier, quote_string_literal}; use super::core::NodeDbRemote; +/// Output column holding the whole row as one JSON object. +const DOCUMENT_COLUMN: &str = "document"; + impl NodeDbRemote { + /// Read every top-level field of the row whose identity column holds + /// `id`, each with its stored type. + /// + /// `to_jsonb(*)` returns the whole row as one JSON object. Each field + /// decodes to the `Value` kind the native client returns. A key lookup + /// stays a point read: the server evaluates the item on the one row the + /// lookup returns. pub(super) async fn document_get_impl( &self, collection: &str, id: &str, ) -> NodeDbResult> { - let quoted = quote_identifier(collection); - let sql = format!("SELECT id, data FROM {quoted} WHERE id = $1"); - let (_, rows) = self.query_raw(&sql, &[&id]).await?; - - let Some(row) = rows.first() else { - return Ok(None); - }; - - let doc_id = row - .first() - .and_then(|v| v.as_str()) - .unwrap_or(id) - .to_string(); - let mut doc = Document::new(doc_id); - - // The `data` column arrives either already structured (native-typed - // cell) or as JSON text (pgwire renders composite values textually). - // Anything else is a shape this client cannot map, and is reported as - // such: silently yielding a field-less document would be - // indistinguishable from a document that genuinely stores no fields. - match row.get(1) { - None | Some(Value::Null) => {} - Some(Value::Object(fields)) => { - for (k, v) in fields { - doc.set(k.clone(), v.clone()); - } - } - Some(Value::String(json_str)) => { - let parsed: HashMap = sonic_rs::from_str(json_str) - .map_err(|e| { - NodeDbError::serialization( - "json", - format!("document_get data column for '{collection}'/'{id}': {e}"), - ) - })?; - for (k, v) in &parsed { - doc.set(k.clone(), json_to_value(v)); - } - } - Some(other) => { - return Err(NodeDbError::serialization( - "json", - format!( - "document_get data column for '{collection}'/'{id}': \ - expected an object or JSON text, got {other:?}" - ), - )); - } - } - - Ok(Some(doc)) - } - - pub(super) async fn document_put_impl( - &self, - collection: &str, - doc: Document, - ) -> NodeDbResult<()> { - let collection = quote_identifier(collection); - let data_json = sonic_rs::to_string(&doc.fields) - .map_err(|e| NodeDbError::storage(format!("document serialization: {e}")))?; - // NodeDB's SQL planner accepts JSON text values directly into - // the document `data` column — no `::jsonb` cast on the - // expression side, which the planner currently rejects as an - // "unsupported value expression". The server interprets the - // string literal as document JSON when the target column is the - // doc-engine `data` column. + let key = self.identity_column(collection).await?; let sql = format!( - "INSERT INTO {collection} (id, data) VALUES ($1, $2) \ - ON CONFLICT (id) DO UPDATE SET data = $2" + "SELECT to_jsonb(*) AS {DOCUMENT_COLUMN} FROM {} WHERE {} = {}", + quote_identifier(collection), + quote_identifier(&key), + quote_string_literal(id) ); - self.execute_raw(&sql, &[&doc.id, &data_json]).await?; - Ok(()) + let (columns, rows) = self.simple_query_raw(&sql).await?; + document_from_rows(collection, id, &columns, rows) } pub(super) async fn document_delete_impl( @@ -99,9 +48,149 @@ impl NodeDbRemote { collection: &str, id: &str, ) -> NodeDbResult<()> { - let collection = quote_identifier(collection); - let sql = format!("DELETE FROM {collection} WHERE id = $1"); + let key = self.identity_column(collection).await?; + let sql = format!( + "DELETE FROM {} WHERE {} = $1", + quote_identifier(collection), + quote_identifier(&key) + ); self.execute_raw(&sql, &[&id]).await?; Ok(()) } } + +/// The document a `to_jsonb(*)` point read answered with, or `None` when no +/// row matched. +/// +/// The one cell is the JSON text of the row object. More than one row, a +/// missing column, or a cell that is not a JSON object is an error naming +/// the document. +fn document_from_rows( + collection: &str, + id: &str, + columns: &[String], + rows: Vec>, +) -> NodeDbResult> { + let mut rows = rows.into_iter(); + let Some(row) = rows.next() else { + return Ok(None); + }; + let extra = rows.count(); + if extra > 0 { + return Err(malformed( + collection, + id, + format!("expected at most one row, got {}", extra + 1), + )); + } + let index = columns + .iter() + .position(|c| c == DOCUMENT_COLUMN) + .ok_or_else(|| { + malformed( + collection, + id, + format!("no '{DOCUMENT_COLUMN}' column; columns are {columns:?}"), + ) + })?; + let text = match row.into_iter().nth(index) { + Some(Value::String(text)) => text, + other => { + return Err(malformed( + collection, + id, + format!("'{DOCUMENT_COLUMN}' cell is not JSON text: {other:?}"), + )); + } + }; + let json: serde_json::Value = sonic_rs::from_str(&text) + .map_err(|e| malformed(collection, id, format!("row JSON does not parse: {e}")))?; + let serde_json::Value::Object(fields) = json else { + return Err(malformed( + collection, + id, + format!("row JSON is not an object: {text}"), + )); + }; + + let mut doc = Document::new(id); + for (name, field) in &fields { + let value = json_to_value(field) + .map_err(|e| malformed(collection, id, format!("field '{name}': {e}")))?; + if !is_identity_cell(name, &value, id) { + doc.set(name.clone(), value); + } + } + Ok(Some(doc)) +} + +fn malformed(collection: &str, id: &str, detail: impl std::fmt::Display) -> NodeDbError { + NodeDbError::serialization( + "pgwire", + format!("document_get '{collection}'/'{id}': {detail}"), + ) +} + +#[cfg(test)] +mod tests { + use super::*; + + fn read(cells: Vec>) -> NodeDbResult> { + document_from_rows("docs", "d1", &[DOCUMENT_COLUMN.to_string()], cells) + } + + fn json_cell(text: &str) -> Vec> { + vec![vec![Value::String(text.into())]] + } + + #[test] + fn every_field_keeps_its_type() { + let doc = read(json_cell( + r#"{"id":"d1","count":5,"ratio":1.5,"whole":2.0,"active":true,"name":"n", + "tags":["a",2],"missing":null,"meta":{"k":"v"}}"#, + )) + .expect("a well-formed row parses") + .expect("a row is a document"); + assert_eq!(doc.id, "d1"); + assert_eq!(doc.get("count"), Some(&Value::Integer(5))); + assert_eq!(doc.get("ratio"), Some(&Value::Float(1.5))); + assert_eq!(doc.get("whole"), Some(&Value::Float(2.0))); + assert_eq!(doc.get("active"), Some(&Value::Bool(true))); + assert_eq!(doc.get("name"), Some(&Value::String("n".into()))); + assert_eq!( + doc.get("tags"), + Some(&Value::Array(vec![ + Value::String("a".into()), + Value::Integer(2) + ])) + ); + assert_eq!(doc.get("missing"), Some(&Value::Null)); + assert!(matches!(doc.get("meta"), Some(Value::Object(_)))); + assert_eq!(doc.get("id"), None, "the identity cell is not a field"); + } + + #[test] + fn no_row_is_no_document() { + assert_eq!(read(Vec::new()).expect("a miss is not a fault"), None); + } + + #[test] + fn a_cell_that_is_not_a_json_object_is_an_error() { + let err = read(json_cell("[1,2]")).expect_err("an array is not a row"); + assert!(err.to_string().contains("not an object"), "{err}"); + let err = read(vec![vec![Value::Null]]).expect_err("a NULL cell is not a row"); + assert!(err.to_string().contains("not JSON text"), "{err}"); + let err = read(json_cell("{")).expect_err("broken JSON"); + assert!(err.to_string().contains("does not parse"), "{err}"); + } + + #[test] + fn several_rows_are_an_error() { + let err = read(vec![ + vec![Value::String("{}".into())], + vec![Value::String("{}".into())], + ]) + .expect_err("a point read never picks one of several rows"); + assert!(err.to_string().contains("one row"), "{err}"); + } +} diff --git a/nodedb-client/src/remote/client/document_bind.rs b/nodedb-client/src/remote/client/document_bind.rs new file mode 100644 index 000000000..3e34da079 --- /dev/null +++ b/nodedb-client/src/remote/client/document_bind.rs @@ -0,0 +1,240 @@ +// SPDX-License-Identifier: Apache-2.0 + +//! The parameterized `INSERT` that writes one whole document over pgwire. +//! +//! Every top-level field is its own column, so SQL filters, projects and +//! text-searches the field by name. Every value is a bound parameter with a +//! declared type: the statement text carries quoted identifiers and +//! placeholders only. An array field binds as `ARRAY[$a, $b, ...]`. + +use nodedb_types::document::Document; +use nodedb_types::error::{NodeDbError, NodeDbResult}; +use nodedb_types::value::Value; +use tokio_postgres::types::{ToSql, Type}; + +use crate::sql_escape::quote_identifier; + +/// One bound parameter and the type the statement declares for it. +pub(super) type TypedParam = (Box, Type); + +/// A whole-document `INSERT` and its typed parameters, in placeholder order. +pub(super) struct DocumentInsert { + pub sql: String, + pub params: Vec, +} + +impl DocumentInsert { + /// The statement's declared parameter types, in placeholder order. + pub fn types(&self) -> Vec { + self.params.iter().map(|(_, ty)| ty.clone()).collect() + } + + /// The parameters as the driver borrows them, in placeholder order. + pub fn values(&self) -> Vec<&(dyn ToSql + Sync)> { + self.params + .iter() + .map(|(value, _)| { + let value: &(dyn ToSql + Sync) = value.as_ref(); + value + }) + .collect() + } +} + +/// Refuse a document whose identity-column field names a different document. +/// +/// `key` is the collection's identity column. A field under it that differs +/// from `Document::id` makes SQL and key reads disagree on which document the +/// row is. +pub(super) fn check_identity_field( + collection: &str, + doc: &Document, + key: &str, +) -> NodeDbResult<()> { + match doc.fields.get(key) { + None => Ok(()), + Some(Value::String(s)) if *s == doc.id => Ok(()), + Some(other) => Err(NodeDbError::bad_request(format!( + "document_put '{collection}'/'{}': field '{key}' holds {other:?}, which differs \ + from the document id; '{key}' is the collection's identity column; remove the \ + field or set it to the document id", + doc.id + ))), + } +} + +/// Build `INSERT INTO (, ) VALUES ($1, ...)`. +/// +/// `key` is the collection's identity column and `$1` is the document id. A +/// field named `key` is that same column: the caller refuses one that +/// differs from the document id first. Fields bind in name order, so one +/// field set always yields one statement text. +pub(super) fn document_insert( + collection: &str, + key: &str, + doc: &Document, +) -> NodeDbResult { + let mut fields: Vec<(&String, &Value)> = doc + .fields + .iter() + .filter(|(name, _)| name.as_str() != key) + .collect(); + fields.sort_by(|a, b| a.0.cmp(b.0)); + + let id_param: TypedParam = (Box::new(doc.id.clone()), Type::TEXT); + let mut params: Vec = vec![id_param]; + let mut columns = vec![quote_identifier(key)]; + let mut values = vec!["$1".to_string()]; + for (name, value) in fields { + columns.push(quote_identifier(name)); + values.push(bind_value(collection, &doc.id, name, value, &mut params)?); + } + + Ok(DocumentInsert { + sql: format!( + "INSERT INTO {} ({}) VALUES ({})", + quote_identifier(collection), + columns.join(", "), + values.join(", ") + ), + params, + }) +} + +/// Push `value`'s parameters and return the SQL expression that names them. +fn bind_value( + collection: &str, + doc_id: &str, + field: &str, + value: &Value, + params: &mut Vec, +) -> NodeDbResult { + let param: TypedParam = match value { + Value::Null => (Box::new(None::), Type::TEXT), + Value::Bool(b) => (Box::new(*b), Type::BOOL), + Value::Integer(i) => (Box::new(*i), Type::INT8), + Value::Float(f) => (Box::new(*f), Type::FLOAT8), + Value::String(s) => (Box::new(s.clone()), Type::TEXT), + Value::Array(items) => { + let elements = items + .iter() + .map(|item| bind_value(collection, doc_id, field, item, params)) + .collect::>>()?; + return Ok(format!("ARRAY[{}]", elements.join(", "))); + } + other => { + return Err(NodeDbError::bad_request(format!( + "document_put '{collection}'/'{doc_id}': field '{field}' holds {other:?}, \ + which pgwire SQL cannot store as a typed document field; \ + write this document with NativeClient" + ))); + } + }; + params.push(param); + Ok(format!("${}", params.len())) +} + +#[cfg(test)] +mod tests { + use super::*; + + fn doc(fields: &[(&str, Value)]) -> Document { + let mut doc = Document::new("d1"); + for (name, value) in fields { + doc.set(*name, value.clone()); + } + doc + } + + #[test] + fn every_field_is_its_own_bound_column() { + let insert = document_insert( + "docs", + "id", + &doc(&[ + ("body", Value::String("x".into())), + ("n", Value::Integer(3)), + ]), + ) + .expect("scalar fields bind"); + assert_eq!( + insert.sql, + "INSERT INTO \"docs\" (\"id\", \"body\", \"n\") VALUES ($1, $2, $3)" + ); + assert_eq!(insert.types(), vec![Type::TEXT, Type::TEXT, Type::INT8]); + assert_eq!(insert.values().len(), 3); + } + + #[test] + fn an_array_field_binds_each_element() { + let insert = document_insert( + "docs", + "id", + &doc(&[( + "tags", + Value::Array(vec![Value::String("a".into()), Value::String("b".into())]), + )]), + ) + .expect("array of scalars binds"); + assert_eq!( + insert.sql, + "INSERT INTO \"docs\" (\"id\", \"tags\") VALUES ($1, ARRAY[$2, $3])" + ); + } + + #[test] + fn an_id_field_is_the_identity_column_not_a_second_one() { + let insert = document_insert("docs", "id", &doc(&[("id", Value::String("d1".into()))])) + .expect("matching id binds"); + assert_eq!(insert.sql, "INSERT INTO \"docs\" (\"id\") VALUES ($1)"); + } + + #[test] + fn an_identity_field_naming_another_document_is_refused() { + assert!(check_identity_field("c", &doc(&[]), "id").is_ok()); + assert!( + check_identity_field("c", &doc(&[("id", Value::String("d1".into()))]), "id").is_ok() + ); + let err = check_identity_field("c", &doc(&[("id", Value::String("d2".into()))]), "id") + .expect_err("conflicting id"); + assert!(err.to_string().contains("d1"), "{err}"); + } + + #[test] + fn a_declared_key_is_checked_and_id_is_an_ordinary_field() { + let other_id = doc(&[("id", Value::String("other".into()))]); + assert!(check_identity_field("c", &other_id, "sku").is_ok()); + let err = check_identity_field("c", &doc(&[("sku", Value::String("d2".into()))]), "sku") + .expect_err("conflicting key"); + assert!(err.to_string().contains("'sku'"), "{err}"); + } + + #[test] + fn a_declared_key_is_the_identity_column_and_id_is_a_field() { + let insert = document_insert( + "items", + "sku", + &doc(&[ + ("sku", Value::String("d1".into())), + ("id", Value::String("other".into())), + ]), + ) + .expect("matching key binds"); + assert_eq!( + insert.sql, + "INSERT INTO \"items\" (\"sku\", \"id\") VALUES ($1, $2)" + ); + } + + #[test] + fn a_nested_object_is_refused_naming_the_field() { + let Err(err) = document_insert( + "docs", + "id", + &doc(&[("meta", Value::Object(std::collections::HashMap::new()))]), + ) else { + panic!("object fields have no SQL value form"); + }; + assert!(err.to_string().contains("meta")); + } +} diff --git a/nodedb-client/src/remote/client/document_put.rs b/nodedb-client/src/remote/client/document_put.rs new file mode 100644 index 000000000..fabc6353a --- /dev/null +++ b/nodedb-client/src/remote/client/document_put.rs @@ -0,0 +1,176 @@ +// SPDX-License-Identifier: Apache-2.0 + +//! The remote whole-document replace. +//! +//! The replace is a `DELETE` then an `INSERT` of the new field set. An insert +//! merged over the old row would keep fields the new document dropped. +//! +//! The pair is atomic in both connection states: +//! - Inside the caller's transaction block, it runs under a savepoint. The put +//! never commits or rolls back the caller's block. +//! - Outside a block, it runs in its own `BEGIN` / `COMMIT`. +//! +//! The driver does not expose the server's transaction status. The put opens +//! its savepoint first. The server refuses a savepoint outside a block with +//! SQLSTATE `25P01`, and that refusal selects the standalone path. +//! +//! A failed put inside a block rolls back to its savepoint and releases it. +//! The caller's block stays usable and keeps every write made before the put. +//! A put into an already aborted block is refused with SQLSTATE `25P02`, and +//! the block stays aborted. + +use nodedb_types::document::Document; +use nodedb_types::error::{NodeDbError, NodeDbResult}; +use tokio_postgres::Client; +use tokio_postgres::error::SqlState; +use tokio_postgres::types::Type; + +use crate::sql_escape::quote_identifier; + +use super::core::{NodeDbRemote, pg_error_detail}; +use super::document_bind::{DocumentInsert, check_identity_field, document_insert}; + +/// The savepoint a put opens inside the caller's transaction block. +const PUT_SAVEPOINT: &str = "nodedb_document_put"; + +/// What the put opened around its statements, and so what closes them. +enum PutScope { + /// A savepoint inside the caller's transaction block. + Savepoint, + /// The put's own transaction block. + Transaction, +} + +/// The document a put writes. +struct PutTarget<'a> { + collection: &'a str, + /// The collection's identity column. + key: &'a str, + id: &'a str, +} + +impl PutTarget<'_> { + fn error(&self, step: &str, e: &tokio_postgres::Error) -> NodeDbError { + NodeDbError::storage(format!( + "document_put '{}'/'{}' {step} failed: {}", + self.collection, + self.id, + pg_error_detail(e) + )) + } + + fn undo_error( + &self, + cause: &NodeDbError, + step: &str, + e: &tokio_postgres::Error, + ) -> NodeDbError { + NodeDbError::storage(format!( + "document_put '{}'/'{}' failed: {cause}; the {step} that undoes it also failed: {}", + self.collection, + self.id, + pg_error_detail(e) + )) + } +} + +impl NodeDbRemote { + /// Replace the document's whole field set, creating it when absent. + pub(super) async fn document_put_impl( + &self, + collection: &str, + doc: Document, + ) -> NodeDbResult<()> { + let key = self.identity_column(collection).await?; + check_identity_field(collection, &doc, &key)?; + let insert = document_insert(collection, &key, &doc)?; + let target = PutTarget { + collection, + key: &key, + id: &doc.id, + }; + let client = self.client.lock().await; + let scope = open_scope(&client, &target).await?; + let written = replace(&client, &target, &insert).await; + close_scope(&client, &target, scope, written).await + } +} + +/// Open a savepoint inside the caller's block, else the put's own block. +async fn open_scope(client: &Client, target: &PutTarget<'_>) -> NodeDbResult { + match client + .batch_execute(&format!("SAVEPOINT {PUT_SAVEPOINT}")) + .await + { + Ok(()) => Ok(PutScope::Savepoint), + Err(e) if e.code() == Some(&SqlState::NO_ACTIVE_SQL_TRANSACTION) => { + client + .batch_execute("BEGIN") + .await + .map_err(|e| target.error("begin", &e))?; + Ok(PutScope::Transaction) + } + Err(e) => Err(target.error("savepoint", &e)), + } +} + +/// Delete the old row and insert the new field set. +async fn replace( + client: &Client, + target: &PutTarget<'_>, + insert: &DocumentInsert, +) -> NodeDbResult<()> { + let delete_sql = format!( + "DELETE FROM {} WHERE {} = $1", + quote_identifier(target.collection), + quote_identifier(target.key) + ); + let delete = client + .prepare_typed(&delete_sql, &[Type::TEXT]) + .await + .map_err(|e| target.error("delete prepare", &e))?; + client + .execute(&delete, &[&target.id]) + .await + .map_err(|e| target.error("delete", &e))?; + let statement = client + .prepare_typed(&insert.sql, &insert.types()) + .await + .map_err(|e| target.error("insert prepare", &e))?; + client + .execute(&statement, &insert.values()) + .await + .map_err(|e| target.error("insert", &e))?; + Ok(()) +} + +/// Keep the write on success, undo it on error. +async fn close_scope( + client: &Client, + target: &PutTarget<'_>, + scope: PutScope, + written: NodeDbResult<()>, +) -> NodeDbResult<()> { + match (scope, written) { + (PutScope::Savepoint, Ok(())) => client + .batch_execute(&format!("RELEASE SAVEPOINT {PUT_SAVEPOINT}")) + .await + .map_err(|e| target.error("release savepoint", &e)), + (PutScope::Transaction, Ok(())) => client + .batch_execute("COMMIT") + .await + .map_err(|e| target.error("commit", &e)), + (PutScope::Savepoint, Err(cause)) => { + let undo = + format!("ROLLBACK TO SAVEPOINT {PUT_SAVEPOINT}; RELEASE SAVEPOINT {PUT_SAVEPOINT}"); + match client.batch_execute(&undo).await { + Ok(()) => Err(cause), + Err(e) => Err(target.undo_error(&cause, "rollback to savepoint", &e)), + } + } + (PutScope::Transaction, Err(cause)) => match client.batch_execute("ROLLBACK").await { + Ok(()) => Err(cause), + Err(e) => Err(target.undo_error(&cause, "rollback", &e)), + }, + } +} diff --git a/nodedb-client/src/remote/client/graph.rs b/nodedb-client/src/remote/client/graph.rs index eae65d91d..fab3ba9d5 100644 --- a/nodedb-client/src/remote/client/graph.rs +++ b/nodedb-client/src/remote/client/graph.rs @@ -2,17 +2,14 @@ //! Graph operation implementations for `NodeDbRemote`. -use std::collections::HashMap; - use nodedb_types::document::Document; use nodedb_types::error::{NodeDbError, NodeDbResult}; use nodedb_types::filter::EdgeFilter; use nodedb_types::graph::GraphStats; use nodedb_types::id::{EdgeId, NodeId}; -use nodedb_types::result::{SubGraph, SubGraphEdge, SubGraphNode}; +use nodedb_types::result::SubGraph; use nodedb_types::value::Value; -use super::super::parse::parse_graph_traverse_json; use super::core::NodeDbRemote; use crate::sql_escape::quote_string_literal; @@ -22,71 +19,18 @@ impl NodeDbRemote { collection: &str, start: &NodeId, depth: u8, + direction: nodedb_types::graph::Direction, edge_filter: Option<&EdgeFilter>, ) -> NodeDbResult { - // Server-side DSL: `GRAPH TRAVERSE IN '' FROM '' - // DEPTH [LABEL '']`. The collection is not decorative: the - // server authorizes the traversal against it, and without it the walk - // would cross every collection in the tenant. - let label_clause = edge_filter - .and_then(|f| f.labels.first()) - .map(|l| format!(" LABEL {}", quote_string_literal(l))) - .unwrap_or_default(); - let collection_lit = quote_string_literal(collection); - let start_lit = quote_string_literal(start.as_str()); - let sql = format!( - "GRAPH TRAVERSE IN {collection_lit} FROM {start_lit} DEPTH {depth}{label_clause}" - ); - + let sql = crate::graph_dsl::build_graph_traverse_sql( + collection, + start, + depth, + direction, + edge_filter, + )?; let (columns, rows) = self.simple_query_raw(&sql).await?; - - if columns.len() == 1 && columns[0] == "result" { - if let Some(row) = rows.first() - && let Some(Value::String(json_text)) = row.first() - { - return parse_graph_traverse_json(json_text); - } - return Ok(SubGraph::empty()); - } - - // Structured: node_id, depth, edge_src, edge_dst, edge_label columns. - let mut nodes = Vec::new(); - let mut edges = Vec::new(); - let mut seen_nodes = std::collections::HashSet::new(); - - for row in &rows { - let node_id_str = row.first().and_then(|v| v.as_str()).unwrap_or(""); - let d = row.get(1).and_then(|v| v.as_i64()).unwrap_or(0) as u8; - - if seen_nodes.insert(node_id_str.to_string()) { - nodes.push(SubGraphNode { - id: NodeId::from_validated(node_id_str.to_owned()), - depth: d, - properties: HashMap::new(), - }); - } - - if let (Some(src), Some(dst), Some(label)) = ( - row.get(2).and_then(|v| v.as_str()), - row.get(3).and_then(|v| v.as_str()), - row.get(4).and_then(|v| v.as_str()), - ) { - edges.push(SubGraphEdge { - id: EdgeId::try_first( - NodeId::from_validated(src.to_owned()), - NodeId::from_validated(dst.to_owned()), - label, - ) - .expect("server wire label already validated"), - from: NodeId::from_validated(src.to_owned()), - to: NodeId::from_validated(dst.to_owned()), - label: label.to_string(), - properties: HashMap::new(), - }); - } - } - - Ok(SubGraph { nodes, edges }) + crate::graph_dsl::decode_traverse_result(&columns, &rows) } pub(super) async fn graph_insert_edge_impl( @@ -103,9 +47,12 @@ impl NodeDbRemote { // the collection's graph overlay. simple_query is required // because the DSL doesn't fit the extended-query row-description // shape (it returns CommandComplete only). + // The property object is the document's fields, as the native client + // sends them. The document id is not an edge property. let props_clause = match properties { Some(d) => { - let json = sonic_rs::to_string(&d) + let object = serde_json::Value::from(Value::Object(d.fields)); + let json = sonic_rs::to_string(&object) .map_err(|e| NodeDbError::storage(format!("properties serialization: {e}")))?; format!(" PROPERTIES {}", quote_string_literal(&json)) } @@ -179,42 +126,13 @@ impl NodeDbRemote { max_depth: u8, edge_filter: Option<&EdgeFilter>, ) -> NodeDbResult>> { - // Use the server's `GRAPH PATH` operator instead of the trait - // default's per-hop BFS — one round-trip vs O(path_length). - // Like `graph_traverse`, the path is scoped to `collection`: the server - // authorizes against it and walks only that collection's edges. - let label_clause = edge_filter - .and_then(|f| f.labels.first()) - .map(|l| format!(" LABEL {}", quote_string_literal(l))) - .unwrap_or_default(); - let collection_lit = quote_string_literal(collection); - let from_s = quote_string_literal(from.as_str()); - let to_s = quote_string_literal(to.as_str()); - let sql = format!( - "GRAPH PATH IN {collection_lit} FROM {from_s} TO {to_s} \ - MAX_DEPTH {max_depth}{label_clause}" - ); - - let (_columns, rows) = self.simple_query_raw(&sql).await?; - // Server emits a single `result` column carrying a JSON array - // of node ids — empty array means unreachable. - let Some(row) = rows.first() else { - return Ok(None); - }; - let Some(Value::String(json_text)) = row.first() else { - return Ok(None); - }; - let parsed: Vec = sonic_rs::from_str(json_text) - .map_err(|e| NodeDbError::storage(format!("graph shortest path response: {e}")))?; - if parsed.is_empty() { - return Ok(None); - } - Ok(Some( - parsed - .into_iter() - .map(NodeId::from_validated) - .collect::>(), - )) + // The server's `GRAPH PATH` answers in one round trip. The path is + // scoped to `collection`: the server authorizes against it and walks + // only that collection's edges. + let sql = + crate::graph_dsl::build_graph_path_sql(collection, from, to, max_depth, edge_filter)?; + let (columns, rows) = self.simple_query_raw(&sql).await?; + crate::graph_dsl::decode_path_result(&columns, &rows) } } diff --git a/nodedb-client/src/remote/client/identity.rs b/nodedb-client/src/remote/client/identity.rs new file mode 100644 index 000000000..07710e89f --- /dev/null +++ b/nodedb-client/src/remote/client/identity.rs @@ -0,0 +1,48 @@ +// SPDX-License-Identifier: Apache-2.0 + +//! The identity column of a collection, read from the server catalog. +//! +//! The SQL the remote client generates names the key column: the declared +//! primary key, else `id`. The client reads it from `DESCRIBE` on every +//! call that needs it, so a dropped and recreated collection never answers +//! with a stale key. + +use nodedb_types::error::NodeDbResult; +use nodedb_types::value::Value; + +use crate::document_identity::identity_column_from_describe; +use crate::sql_escape::quote_identifier; + +use super::core::NodeDbRemote; + +impl NodeDbRemote { + /// The identity column of `collection`. + pub(super) async fn identity_column(&self, collection: &str) -> NodeDbResult { + let sql = format!("DESCRIBE {}", quote_identifier(collection)); + let (columns, rows) = self.simple_query_raw(&sql).await?; + identity_column_from_describe(collection, &columns, &rows, pg_text_bool) + } +} + +/// A `bool` cell in the PostgreSQL text format the simple-query protocol +/// carries: `t` or `f`. +fn pg_text_bool(cell: &Value) -> Option { + match cell { + Value::String(s) if s == "t" => Some(true), + Value::String(s) if s == "f" => Some(false), + _ => None, + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn only_pg_text_bools_decode() { + assert_eq!(pg_text_bool(&Value::String("t".into())), Some(true)); + assert_eq!(pg_text_bool(&Value::String("f".into())), Some(false)); + assert_eq!(pg_text_bool(&Value::String("true".into())), None); + assert_eq!(pg_text_bool(&Value::Bool(true)), None); + } +} diff --git a/nodedb-client/src/remote/client/mod.rs b/nodedb-client/src/remote/client/mod.rs index 15aac9ade..b363a685f 100644 --- a/nodedb-client/src/remote/client/mod.rs +++ b/nodedb-client/src/remote/client/mod.rs @@ -3,8 +3,12 @@ pub mod core; mod dispatch; mod document; +mod document_bind; +mod document_put; mod graph; +mod identity; mod sql_lifecycle; +mod text_search; mod vector; pub use core::NodeDbRemote; diff --git a/nodedb-client/src/remote/client/sql_lifecycle.rs b/nodedb-client/src/remote/client/sql_lifecycle.rs index 6138d1f95..e65da11f2 100644 --- a/nodedb-client/src/remote/client/sql_lifecycle.rs +++ b/nodedb-client/src/remote/client/sql_lifecycle.rs @@ -2,16 +2,13 @@ //! SQL execution and collection lifecycle implementations for `NodeDbRemote`. -use std::collections::HashMap; - use nodedb_types::dropped_collection::DroppedCollection; use nodedb_types::error::NodeDbResult; -use nodedb_types::result::{QueryResult, SearchResult}; -use nodedb_types::text_search::TextSearchParams; +use nodedb_types::result::QueryResult; use nodedb_types::value::Value; use crate::row_decode::parse_dropped_collection_rows; -use crate::sql_escape::{quote_identifier, quote_string_literal}; +use crate::sql_escape::quote_identifier; use super::super::sql::translate_params; use super::core::NodeDbRemote; @@ -74,70 +71,4 @@ impl NodeDbRemote { let (_columns, rows) = self.query_raw(sql, &[]).await?; parse_dropped_collection_rows(&rows) } - - pub(super) async fn text_search_impl( - &self, - collection: &str, - field: &str, - query: &str, - top_k: usize, - params: TextSearchParams, - ) -> NodeDbResult> { - // Server-side FTS query SQL: `text_match(, '')` in - // a WHERE clause selects matching ids; `bm25_score(, - // '')` in the SELECT list exposes the score so callers - // can order/rank. The planner pattern-matches this shape and - // dispatches `SqlPlan::TextSearch`. - // - // `params` (mode, fuzzy, prefix, etc.) is intentionally ignored - // for now — every supported option is also expressible in the - // SQL form, but threading them through the DSL string is its - // own widening. The defaults (Plain query with fuzzy=true) cover - // the common case the trait's spec calls out. - let _ = params; - let coll = quote_identifier(collection); - let field_quoted = quote_identifier(field); - let q_lit = quote_string_literal(query); - let sql = format!( - "SELECT id, bm25_score({field_quoted}, {q_lit}) AS score \ - FROM {coll} \ - WHERE text_match({field_quoted}, {q_lit}) \ - LIMIT {top_k}" - ); - - let (columns, rows) = self.simple_query_raw(&sql).await?; - let id_idx = columns.iter().position(|c| c == "id").unwrap_or(0); - let score_idx = columns.iter().position(|c| c == "score").unwrap_or(1); - - let mut results = Vec::with_capacity(rows.len()); - for row in &rows { - let id = row - .get(id_idx) - .and_then(|v| v.as_str()) - .unwrap_or("") - .to_string(); - // simple_query returns text — score arrives as a stringified - // float. Parse defensively so a missing/malformed score does - // not torpedo the whole result set; callers prefer ordered - // ids with score 0.0 over an Err. - let score = row - .get(score_idx) - .and_then(|v| v.as_str()) - .and_then(|s| s.parse::().ok()) - .or_else(|| { - row.get(score_idx) - .and_then(|v| v.as_f64()) - .map(|f| f as f32) - }) - .unwrap_or(0.0); - - results.push(SearchResult { - id, - node_id: None, - distance: score, - metadata: HashMap::new(), - }); - } - Ok(results) - } } diff --git a/nodedb-client/src/remote/client/text_search.rs b/nodedb-client/src/remote/client/text_search.rs new file mode 100644 index 000000000..dfd3d472c --- /dev/null +++ b/nodedb-client/src/remote/client/text_search.rs @@ -0,0 +1,47 @@ +// SPDX-License-Identifier: Apache-2.0 + +//! Remote full-text search over pgwire. +//! +//! The statement comes from the shared builder in `search_sql`, so the +//! native client sends the same one. + +use std::collections::HashSet; + +use nodedb_types::error::NodeDbResult; +use nodedb_types::result::SearchResult; +use nodedb_types::text_search::TextSearchParams; + +use crate::row_decode::decode_search_hits; +use crate::search_sql::{TextSearchRequest, text_hit_source, text_search_sql}; + +use super::core::NodeDbRemote; + +impl NodeDbRemote { + pub(super) async fn text_search_impl( + &self, + collection: &str, + field: &str, + query: &str, + top_k: usize, + params: TextSearchParams, + allowed_ids: Option<&HashSet>, + ) -> NodeDbResult> { + if allowed_ids.is_some_and(HashSet::is_empty) || top_k == 0 { + return Ok(Vec::new()); + } + let key = match allowed_ids { + Some(_) => Some(self.identity_column(collection).await?), + None => None, + }; + let sql = text_search_sql(&TextSearchRequest { + collection, + field, + query, + top_k, + params: ¶ms, + allowed: key.as_deref().zip(allowed_ids), + }); + let (columns, rows) = self.simple_query_raw(&sql).await?; + decode_search_hits(&text_hit_source(collection), &columns, &rows) + } +} diff --git a/nodedb-client/src/remote/client/vector.rs b/nodedb-client/src/remote/client/vector.rs index 8476b0f8f..7e374489a 100644 --- a/nodedb-client/src/remote/client/vector.rs +++ b/nodedb-client/src/remote/client/vector.rs @@ -2,7 +2,7 @@ //! Vector operation implementations for `NodeDbRemote`. -use std::collections::HashMap; +use std::collections::HashSet; use nodedb_types::document::Document; use nodedb_types::error::{NodeDbError, NodeDbResult}; @@ -10,9 +10,10 @@ use nodedb_types::filter::MetadataFilter; use nodedb_types::result::SearchResult; use crate::remote_parse::format_vector_array; +use crate::row_decode::search_hit::{DISTANCE_COLUMN, ID_COLUMN}; +use crate::row_decode::{HitSource, decode_search_hits}; use crate::sql_escape::quote_identifier; -use super::super::parse::parse_vector_search_json; use super::super::sql::{build_vector_search_sql, render_metadata_filter_public}; use super::core::NodeDbRemote; @@ -23,43 +24,31 @@ impl NodeDbRemote { query: &[f32], k: usize, filter: Option<&MetadataFilter>, + allowed_ids: Option<&HashSet>, ) -> NodeDbResult> { - let sql = build_vector_search_sql(collection, query, k, filter)?; - - let (columns, rows) = self.query_raw(&sql, &[]).await?; - - // The DSL path returns JSON in a single "result" column. - if columns.len() == 1 && columns[0] == "result" { - if let Some(row) = rows.first() - && let Some(nodedb_types::value::Value::String(json_text)) = row.first() - { - return parse_vector_search_json(json_text); - } + // An empty allowed set admits no candidate. + if allowed_ids.is_some_and(HashSet::is_empty) { return Ok(Vec::new()); } + // The allowed ids restrict the key column, so the server lowers them + // to the candidate set the index search ranks within. + let key = match allowed_ids { + Some(_) => Some(self.identity_column(collection).await?), + None => None, + }; + let allowed = key.as_deref().zip(allowed_ids); + let sql = build_vector_search_sql(collection, query, k, filter, allowed)?; - // Structured result set: id, distance columns. - let mut results = Vec::with_capacity(rows.len()); - let id_idx = columns.iter().position(|c| c == "id").unwrap_or(0); - let dist_idx = columns.iter().position(|c| c == "distance").unwrap_or(1); - - for row in &rows { - let id = row - .get(id_idx) - .and_then(|v| v.as_str()) - .unwrap_or("") - .to_string(); - let distance = row.get(dist_idx).and_then(|v| v.as_f64()).unwrap_or(0.0) as f32; - - results.push(SearchResult { - id, - node_id: None, - distance, - metadata: HashMap::new(), - }); - } - - Ok(results) + let (columns, rows) = self.query_raw(&sql, &[]).await?; + decode_search_hits( + &HitSource { + op: "vector_search", + collection, + score_column: DISTANCE_COLUMN, + }, + &columns, + &rows, + ) } pub(super) async fn vector_insert_field_impl( @@ -70,20 +59,11 @@ impl NodeDbRemote { embedding: &[f32], metadata: Option, ) -> NodeDbResult<()> { - // Field-aware path: emit `INSERT INTO (id, [, - // metadata]) VALUES ($1, ARRAY[...]<, $2>)` so the vector lands - // on the column named by the trait — not on whichever vector - // column the planner picks when the column name is omitted. - let coll = quote_identifier(collection); - let field = quote_identifier(field_name); - let vec_lit = format_vector_array(embedding); - - let sql = match metadata { - Some(_) => { - format!("INSERT INTO {coll} (id, {field}, metadata) VALUES ($1, {vec_lit}, $2)") - } - None => format!("INSERT INTO {coll} (id, {field}) VALUES ($1, {vec_lit})"), - }; + // The vector lands on the column the caller names, not on whichever + // vector column the planner picks when the column name is omitted. + let key = self.identity_column(collection).await?; + let metadata_param = metadata.as_ref().map(|_| "$2"); + let sql = vector_insert_sql(collection, &key, field_name, embedding, metadata_param); if let Some(d) = metadata { let meta_json = sonic_rs::to_string(&d) @@ -118,32 +98,22 @@ impl NodeDbRemote { None => String::new(), }; let sql = format!( - "SELECT id, vector_distance({field}, {vec_lit}) AS distance \ + "SELECT {ID_COLUMN}, vector_distance({field}, {vec_lit}) AS {DISTANCE_COLUMN} \ FROM {coll}{where_clause} \ ORDER BY vector_distance({field}, {vec_lit}) \ LIMIT {k}" ); let (columns, rows) = self.query_raw(&sql, &[]).await?; - let id_idx = columns.iter().position(|c| c == "id").unwrap_or(0); - let dist_idx = columns.iter().position(|c| c == "distance").unwrap_or(1); - - let mut results = Vec::with_capacity(rows.len()); - for row in &rows { - let id = row - .get(id_idx) - .and_then(|v| v.as_str()) - .unwrap_or("") - .to_string(); - let distance = row.get(dist_idx).and_then(|v| v.as_f64()).unwrap_or(0.0) as f32; - results.push(SearchResult { - id, - node_id: None, - distance, - metadata: HashMap::new(), - }); - } - Ok(results) + decode_search_hits( + &HitSource { + op: "vector_search_field", + collection, + score_column: DISTANCE_COLUMN, + }, + &columns, + &rows, + ) } pub(super) async fn vector_insert_impl( @@ -153,25 +123,77 @@ impl NodeDbRemote { embedding: &[f32], metadata: Option, ) -> NodeDbResult<()> { - let collection = quote_identifier(collection); + let key = self.identity_column(collection).await?; let meta_json = match metadata { Some(d) => sonic_rs::to_string(&d) .map_err(|e| NodeDbError::storage(format!("metadata serialization: {e}")))?, None => "{}".into(), }; - - let sql = format!( - "INSERT INTO {collection} (id, embedding, metadata) VALUES ($1, {}, $2::jsonb)", - format_vector_array(embedding), + let sql = vector_insert_sql( + collection, + &key, + DEFAULT_VECTOR_COLUMN, + embedding, + Some("$2::jsonb"), ); self.execute_raw(&sql, &[&id, &meta_json]).await?; Ok(()) } pub(super) async fn vector_delete_impl(&self, collection: &str, id: &str) -> NodeDbResult<()> { - let collection = quote_identifier(collection); - let sql = format!("DELETE FROM {collection} WHERE id = $1"); + let key = self.identity_column(collection).await?; + let sql = format!( + "DELETE FROM {} WHERE {} = $1", + quote_identifier(collection), + quote_identifier(&key) + ); self.execute_raw(&sql, &[&id]).await?; Ok(()) } } + +/// The vector column a field-less `vector_insert` writes. +const DEFAULT_VECTOR_COLUMN: &str = "embedding"; + +/// `INSERT INTO (, [, metadata]) VALUES ($1, +/// ARRAY[...][, ])`. +/// +/// `key` is the collection's identity column, which holds the id bound as +/// `$1`. `metadata` is the SQL expression of the metadata parameter, or +/// `None` for no metadata column. +fn vector_insert_sql( + collection: &str, + key: &str, + field: &str, + embedding: &[f32], + metadata: Option<&str>, +) -> String { + let collection = quote_identifier(collection); + let key = quote_identifier(key); + let field = quote_identifier(field); + let vector = format_vector_array(embedding); + match metadata { + Some(param) => format!( + "INSERT INTO {collection} ({key}, {field}, metadata) VALUES ($1, {vector}, {param})" + ), + None => format!("INSERT INTO {collection} ({key}, {field}) VALUES ($1, {vector})"), + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn a_vector_insert_writes_the_id_under_the_identity_column() { + assert_eq!( + vector_insert_sql("vecs", "sku", "embedding", &[1.0, 0.5], Some("$2::jsonb")), + "INSERT INTO \"vecs\" (\"sku\", \"embedding\", metadata) \ + VALUES ($1, ARRAY[1,0.5], $2::jsonb)" + ); + assert_eq!( + vector_insert_sql("vecs", "id", "img", &[2.0], None), + "INSERT INTO \"vecs\" (\"id\", \"img\") VALUES ($1, ARRAY[2])" + ); + } +} diff --git a/nodedb-client/src/remote/mod.rs b/nodedb-client/src/remote/mod.rs index 1c7e00cf0..35d0f05cb 100644 --- a/nodedb-client/src/remote/mod.rs +++ b/nodedb-client/src/remote/mod.rs @@ -4,11 +4,11 @@ //! into SQL/DSL and sends them to the NodeDB Origin. //! //! Split into per-concern files: connection lifecycle and trait impl in -//! `client`, SQL/param translation seams in `sql`, JSON response parsing -//! in `parse`. Each file holds the unit tests for the code it owns. +//! `client`, SQL/param translation seams in `sql`. Search-hit rows decode +//! through the crate-level `row_decode` module. Each file holds the unit +//! tests for the code it owns. pub mod client; -mod parse; mod sql; pub use client::NodeDbRemote; diff --git a/nodedb-client/src/remote/parse.rs b/nodedb-client/src/remote/parse.rs deleted file mode 100644 index 9208f977c..000000000 --- a/nodedb-client/src/remote/parse.rs +++ /dev/null @@ -1,122 +0,0 @@ -// SPDX-License-Identifier: Apache-2.0 - -//! JSON response parsing for the remote client. -//! -//! The DSL paths (`SEARCH ... USING VECTOR`, graph traversal) return -//! results as JSON in a single text column. These helpers decode that -//! into typed `SearchResult` / `SubGraph` values for the trait surface. -//! Row-shaped responses (system catalog tables) go through -//! [`crate::row_decode`] instead so the remote and trait-default decoders -//! share one parser. - -use std::collections::HashMap; - -use nodedb_types::error::{NodeDbError, NodeDbResult}; -use nodedb_types::id::{EdgeId, NodeId}; -use nodedb_types::result::{SearchResult, SubGraph, SubGraphEdge, SubGraphNode}; - -/// Parse a JSON string from the DSL's "result" column into `Vec`. -pub(super) fn parse_vector_search_json(json_text: &str) -> NodeDbResult> { - let parsed: serde_json::Value = sonic_rs::from_str(json_text) - .map_err(|e| NodeDbError::serialization("json", e.to_string()))?; - - let mut results = Vec::new(); - if let Some(arr) = parsed.as_array() { - for item in arr { - let id = item - .get("id") - .and_then(|v| v.as_str()) - .unwrap_or("") - .to_string(); - let distance = item.get("distance").and_then(|v| v.as_f64()).unwrap_or(0.0) as f32; - results.push(SearchResult { - id, - node_id: None, - distance, - metadata: HashMap::new(), - }); - } - } - - Ok(results) -} - -/// Parse a JSON string from graph_traverse into `SubGraph`. -pub(super) fn parse_graph_traverse_json(json_text: &str) -> NodeDbResult { - let parsed: serde_json::Value = sonic_rs::from_str(json_text) - .map_err(|e| NodeDbError::serialization("json", e.to_string()))?; - - let mut nodes = Vec::new(); - let mut edges = Vec::new(); - - if let Some(n) = parsed.get("nodes").and_then(|v| v.as_array()) { - for item in n { - let id = item - .get("id") - .and_then(|v| v.as_str()) - .unwrap_or("") - .to_string(); - let depth = item.get("depth").and_then(|v| v.as_u64()).unwrap_or(0) as u8; - nodes.push(SubGraphNode { - id: NodeId::from_validated(id), - depth, - properties: HashMap::new(), - }); - } - } - - if let Some(e) = parsed.get("edges").and_then(|v| v.as_array()) { - for item in e { - let src = item.get("from").and_then(|v| v.as_str()).unwrap_or(""); - let dst = item.get("to").and_then(|v| v.as_str()).unwrap_or(""); - let label = item.get("label").and_then(|v| v.as_str()).unwrap_or(""); - edges.push(SubGraphEdge { - id: EdgeId::try_first( - NodeId::from_validated(src.to_owned()), - NodeId::from_validated(dst.to_owned()), - label, - ) - .expect("server wire label already validated"), - from: NodeId::from_validated(src.to_owned()), - to: NodeId::from_validated(dst.to_owned()), - label: label.to_string(), - properties: HashMap::new(), - }); - } - } - - Ok(SubGraph { nodes, edges }) -} - -#[cfg(test)] -mod tests { - use super::*; - - #[test] - fn parse_vector_search_json_works() { - let json = r#"[{"id":"v1","distance":0.1},{"id":"v2","distance":0.5}]"#; - let results = parse_vector_search_json(json).unwrap(); - assert_eq!(results.len(), 2); - assert_eq!(results[0].id, "v1"); - assert!((results[0].distance - 0.1).abs() < 0.001); - assert_eq!(results[1].id, "v2"); - } - - #[test] - fn parse_graph_traverse_json_works() { - let json = r#"{ - "nodes": [{"id":"a","depth":0},{"id":"b","depth":1}], - "edges": [{"from":"a","to":"b","label":"KNOWS"}] - }"#; - let sg = parse_graph_traverse_json(json).unwrap(); - assert_eq!(sg.node_count(), 2); - assert_eq!(sg.edge_count(), 1); - assert_eq!(sg.edges[0].label, "KNOWS"); - } - - #[test] - fn parse_empty_search_json() { - let results = parse_vector_search_json("[]").unwrap(); - assert!(results.is_empty()); - } -} diff --git a/nodedb-client/src/remote/sql.rs b/nodedb-client/src/remote/sql.rs index 14b1b0b2d..98f2837df 100644 --- a/nodedb-client/src/remote/sql.rs +++ b/nodedb-client/src/remote/sql.rs @@ -10,34 +10,53 @@ //! pgwire driver; the trait method `execute_sql` borrows them as //! `&[&(dyn ToSql + Sync)]` for the actual `Client::query` call. +use std::collections::HashSet; + use nodedb_types::error::{NodeDbError, NodeDbResult}; use nodedb_types::filter::MetadataFilter; use nodedb_types::value::Value; use crate::remote_parse::format_vector_array; +use crate::row_decode::search_hit::{DISTANCE_COLUMN, ID_COLUMN}; +use crate::search_sql::key_filter::key_in_list; use crate::sql_escape::{quote_identifier, quote_string_literal}; /// Build the SQL for a vector search. /// /// Always emits the canonical -/// `SELECT * FROM [ WHERE ] ORDER BY vector_distance(ARRAY[...]) LIMIT ` -/// shape so the optional `WHERE` clause precedes `ORDER BY` (the SEARCH -/// preprocessor's "trailing append" form would have placed `WHERE` after -/// `ORDER BY`, which is invalid SQL). +/// `SELECT id, vector_distance(ARRAY[...]) AS distance FROM [ WHERE ] +/// ORDER BY vector_distance(ARRAY[...]) LIMIT ` shape. The optional +/// `WHERE` clause precedes `ORDER BY`. The projection names the two columns +/// a hit decodes from. +/// +/// `allowed_ids` pairs the allowed ids with the collection's identity +/// column. It adds a top-level ` IN (...)` conjunct. The server lowers +/// that key conjunct to the candidate bitmap the index search honors, so the +/// top-k is drawn from the allowed ids only. pub(super) fn build_vector_search_sql( collection: &str, query: &[f32], k: usize, filter: Option<&MetadataFilter>, + allowed_ids: Option<(&str, &HashSet)>, ) -> NodeDbResult { let collection = quote_identifier(collection); - let where_clause = match filter { - Some(f) => format!(" WHERE {}", render_metadata_filter(f)?), - None => String::new(), + let mut conjuncts: Vec = Vec::new(); + if let Some(f) = filter { + conjuncts.push(render_metadata_filter(f)?); + } + if let Some((key, ids)) = allowed_ids { + conjuncts.push(key_in_list(key, ids)); + } + let where_clause = if conjuncts.is_empty() { + String::new() + } else { + format!(" WHERE {}", conjuncts.join(" AND ")) }; + let query = format_vector_array(query); Ok(format!( - "SELECT * FROM {collection}{where_clause} ORDER BY vector_distance({}) LIMIT {k}", - format_vector_array(query), + "SELECT {ID_COLUMN}, vector_distance({query}) AS {DISTANCE_COLUMN} \ + FROM {collection}{where_clause} ORDER BY vector_distance({query}) LIMIT {k}", )) } @@ -180,17 +199,12 @@ mod tests { #[test] fn vector_search_sql_without_filter_renders_basic_form() { - // No-filter path: SELECT * FROM ORDER BY vector_distance(ARRAY[..]) LIMIT k. - let sql = - build_vector_search_sql("docs", &[0.1, 0.2, 0.3], 5, None).expect("no-filter is fine"); - assert!(sql.contains("SELECT")); - assert!(sql.contains("docs")); - assert!(sql.contains("vector_distance")); - assert!(sql.contains("ARRAY[0.1,0.2,0.3]")); - assert!(sql.contains("LIMIT 5")); - assert!( - !sql.contains(" WHERE "), - "no-filter SQL must not have WHERE; got: {sql}" + let sql = build_vector_search_sql("docs", &[0.1, 0.2, 0.3], 5, None, None) + .expect("no-filter is fine"); + assert_eq!( + sql, + "SELECT id, vector_distance(ARRAY[0.1,0.2,0.3]) AS distance FROM \"docs\" \ + ORDER BY vector_distance(ARRAY[0.1,0.2,0.3]) LIMIT 5" ); } @@ -199,7 +213,7 @@ mod tests { // Spec: a non-None Eq filter renders into a server-side predicate // referencing both the field and the value. let filter = MetadataFilter::eq("category", Value::String("ai".into())); - let sql = build_vector_search_sql("docs", &[0.1, 0.2], 5, Some(&filter)) + let sql = build_vector_search_sql("docs", &[0.1, 0.2], 5, Some(&filter), None) .expect("non-None metadata filter must be accepted client-side, not rejected"); assert!( sql.contains("category"), @@ -233,7 +247,7 @@ mod tests { value: Value::Float(0.5), }, ]); - let sql = build_vector_search_sql("docs", &[0.1], 3, Some(&filter)) + let sql = build_vector_search_sql("docs", &[0.1], 3, Some(&filter), None) .expect("compound metadata filter must be rendered, not rejected"); assert!( sql.contains("category"), @@ -255,12 +269,30 @@ mod tests { Value::String("databases".into()), ], }; - let sql = build_vector_search_sql("docs", &[0.0], 1, Some(&filter)).unwrap(); + let sql = build_vector_search_sql("docs", &[0.0], 1, Some(&filter), None).unwrap(); assert!(sql.contains(" IN (")); assert!(sql.contains("'rust'")); assert!(sql.contains("'databases'")); } + #[test] + fn vector_search_sql_restricts_to_allowed_ids() { + let ids: HashSet = ["b", "a'"].iter().map(|s| s.to_string()).collect(); + let sql = build_vector_search_sql("docs", &[0.0], 2, None, Some(("id", &ids))).unwrap(); + assert!( + sql.contains(" WHERE \"id\" IN ('a''', 'b') ORDER BY "), + "allowed ids render as a key conjunct; got: {sql}" + ); + let filter = MetadataFilter::eq("category", Value::String("ai".into())); + let sql = + build_vector_search_sql("docs", &[0.0], 2, Some(&filter), Some(("sku", &ids))).unwrap(); + assert!( + sql.contains(" WHERE \"category\" = 'ai' AND \"sku\" IN ('a''', 'b') ORDER BY "), + "the key conjunct names the declared key and joins the filter at the top level; \ + got: {sql}" + ); + } + #[test] fn render_sql_literal_escapes_single_quotes() { let sql = render_sql_literal(&Value::String("o'reilly".into())).unwrap(); diff --git a/nodedb-client/src/remote_parse.rs b/nodedb-client/src/remote_parse.rs index 0d270624c..d6c2b7dfc 100644 --- a/nodedb-client/src/remote_parse.rs +++ b/nodedb-client/src/remote_parse.rs @@ -2,88 +2,70 @@ //! Value conversion and SQL formatting helpers for the remote client. //! -//! Extracts pgwire row/column conversion, JSON-to-Value mapping, and -//! SQL identifier quoting into a focused module. Shared row decoders +//! Holds the pgwire row-cell entry point (the decoders live in +//! [`crate::pg_cell`]), JSON-to-Value mapping, and SQL array formatting. Shared row decoders //! (column-level `value_as_*`, `_system.dropped_collections` parsing) //! live in [`crate::row_decode`] instead so the trait default impls can //! reuse them without dragging in pgwire types. use nodedb_types::Value; +use nodedb_types::error::{NodeDbError, NodeDbResult}; -/// Convert a `tokio_postgres` column value to `nodedb_types::Value`. +use crate::pg_cell::{RawCell, decode_cell}; + +/// Convert cell `idx` of a `tokio_postgres` row, whose column is `column`, +/// to `nodedb_types::Value` by the column type. A cell that does not decode +/// is an error naming the column, the PG type and the cell text, never a +/// NULL. pub(crate) fn pg_value_to_value( row: &tokio_postgres::Row, idx: usize, - ty: &tokio_postgres::types::Type, -) -> Value { - use tokio_postgres::types::Type; - - match *ty { - Type::BOOL => row - .try_get::<_, bool>(idx) - .map(Value::Bool) - .unwrap_or(Value::Null), - Type::INT2 => row - .try_get::<_, i16>(idx) - .map(|v| Value::Integer(v as i64)) - .unwrap_or(Value::Null), - Type::INT4 => row - .try_get::<_, i32>(idx) - .map(|v| Value::Integer(v as i64)) - .unwrap_or(Value::Null), - Type::INT8 => row - .try_get::<_, i64>(idx) - .map(Value::Integer) - .unwrap_or(Value::Null), - Type::FLOAT4 => row - .try_get::<_, f32>(idx) - .map(|v| Value::Float(v as f64)) - .unwrap_or(Value::Null), - Type::FLOAT8 => row - .try_get::<_, f64>(idx) - .map(Value::Float) - .unwrap_or(Value::Null), - Type::TEXT | Type::VARCHAR | Type::NAME => row - .try_get::<_, String>(idx) - .map(Value::String) - .unwrap_or(Value::Null), - Type::BYTEA => row - .try_get::<_, Vec>(idx) - .map(Value::Bytes) - .unwrap_or(Value::Null), - Type::JSON | Type::JSONB => row - .try_get::<_, serde_json::Value>(idx) - .map(|v| json_to_value(&v)) - .unwrap_or(Value::Null), - _ => { - // Fallback: try as string. - row.try_get::<_, String>(idx) - .map(Value::String) - .unwrap_or(Value::Null) - } - } + column: &tokio_postgres::Column, +) -> NodeDbResult { + let raw = row.try_get::<_, Option>>(idx).map_err(|e| { + NodeDbError::serialization( + "pgwire", + format!( + "column \"{}\" of type {}: cannot read the cell: {e}", + column.name(), + column.type_().name() + ), + ) + })?; + decode_cell(column.name(), column.type_(), raw.map(|cell| cell.0)) } /// Convert `serde_json::Value` to `nodedb_types::Value`. -pub(crate) fn json_to_value(v: &serde_json::Value) -> Value { - match v { +/// +/// A JSON number keeps the kind the native client decodes it to: an integer +/// in `i64` range is an `Integer`, a larger unsigned integer is the `Decimal` +/// of [`Value::from_u64`], and every other number is a `Float`. A number +/// with no `f64` form is an error naming it, never a guessed value. +pub(crate) fn json_to_value(v: &serde_json::Value) -> NodeDbResult { + Ok(match v { serde_json::Value::Null => Value::Null, serde_json::Value::Bool(b) => Value::Bool(*b), - serde_json::Value::Number(n) => { - if let Some(i) = n.as_i64() { - Value::Integer(i) - } else { - Value::Float(n.as_f64().unwrap_or(0.0)) + serde_json::Value::Number(n) => match (n.as_i64(), n.as_u64(), n.as_f64()) { + (Some(i), _, _) => Value::Integer(i), + (None, Some(u), _) => Value::from_u64(u), + (None, None, Some(f)) => Value::Float(f), + (None, None, None) => { + return Err(NodeDbError::serialization( + "json", + format!("number {n} has no integer or f64 form"), + )); } - } + }, serde_json::Value::String(s) => Value::String(s.clone()), - serde_json::Value::Array(a) => Value::Array(a.iter().map(json_to_value).collect()), + serde_json::Value::Array(a) => { + Value::Array(a.iter().map(json_to_value).collect::>()?) + } serde_json::Value::Object(m) => Value::Object( m.iter() - .map(|(k, v)| (k.clone(), json_to_value(v))) - .collect(), + .map(|(k, v)| Ok((k.clone(), json_to_value(v)?))) + .collect::>()?, ), - } + }) } /// Format an f32 slice as a SQL ARRAY literal: `ARRAY[0.1,0.2,0.3]`. @@ -108,21 +90,33 @@ mod tests { assert_eq!(arr, "ARRAY[]"); } + fn decoded(v: serde_json::Value) -> Value { + json_to_value(&v).expect("every serde_json number has an f64 form") + } + #[test] fn json_to_value_primitives() { - assert_eq!(json_to_value(&serde_json::json!(null)), Value::Null); - assert_eq!(json_to_value(&serde_json::json!(true)), Value::Bool(true)); - assert_eq!(json_to_value(&serde_json::json!(42)), Value::Integer(42)); - assert_eq!(json_to_value(&serde_json::json!(2.5)), Value::Float(2.5)); + assert_eq!(decoded(serde_json::json!(null)), Value::Null); + assert_eq!(decoded(serde_json::json!(true)), Value::Bool(true)); + assert_eq!(decoded(serde_json::json!(42)), Value::Integer(42)); + assert_eq!(decoded(serde_json::json!(2.5)), Value::Float(2.5)); assert_eq!( - json_to_value(&serde_json::json!("hello")), + decoded(serde_json::json!("hello")), Value::String("hello".into()) ); } + #[test] + fn json_to_value_keeps_a_u64_above_i64_max_exact() { + assert_eq!( + decoded(serde_json::json!(u64::MAX)), + Value::from_u64(u64::MAX) + ); + } + #[test] fn json_to_value_nested() { - let v = json_to_value(&serde_json::json!({"a": [1, 2]})); + let v = decoded(serde_json::json!({"a": [1, 2]})); assert!(matches!(v, Value::Object(_))); } } diff --git a/nodedb-client/src/row_decode/mod.rs b/nodedb-client/src/row_decode/mod.rs index 1cb5f6fca..282d1f17c 100644 --- a/nodedb-client/src/row_decode/mod.rs +++ b/nodedb-client/src/row_decode/mod.rs @@ -9,6 +9,8 @@ //! system-catalog column layout changes. pub(crate) mod dropped_collection; +pub(crate) mod search_hit; pub(crate) mod value; pub(crate) use dropped_collection::parse_dropped_collection_rows; +pub(crate) use search_hit::{HitSource, decode_search_hits}; diff --git a/nodedb-client/src/row_decode/search_hit.rs b/nodedb-client/src/row_decode/search_hit.rs new file mode 100644 index 000000000..c14c5e0b1 --- /dev/null +++ b/nodedb-client/src/row_decode/search_hit.rs @@ -0,0 +1,196 @@ +// SPDX-License-Identifier: Apache-2.0 + +//! Decoder for search-hit rows: one `SearchResult` per row, read by column +//! name. +//! +//! Vector and text search return the hit id in the `id` column and the +//! rank value in a named score column (`distance` for vector search, +//! `score` for text search). Other columns (document fields, `_surrogate`) +//! are ignored. Both clients decode through here: the native protocol +//! carries typed cells, while pgwire can carry every cell as text. +//! +//! A row set with rows but without the id or score column is an error, and +//! so is a row whose id or score does not decode. A hit dropped or defaulted +//! here reads, to the caller, as "the search matched less". + +use std::collections::HashMap; + +use nodedb_types::error::{NodeDbError, NodeDbResult}; +use nodedb_types::result::SearchResult; +use nodedb_types::value::Value; + +/// Output column holding the hit id. +pub(crate) const ID_COLUMN: &str = "id"; +/// Output column holding a vector hit's distance. +pub(crate) const DISTANCE_COLUMN: &str = "distance"; + +/// What a hit row set answers, named in decode errors. +pub(crate) struct HitSource<'a> { + /// The client operation, e.g. `vector_search`. + pub op: &'a str, + pub collection: &'a str, + /// The column holding the rank value. + pub score_column: &'a str, +} + +/// Decode every row of a search result into a hit, in row order. +/// +/// An empty row set is no hits whatever its columns: pgwire reports no +/// column names for a statement that returned no rows. +pub(crate) fn decode_search_hits( + source: &HitSource<'_>, + columns: &[String], + rows: &[Vec], +) -> NodeDbResult> { + if rows.is_empty() { + return Ok(Vec::new()); + } + let id_idx = column_index(source, columns, ID_COLUMN)?; + let score_idx = column_index(source, columns, source.score_column)?; + rows.iter() + .enumerate() + .map(|(row_idx, row)| { + let id = hit_id(source, row_idx, row.get(id_idx))?; + let distance = hit_score(source, row_idx, &id, row.get(score_idx))?; + Ok(SearchResult { + id, + node_id: None, + distance, + metadata: HashMap::new(), + }) + }) + .collect() +} + +fn column_index(source: &HitSource<'_>, columns: &[String], name: &str) -> NodeDbResult { + columns.iter().position(|c| c == name).ok_or_else(|| { + decode_error( + source, + format!("result has no '{name}' column; columns are {columns:?}"), + ) + }) +} + +/// A text id, or the integer identity the server reports for a hit bound to +/// no primary key. +fn hit_id(source: &HitSource<'_>, row_idx: usize, cell: Option<&Value>) -> NodeDbResult { + match cell { + Some(Value::String(id)) => Ok(id.clone()), + Some(Value::Integer(id)) => Ok(id.to_string()), + other => Err(decode_error( + source, + format!("row {row_idx} has id {other:?}, expected text or an integer"), + )), + } +} + +/// A number, or the decimal text pgwire carries for one. +/// +/// Text parses as the `f64` the server rendered, then narrows to `f32` like +/// a typed `Float`. Both transports narrow the same `f64`, so a score is +/// identical whichever client read it. +fn hit_score( + source: &HitSource<'_>, + row_idx: usize, + id: &str, + cell: Option<&Value>, +) -> NodeDbResult { + match cell { + Some(Value::Float(score)) => Ok(*score as f32), + Some(Value::Integer(score)) => Ok(*score as f32), + Some(Value::String(text)) => text.parse::().map(|score| score as f32).map_err(|e| { + decode_error( + source, + format!( + "row {row_idx} ('{id}') {} '{text}' is not a number: {e}", + source.score_column + ), + ) + }), + other => Err(decode_error( + source, + format!( + "row {row_idx} ('{id}') has {} {other:?}, expected a number", + source.score_column + ), + )), + } +} + +fn decode_error(source: &HitSource<'_>, detail: String) -> NodeDbError { + NodeDbError::serialization( + "search hit", + format!("{} '{}': {detail}", source.op, source.collection), + ) +} + +#[cfg(test)] +mod tests { + use super::*; + + const SOURCE: HitSource<'static> = HitSource { + op: "vector_search", + collection: "c", + score_column: "distance", + }; + + fn names(columns: &[&str]) -> Vec { + columns.iter().map(|c| c.to_string()).collect() + } + + #[test] + fn hits_decode_by_column_name_in_row_order() { + let columns = names(&["embedding", "distance", "id"]); + let rows = vec![ + vec![Value::Null, Value::Float(0.5), Value::String("a".into())], + vec![Value::Null, Value::Integer(2), Value::String("b".into())], + ]; + let hits = decode_search_hits(&SOURCE, &columns, &rows).expect("decode"); + let got: Vec<(&str, f32)> = hits.iter().map(|h| (h.id.as_str(), h.distance)).collect(); + assert_eq!(got, vec![("a", 0.5), ("b", 2.0)]); + } + + #[test] + fn pgwire_text_cells_decode() { + let rows = vec![vec![ + Value::String("a".into()), + Value::String("0.25".into()), + ]]; + let hits = decode_search_hits(&SOURCE, &names(&["id", "distance"]), &rows).expect("decode"); + assert_eq!(hits[0].distance, 0.25); + } + + #[test] + fn an_integer_identity_decodes_as_its_text() { + let rows = vec![vec![Value::Integer(7), Value::Float(1.0)]]; + let hits = decode_search_hits(&SOURCE, &names(&["id", "distance"]), &rows).expect("decode"); + assert_eq!(hits[0].id, "7"); + } + + #[test] + fn no_rows_is_no_hits_without_columns() { + assert!( + decode_search_hits(&SOURCE, &[], &[]) + .expect("empty") + .is_empty() + ); + } + + #[test] + fn a_missing_column_is_an_error_naming_it() { + let rows = vec![vec![Value::String("a".into())]]; + let err = decode_search_hits(&SOURCE, &names(&["id"]), &rows).expect_err("no distance"); + assert!(err.to_string().contains("'distance'"), "{err}"); + } + + #[test] + fn a_null_id_or_bad_score_is_an_error() { + let columns = names(&["id", "distance"]); + let null_id = vec![vec![Value::Null, Value::Float(1.0)]]; + assert!(decode_search_hits(&SOURCE, &columns, &null_id).is_err()); + let bad_score = vec![vec![Value::String("a".into()), Value::String("x".into())]]; + assert!(decode_search_hits(&SOURCE, &columns, &bad_score).is_err()); + let null_score = vec![vec![Value::String("a".into()), Value::Null]]; + assert!(decode_search_hits(&SOURCE, &columns, &null_score).is_err()); + } +} diff --git a/nodedb-client/src/row_decode/value.rs b/nodedb-client/src/row_decode/value.rs index 8d03e6373..efd28752c 100644 --- a/nodedb-client/src/row_decode/value.rs +++ b/nodedb-client/src/row_decode/value.rs @@ -13,12 +13,16 @@ use nodedb_types::error::{NodeDbError, NodeDbResult}; use nodedb_types::value::Value; -/// Decode a row value as `u64`. Accepts both `Value::Integer` (extended- -/// query / native path) and `Value::String` containing a base-10 integer -/// (pgwire simple-query path, which returns every column as text). +/// Decode a row value as `u64`. Accepted shapes: +/// - `Value::Integer`: extended-query or native path, up to `i64::MAX`. +/// - `Value::Decimal` with scale 0: a msgpack `uint64` above `i64::MAX`. +/// - `Value::String` holding a base-10 integer: pgwire simple-query text. pub(crate) fn value_as_u64(v: &Value) -> NodeDbResult { match v { - Value::Integer(i) => Ok(*i as u64), + Value::Integer(i) => u64::try_from(*i) + .map_err(|_| NodeDbError::storage(format!("expected u64 column, got negative {i}"))), + Value::Decimal(d) => Value::decimal_as_wide_u64(d) + .ok_or_else(|| NodeDbError::storage(format!("expected u64 column, got decimal {d}"))), Value::String(s) => s .parse::() .map_err(|e| NodeDbError::storage(format!("parse u64 from '{s}': {e}"))), @@ -50,6 +54,26 @@ mod tests { assert_eq!(value_as_u64(&Value::Integer(42)).unwrap(), 42u64); } + #[test] + fn value_as_u64_rejects_a_negative_integer() { + let err = value_as_u64(&Value::Integer(-1)).unwrap_err(); + assert!(err.to_string().contains("negative -1"), "{err}"); + } + + #[test] + fn value_as_u64_accepts_a_wide_uint64_decimal() { + let wide = Value::from_u64(u64::MAX); + assert!(matches!(wide, Value::Decimal(_))); + assert_eq!(value_as_u64(&wide).unwrap(), u64::MAX); + } + + #[test] + fn value_as_u64_rejects_a_scaled_decimal() { + let scaled = Value::Decimal(rust_decimal::Decimal::new(15, 1)); + let err = value_as_u64(&scaled).unwrap_err(); + assert!(err.to_string().contains("decimal 1.5"), "{err}"); + } + #[test] fn value_as_u64_accepts_numeric_string() { // Simple-query path returns integer columns as text — must parse. diff --git a/nodedb-client/src/search_sql/key_filter.rs b/nodedb-client/src/search_sql/key_filter.rs new file mode 100644 index 000000000..5c720f980 --- /dev/null +++ b/nodedb-client/src/search_sql/key_filter.rs @@ -0,0 +1,34 @@ +// SPDX-License-Identifier: Apache-2.0 + +//! The ` IN (...)` conjunct that restricts a search to allowed ids. + +use std::collections::HashSet; + +use crate::sql_escape::{quote_identifier, quote_string_literal}; + +/// ` IN ('a', 'b', ...)` over `ids` in sorted order. +/// +/// `key` is the collection's identity column: the declared primary key, +/// else `id`. The server lowers a top-level conjunct on that column to the +/// candidate set a search ranks within. A conjunct on any other column +/// filters after the ranking cut instead. +pub(crate) fn key_in_list(key: &str, ids: &HashSet) -> String { + let mut ids: Vec<&String> = ids.iter().collect(); + ids.sort(); + let list: Vec = ids + .iter() + .map(|id| quote_string_literal(id.as_str())) + .collect(); + format!("{} IN ({})", quote_identifier(key), list.join(", ")) +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn ids_render_sorted_and_quoted_under_the_key_column() { + let ids: HashSet = ["b", "a'"].iter().map(|s| s.to_string()).collect(); + assert_eq!(key_in_list("sku", &ids), "\"sku\" IN ('a''', 'b')"); + } +} diff --git a/nodedb-client/src/search_sql/mod.rs b/nodedb-client/src/search_sql/mod.rs new file mode 100644 index 000000000..1815690c1 --- /dev/null +++ b/nodedb-client/src/search_sql/mod.rs @@ -0,0 +1,6 @@ +// SPDX-License-Identifier: Apache-2.0 + +pub(crate) mod key_filter; +pub(crate) mod text_search; + +pub(crate) use text_search::{TextSearchRequest, text_hit_source, text_search_sql}; diff --git a/nodedb-client/src/search_sql/text_search.rs b/nodedb-client/src/search_sql/text_search.rs new file mode 100644 index 000000000..eb3c55026 --- /dev/null +++ b/nodedb-client/src/search_sql/text_search.rs @@ -0,0 +1,187 @@ +// SPDX-License-Identifier: Apache-2.0 + +//! The full-text search statement both clients send. +//! +//! The search is one statement: +//! `SELECT id, bm25_score(f, q[, opts]) AS score FROM c +//! WHERE text_match(f, q[, opts]) [AND IN (...)] ORDER BY score DESC LIMIT k`. +//! +//! `opts` are the named options `mode => 'and'` and `fuzzy => true` for each +//! `params` value that differs from [`TextSearchParams::default`]. The server +//! runs that default for an omitted option. A phrase query (`"..."`) takes no +//! option on the server, so a default-params phrase search renders none. + +use std::collections::HashSet; + +use nodedb_types::text_search::TextSearchParams; + +use crate::row_decode::HitSource; +use crate::row_decode::search_hit::ID_COLUMN; +use crate::sql_escape::{quote_identifier, quote_string_literal}; + +use super::key_filter::key_in_list; + +/// Output column holding the BM25 score. +const SCORE_COLUMN: &str = "score"; + +/// The hit rows of a text search of `collection`. +pub(crate) fn text_hit_source(collection: &str) -> HitSource<'_> { + HitSource { + op: "text_search", + collection, + score_column: SCORE_COLUMN, + } +} + +/// One text search request. +pub(crate) struct TextSearchRequest<'a> { + pub collection: &'a str, + /// The indexed field. Empty searches the whole-document index. + pub field: &'a str, + pub query: &'a str, + pub top_k: usize, + pub params: &'a TextSearchParams, + /// The collection's identity column and the ids the search may return. + pub allowed: Option<(&'a str, &'a HashSet)>, +} + +/// The search statement. Every user string is a quoted literal or identifier. +pub(crate) fn text_search_sql(request: &TextSearchRequest<'_>) -> String { + let target = text_search_target(request.field); + let q = quote_string_literal(request.query); + let opts = text_search_options(request.params); + let allowed = match request.allowed { + Some((key, ids)) => format!(" AND {}", key_in_list(key, ids)), + None => String::new(), + }; + format!( + "SELECT {ID_COLUMN}, bm25_score({target}, {q}{opts}) AS {SCORE_COLUMN} \ + FROM {} WHERE text_match({target}, {q}{opts}){allowed} \ + ORDER BY {SCORE_COLUMN} DESC LIMIT {}", + quote_identifier(request.collection), + request.top_k + ) +} + +/// The named options of `params` that differ from the default, each led by +/// `, `. Empty when `params` is the default. +fn text_search_options(params: &TextSearchParams) -> String { + let defaults = TextSearchParams::default(); + let mut options = String::new(); + if params.mode != defaults.mode { + options.push_str(&format!(", mode => '{}'", params.mode.as_str())); + } + if params.fuzzy != defaults.fuzzy { + options.push_str(&format!(", fuzzy => {}", params.fuzzy)); + } + options +} + +/// The first argument of `text_match` / `bm25_score`: the quoted column, or +/// `*` (the whole-document index) when `field` is empty. +fn text_search_target(field: &str) -> String { + if field.is_empty() { + "*".to_string() + } else { + quote_identifier(field) + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::row_decode::decode_search_hits; + use nodedb_types::text_search::QueryMode; + use nodedb_types::value::Value; + + fn request<'a>( + query: &'a str, + top_k: usize, + params: &'a TextSearchParams, + allowed: Option<(&'a str, &'a HashSet)>, + ) -> TextSearchRequest<'a> { + TextSearchRequest { + collection: "docs", + field: "body", + query, + top_k, + params, + allowed, + } + } + + #[test] + fn an_empty_field_searches_the_whole_document() { + assert_eq!(text_search_target(""), "*"); + assert_eq!(text_search_target("body"), "\"body\""); + assert_eq!(text_search_target("we\"ird"), "\"we\"\"ird\""); + } + + #[test] + fn the_statement_ranks_best_first_before_the_limit() { + let params = TextSearchParams::default(); + let sql = text_search_sql(&request("it's", 5, ¶ms, None)); + assert_eq!( + sql, + "SELECT id, bm25_score(\"body\", 'it''s') AS score FROM \"docs\" \ + WHERE text_match(\"body\", 'it''s') ORDER BY score DESC LIMIT 5" + ); + } + + #[test] + fn allowed_ids_restrict_the_candidates_on_the_key_column() { + let ids: HashSet = ["b", "a'"].iter().map(|s| s.to_string()).collect(); + let params = TextSearchParams::default(); + let sql = text_search_sql(&request("q", 3, ¶ms, Some(("sku", &ids)))); + assert!( + sql.contains("WHERE text_match(\"body\", 'q') AND \"sku\" IN ('a''', 'b') ORDER BY"), + "{sql}" + ); + } + + #[test] + fn default_params_render_no_option() { + assert_eq!(text_search_options(&TextSearchParams::default()), ""); + } + + #[test] + fn non_default_params_render_on_both_calls() { + let params = TextSearchParams { + mode: QueryMode::And, + fuzzy: true, + }; + assert_eq!( + text_search_options(¶ms), + ", mode => 'and', fuzzy => true" + ); + let sql = text_search_sql(&request("q", 3, ¶ms, None)); + assert!( + sql.contains("bm25_score(\"body\", 'q', mode => 'and', fuzzy => true) AS score"), + "{sql}" + ); + assert!( + sql.contains("WHERE text_match(\"body\", 'q', mode => 'and', fuzzy => true) ORDER"), + "{sql}" + ); + let fuzzy_only = TextSearchParams { + mode: QueryMode::Or, + fuzzy: true, + }; + assert_eq!(text_search_options(&fuzzy_only), ", fuzzy => true"); + } + + #[test] + fn hits_decode_from_the_score_column() { + let rows = vec![vec![ + Value::String("0.5".into()), + Value::String("d1".into()), + ]]; + let names = vec!["score".to_string(), "id".to_string()]; + let hits = decode_search_hits(&text_hit_source("c"), &names, &rows).expect("decode"); + assert_eq!(hits[0].id, "d1"); + assert_eq!(hits[0].distance, 0.5); + let err = decode_search_hits(&text_hit_source("c"), &names[1..], &rows) + .expect_err("no score column"); + assert!(err.to_string().contains("'score'"), "{err}"); + } +} diff --git a/nodedb-client/src/traits/core/default_impls.rs b/nodedb-client/src/traits/core/default_impls.rs deleted file mode 100644 index 3c3c7fb04..000000000 --- a/nodedb-client/src/traits/core/default_impls.rs +++ /dev/null @@ -1,245 +0,0 @@ -// SPDX-License-Identifier: Apache-2.0 - -//! Default-provided method bodies for [`crate::traits::core::NodeDb`]. -//! -//! Factored out of the trait declaration in `trait_def.rs` so that file -//! stays signatures + docs; each function here backs exactly one -//! default-provided `NodeDb` method, and the trait method delegates to it -//! in one line. Pure relocation — no behavior change. - -use std::collections::{HashMap, HashSet}; - -use nodedb_types::document::Document; -use nodedb_types::dropped_collection::DroppedCollection; -use nodedb_types::error::{NodeDbError, NodeDbResult}; -use nodedb_types::filter::{EdgeFilter, MetadataFilter}; -use nodedb_types::id::NodeId; -use nodedb_types::result::SearchResult; -use nodedb_types::text_search::TextSearchParams; - -use super::quote::quote_ident; -use super::trait_def::NodeDb; - -pub(super) fn graph_pagerank_default( - collection: &str, - personalization: Option>, - damping: Option, - max_iterations: Option, -) -> NodeDbResult> { - let _ = (collection, personalization, damping, max_iterations); - Err(NodeDbError::storage( - "graph_pagerank is not implemented for this NodeDb backend", - )) -} - -pub(super) async fn document_put_with_vector_default( - this: &T, - doc_collection: &str, - doc: Document, - vector_collection: &str, - id: &str, - embedding: &[f32], -) -> NodeDbResult<()> { - this.document_put(doc_collection, doc).await?; - if !embedding.is_empty() { - this.vector_insert(vector_collection, id, embedding, None) - .await?; - } - Ok(()) -} - -pub(super) fn document_get_as_of_default( - collection: &str, - id: &str, - as_of_ms: Option, - valid_time_ms: Option, -) -> NodeDbResult> { - let _ = (collection, id, as_of_ms, valid_time_ms); - Err(NodeDbError::storage( - "document_get_as_of is not implemented on this client", - )) -} - -pub(super) fn document_put_with_valid_time_default( - collection: &str, - doc: Document, - valid_from_ms: Option, - valid_until_ms: Option, -) -> NodeDbResult<()> { - let _ = (collection, doc, valid_from_ms, valid_until_ms); - Err(NodeDbError::storage( - "document_put_with_valid_time is not implemented on this client", - )) -} - -pub(super) fn vector_insert_field_default( - collection: &str, - field_name: &str, - id: &str, - embedding: &[f32], - metadata: Option, -) -> NodeDbResult<()> { - let _ = (collection, id, embedding, metadata); - Err(NodeDbError::storage(format!( - "vector_insert_field is not implemented on this client; \ - field_name={field_name} would have been silently dropped" - ))) -} - -pub(super) fn vector_search_field_default( - collection: &str, - field_name: &str, - query: &[f32], - k: usize, - filter: Option<&MetadataFilter>, -) -> NodeDbResult> { - let _ = (collection, query, k, filter); - Err(NodeDbError::storage(format!( - "vector_search_field is not implemented on this client; \ - field_name={field_name} would have been silently dropped" - ))) -} - -/// Default forward-BFS shortest-path search built on `graph_traverse`. -/// -/// See [`crate::traits::core::NodeDb::graph_shortest_path`] for the -/// full contract. -pub(super) async fn graph_shortest_path_default( - this: &T, - collection: &str, - from: &NodeId, - to: &NodeId, - max_depth: u8, - edge_filter: Option<&EdgeFilter>, -) -> NodeDbResult>> { - if from == to { - return Ok(Some(vec![from.clone()])); - } - if max_depth == 0 { - return Ok(None); - } - - // Map of `node -> parent` used to reconstruct the path once the - // target is reached. The source has no parent entry. - let mut parent: HashMap = HashMap::new(); - let mut frontier: Vec = vec![from.clone()]; - - for _ in 0..max_depth { - let mut next_frontier: Vec = Vec::new(); - for node in &frontier { - let sg = this - .graph_traverse(collection, node, 1, edge_filter) - .await?; - for edge in &sg.edges { - // Only follow edges originating from the current - // node — `graph_traverse` may include adjacent - // edges that don't extend the BFS frontier. - if &edge.from != node { - continue; - } - let dst = &edge.to; - if dst == from || parent.contains_key(dst) { - continue; - } - parent.insert(dst.clone(), node.clone()); - if dst == to { - let mut path = vec![to.clone()]; - let mut cur = to.clone(); - while &cur != from { - let p = parent - .get(&cur) - .expect("BFS reached `to` so all ancestors are tracked") - .clone(); - path.push(p.clone()); - cur = p; - } - path.reverse(); - return Ok(Some(path)); - } - next_frontier.push(dst.clone()); - } - } - if next_frontier.is_empty() { - return Ok(None); - } - frontier = next_frontier; - } - Ok(None) -} - -pub(super) fn text_search_default( - collection: &str, - field: &str, - query: &str, - top_k: usize, - params: TextSearchParams, - allowed_ids: Option<&HashSet>, -) -> NodeDbResult> { - let _ = (collection, field, query, top_k, params, allowed_ids); - Err(NodeDbError::storage( - "text_search is not implemented on this client", - )) -} - -pub(super) async fn batch_vector_insert_default( - this: &T, - collection: &str, - vectors: &[(&str, &[f32])], -) -> NodeDbResult<()> { - for &(id, embedding) in vectors { - this.vector_insert(collection, id, embedding, None).await?; - } - Ok(()) -} - -pub(super) async fn batch_graph_insert_edges_default( - this: &T, - collection: &str, - edges: &[(&str, &str, &str)], -) -> NodeDbResult<()> { - for &(from, to, label) in edges { - let src = NodeId::try_new(from) - .map_err(|e| NodeDbError::storage(format!("invalid node id: {e}")))?; - let dst = NodeId::try_new(to) - .map_err(|e| NodeDbError::storage(format!("invalid node id: {e}")))?; - this.graph_insert_edge(collection, &src, &dst, label, None) - .await?; - } - Ok(()) -} - -pub(super) async fn undrop_collection_default( - this: &T, - name: &str, -) -> NodeDbResult<()> { - let sql = format!("UNDROP COLLECTION {}", quote_ident(name)); - this.execute_sql(&sql, &[]).await?; - Ok(()) -} - -pub(super) async fn drop_collection_purge_default( - this: &T, - name: &str, -) -> NodeDbResult<()> { - let sql = format!("DROP COLLECTION {} PURGE", quote_ident(name)); - this.execute_sql(&sql, &[]).await?; - Ok(()) -} - -pub(super) async fn list_dropped_collections_default( - this: &T, -) -> NodeDbResult> { - let sql = "SELECT tenant_id, name, owner, engine_type, \ - deactivated_at_ns, retention_expires_at_ns \ - FROM _system.dropped_collections"; - let result = this.execute_sql(sql, &[]).await?; - crate::row_decode::parse_dropped_collection_rows(&result.rows) -} - -pub(super) fn on_collection_purged_default() -> NodeDbResult<()> { - Err(NodeDbError::storage( - "on_collection_purged is not supported on this client — \ - requires a push-capable sync connection (NodeDbLite or a \ - sync-enabled remote client)", - )) -} diff --git a/nodedb-client/src/traits/core/default_impls/batch.rs b/nodedb-client/src/traits/core/default_impls/batch.rs new file mode 100644 index 000000000..718fa8a20 --- /dev/null +++ b/nodedb-client/src/traits/core/default_impls/batch.rs @@ -0,0 +1,33 @@ +// SPDX-License-Identifier: Apache-2.0 + +use nodedb_types::error::{NodeDbError, NodeDbResult}; +use nodedb_types::id::NodeId; + +use super::super::trait_def::NodeDb; + +pub(in crate::traits::core) async fn batch_vector_insert_default( + this: &T, + collection: &str, + vectors: &[(&str, &[f32])], +) -> NodeDbResult<()> { + for &(id, embedding) in vectors { + this.vector_insert(collection, id, embedding, None).await?; + } + Ok(()) +} + +pub(in crate::traits::core) async fn batch_graph_insert_edges_default( + this: &T, + collection: &str, + edges: &[(&str, &str, &str)], +) -> NodeDbResult<()> { + for &(from, to, label) in edges { + let src = NodeId::try_new(from) + .map_err(|e| NodeDbError::storage(format!("invalid node id: {e}")))?; + let dst = NodeId::try_new(to) + .map_err(|e| NodeDbError::storage(format!("invalid node id: {e}")))?; + this.graph_insert_edge(collection, &src, &dst, label, None) + .await?; + } + Ok(()) +} diff --git a/nodedb-client/src/traits/core/default_impls/collections.rs b/nodedb-client/src/traits/core/default_impls/collections.rs new file mode 100644 index 000000000..a7ece583b --- /dev/null +++ b/nodedb-client/src/traits/core/default_impls/collections.rs @@ -0,0 +1,43 @@ +// SPDX-License-Identifier: Apache-2.0 + +use super::super::quote::quote_ident; +use nodedb_types::dropped_collection::DroppedCollection; +use nodedb_types::error::{NodeDbError, NodeDbResult}; + +use super::super::trait_def::NodeDb; + +pub(in crate::traits::core) async fn undrop_collection_default( + this: &T, + name: &str, +) -> NodeDbResult<()> { + let sql = format!("UNDROP COLLECTION {}", quote_ident(name)); + this.execute_sql(&sql, &[]).await?; + Ok(()) +} + +pub(in crate::traits::core) async fn drop_collection_purge_default( + this: &T, + name: &str, +) -> NodeDbResult<()> { + let sql = format!("DROP COLLECTION {} PURGE", quote_ident(name)); + this.execute_sql(&sql, &[]).await?; + Ok(()) +} + +pub(in crate::traits::core) async fn list_dropped_collections_default( + this: &T, +) -> NodeDbResult> { + let sql = "SELECT tenant_id, name, owner, engine_type, \ + deactivated_at_ns, retention_expires_at_ns \ + FROM _system.dropped_collections"; + let result = this.execute_sql(sql, &[]).await?; + crate::row_decode::parse_dropped_collection_rows(&result.rows) +} + +pub(in crate::traits::core) fn on_collection_purged_default() -> NodeDbResult<()> { + Err(NodeDbError::storage( + "on_collection_purged is not supported on this client — \ + requires a push-capable sync connection (NodeDbLite or a \ + sync-enabled remote client)", + )) +} diff --git a/nodedb-client/src/traits/core/default_impls/dispatch.rs b/nodedb-client/src/traits/core/default_impls/dispatch.rs new file mode 100644 index 000000000..a0646180a --- /dev/null +++ b/nodedb-client/src/traits/core/default_impls/dispatch.rs @@ -0,0 +1,417 @@ +// SPDX-License-Identifier: Apache-2.0 + +use nodedb_types::document::Document; +use nodedb_types::error::{NodeDbError, NodeDbResult}; +use nodedb_types::filter::MetadataFilter; +use nodedb_types::result::SearchResult; + +use super::super::trait_def::NodeDb; + +pub(in crate::traits::core) async fn document_put_with_vector_default( + this: &T, + doc_collection: &str, + doc: Document, + vector_collection: &str, + id: &str, + embedding: &[f32], +) -> NodeDbResult<()> { + this.document_put(doc_collection, doc).await?; + if !embedding.is_empty() { + this.vector_insert(vector_collection, id, embedding, None) + .await?; + } + Ok(()) +} + +pub(in crate::traits::core) fn document_get_as_of_default( + collection: &str, + id: &str, + as_of_ms: Option, + valid_time_ms: Option, +) -> NodeDbResult> { + let _ = (collection, id, as_of_ms, valid_time_ms); + Err(NodeDbError::storage( + "document_get_as_of is not implemented on this client", + )) +} + +pub(in crate::traits::core) fn document_put_with_valid_time_default( + collection: &str, + doc: Document, + valid_from_ms: Option, + valid_until_ms: Option, +) -> NodeDbResult<()> { + let _ = (collection, doc, valid_from_ms, valid_until_ms); + Err(NodeDbError::storage( + "document_put_with_valid_time is not implemented on this client", + )) +} + +pub(in crate::traits::core) fn vector_insert_field_default( + collection: &str, + field_name: &str, + id: &str, + embedding: &[f32], + metadata: Option, +) -> NodeDbResult<()> { + let _ = (collection, id, embedding, metadata); + Err(NodeDbError::storage(format!( + "vector_insert_field is not implemented on this client; \ + field_name={field_name} would have been silently dropped" + ))) +} + +pub(in crate::traits::core) fn vector_search_field_default( + collection: &str, + field_name: &str, + query: &[f32], + k: usize, + filter: Option<&MetadataFilter>, +) -> NodeDbResult> { + let _ = (collection, query, k, filter); + Err(NodeDbError::storage(format!( + "vector_search_field is not implemented on this client; \ + field_name={field_name} would have been silently dropped" + ))) +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::capabilities::Capabilities; + use async_trait::async_trait; + use nodedb_types::document::Document; + use nodedb_types::error::{NodeDbError, NodeDbResult}; + use nodedb_types::filter::{EdgeFilter, MetadataFilter}; + use nodedb_types::graph::GraphStats; + use nodedb_types::id::{EdgeId, NodeId}; + use nodedb_types::result::{QueryResult, SearchResult, SubGraph}; + use nodedb_types::text_search::TextSearchParams; + use nodedb_types::value::Value; + use std::collections::{HashMap, HashSet}; + + /// Mock implementation to verify the trait is object-safe and + /// can be used as `Arc`. + struct MockDb; + + #[cfg_attr(not(target_arch = "wasm32"), async_trait)] + #[cfg_attr(target_arch = "wasm32", async_trait(?Send))] + impl NodeDb for MockDb { + async fn vector_search( + &self, + _collection: &str, + _query: &[f32], + _k: usize, + _filter: Option<&MetadataFilter>, + _allowed_ids: Option<&HashSet>, + ) -> NodeDbResult> { + Ok(vec![SearchResult { + id: "vec-1".into(), + node_id: None, + distance: 0.1, + metadata: HashMap::new(), + }]) + } + + async fn vector_insert( + &self, + _collection: &str, + _id: &str, + _embedding: &[f32], + _metadata: Option, + ) -> NodeDbResult<()> { + Ok(()) + } + + async fn vector_delete(&self, _collection: &str, _id: &str) -> NodeDbResult<()> { + Ok(()) + } + + async fn graph_traverse( + &self, + _collection: &str, + _start: &NodeId, + _depth: u8, + _direction: nodedb_types::graph::Direction, + _edge_filter: Option<&EdgeFilter>, + ) -> NodeDbResult { + Ok(SubGraph::empty()) + } + + async fn graph_insert_edge( + &self, + _collection: &str, + from: &NodeId, + to: &NodeId, + edge_type: &str, + _properties: Option, + ) -> NodeDbResult { + EdgeId::try_first(from.clone(), to.clone(), edge_type) + .map_err(|e| NodeDbError::storage(format!("invalid edge label: {e}"))) + } + + async fn graph_delete_edge( + &self, + _collection: &str, + _edge_id: &EdgeId, + ) -> NodeDbResult<()> { + Ok(()) + } + + async fn graph_stats( + &self, + collection: Option<&str>, + _as_of: Option, + ) -> NodeDbResult> { + Ok(vec![GraphStats::zero(collection.unwrap_or("mock"))]) + } + + async fn document_get( + &self, + _collection: &str, + id: &str, + ) -> NodeDbResult> { + let mut doc = Document::new(id); + doc.set("title", Value::String("test".into())); + Ok(Some(doc)) + } + + async fn document_put(&self, _collection: &str, _doc: Document) -> NodeDbResult<()> { + Ok(()) + } + + async fn document_delete(&self, _collection: &str, _id: &str) -> NodeDbResult<()> { + Ok(()) + } + + async fn list_insert( + &self, + _collection: &str, + _document_id: &str, + _list_path: &str, + _index: usize, + _fields: &Value, + ) -> NodeDbResult<()> { + Ok(()) + } + + async fn list_delete( + &self, + _collection: &str, + _document_id: &str, + _list_path: &str, + _index: usize, + ) -> NodeDbResult<()> { + Ok(()) + } + + async fn list_move( + &self, + _collection: &str, + _document_id: &str, + _list_path: &str, + _from_index: usize, + _to_index: usize, + ) -> NodeDbResult<()> { + Ok(()) + } + + async fn text_search( + &self, + _collection: &str, + _field: &str, + _query: &str, + _top_k: usize, + _params: TextSearchParams, + _allowed_ids: Option<&HashSet>, + ) -> NodeDbResult> { + Ok(Vec::new()) + } + + async fn execute_sql(&self, _query: &str, _params: &[Value]) -> NodeDbResult { + Ok(QueryResult::empty()) + } + } + + #[test] + fn trait_is_object_safe() { + fn _accepts_dyn(_db: &dyn NodeDb) {} + let db = MockDb; + _accepts_dyn(&db); + } + + #[test] + fn trait_works_with_arc() { + use std::sync::Arc; + let db: Arc = Arc::new(MockDb); + let _ = db; + } + + #[tokio::test] + async fn mock_vector_search() { + let db = MockDb; + let results = db + .vector_search("embeddings", &[0.1, 0.2, 0.3], 5, None, None) + .await + .unwrap(); + assert_eq!(results.len(), 1); + assert_eq!(results[0].id, "vec-1"); + assert!(results[0].distance < 1.0); + } + + #[tokio::test] + async fn mock_vector_insert_and_delete() { + let db = MockDb; + db.vector_insert("coll", "v1", &[1.0, 2.0], None) + .await + .unwrap(); + db.vector_delete("coll", "v1").await.unwrap(); + } + + #[tokio::test] + async fn mock_graph_stats_returns_zero() { + let db = MockDb; + let result = db.graph_stats(Some("social"), None).await.unwrap(); + assert_eq!(result.len(), 1); + let stats = &result[0]; + assert_eq!(stats.collection, "social"); + assert_eq!(stats.node_count, 0); + assert_eq!(stats.edge_count, 0); + assert_eq!(stats.distinct_label_count, 0); + assert!(stats.labels.is_empty()); + } + + #[tokio::test] + async fn mock_graph_stats_tenant_wide_uses_mock_key() { + let db = MockDb; + let result = db.graph_stats(None, None).await.unwrap(); + assert_eq!(result.len(), 1); + assert_eq!(result[0].collection, "mock"); + } + + #[tokio::test] + async fn mock_graph_operations() { + let db = MockDb; + let start = NodeId::try_new("alice").expect("test fixture"); + let subgraph = db + .graph_traverse( + "social", + &start, + 2, + nodedb_types::graph::Direction::Out, + None, + ) + .await + .unwrap(); + assert_eq!(subgraph.node_count(), 0); + + let from = NodeId::try_new("alice").expect("test fixture"); + let to = NodeId::try_new("bob").expect("test fixture"); + let edge_id = db + .graph_insert_edge("social", &from, &to, "KNOWS", None) + .await + .unwrap(); + assert_eq!(edge_id.src.as_str(), "alice"); + assert_eq!(edge_id.dst.as_str(), "bob"); + assert_eq!(edge_id.label, "KNOWS"); + assert_eq!(edge_id.seq, 0); + + db.graph_delete_edge("social", &edge_id).await.unwrap(); + } + + #[tokio::test] + async fn mock_document_operations() { + let db = MockDb; + let doc = db.document_get("notes", "n1").await.unwrap().unwrap(); + assert_eq!(doc.id, "n1"); + assert_eq!(doc.get_str("title"), Some("test")); + + let mut new_doc = Document::new("n2"); + new_doc.set("body", Value::String("hello".into())); + db.document_put("notes", new_doc).await.unwrap(); + + db.document_delete("notes", "n1").await.unwrap(); + } + + #[tokio::test] + async fn mock_execute_sql() { + let db = MockDb; + let result = db.execute_sql("SELECT 1", &[]).await.unwrap(); + assert_eq!(result.row_count(), 0); + } + + /// Verify the full "one API, any runtime" pattern: application + /// code switches between `NodeDbLite` and `NodeDbRemote` only at + /// the construction site. + #[tokio::test] + async fn unified_api_pattern() { + use std::sync::Arc; + + let db: Arc = Arc::new(MockDb); + + let results = db + .vector_search("knowledge_base", &[0.1, 0.2], 5, None, None) + .await + .unwrap(); + assert!(!results.is_empty()); + + let start = NodeId::from_validated(results[0].id.clone()); + let _subgraph = db + .graph_traverse( + "knowledge_base", + &start, + 2, + nodedb_types::graph::Direction::Out, + None, + ) + .await + .unwrap(); + + let doc = Document::new("note-1"); + db.document_put("notes", doc).await.unwrap(); + } + + #[test] + fn default_proto_version_is_zero() { + let db = MockDb; + assert_eq!(db.proto_version(), 0); + } + + #[test] + fn default_capabilities_is_zero() { + let db = MockDb; + assert_eq!(db.capabilities(), 0); + let caps = Capabilities::from_raw(db.capabilities()); + assert!(!caps.supports_streaming()); + assert!(!caps.supports_graphrag()); + } + + #[test] + fn default_server_version_is_empty() { + let db = MockDb; + assert!(db.server_version().is_empty()); + } + + #[test] + fn default_limits_all_none() { + let db = MockDb; + let limits = db.limits(); + assert!(limits.max_vector_dim.is_none()); + assert!(limits.max_top_k.is_none()); + assert!(limits.max_scan_limit.is_none()); + assert!(limits.max_batch_size.is_none()); + assert!(limits.max_crdt_delta_bytes.is_none()); + assert!(limits.max_query_text_bytes.is_none()); + assert!(limits.max_graph_depth.is_none()); + } + + #[test] + fn capabilities_newtype_smoke() { + use nodedb_types::protocol::{CAP_FTS, CAP_STREAMING}; + let caps = Capabilities::from_raw(CAP_STREAMING | CAP_FTS); + assert!(caps.supports_streaming()); + assert!(caps.supports_fts()); + assert!(!caps.supports_graphrag()); + assert!(!caps.supports_crdt()); + } +} diff --git a/nodedb-client/src/traits/core/default_impls/graph.rs b/nodedb-client/src/traits/core/default_impls/graph.rs new file mode 100644 index 000000000..91a16b3b1 --- /dev/null +++ b/nodedb-client/src/traits/core/default_impls/graph.rs @@ -0,0 +1,98 @@ +// SPDX-License-Identifier: Apache-2.0 + +use nodedb_types::error::{NodeDbError, NodeDbResult}; +use nodedb_types::filter::EdgeFilter; +use nodedb_types::id::NodeId; +use std::collections::HashMap; + +use super::super::trait_def::NodeDb; + +pub(in crate::traits::core) fn graph_pagerank_default( + collection: &str, + personalization: Option>, + damping: Option, + max_iterations: Option, +) -> NodeDbResult> { + let _ = (collection, personalization, damping, max_iterations); + Err(NodeDbError::storage( + "graph_pagerank is not implemented for this NodeDb backend", + )) +} + +/// Default forward-BFS shortest-path search built on `graph_traverse`. +/// +/// See [`crate::traits::core::NodeDb::graph_shortest_path`] for the +/// full contract. +pub(in crate::traits::core) async fn graph_shortest_path_default( + this: &T, + collection: &str, + from: &NodeId, + to: &NodeId, + max_depth: u8, + edge_filter: Option<&EdgeFilter>, +) -> NodeDbResult>> { + if from == to { + return Ok(Some(vec![from.clone()])); + } + if max_depth == 0 { + return Ok(None); + } + + // Map of `node -> parent` used to reconstruct the path once the + // target is reached. The source has no parent entry. + let mut parent: HashMap = HashMap::new(); + let mut frontier: Vec = vec![from.clone()]; + + for _ in 0..max_depth { + let mut next_frontier: Vec = Vec::new(); + for node in &frontier { + let sg = this + .graph_traverse( + collection, + node, + 1, + nodedb_types::graph::Direction::Out, + edge_filter, + ) + .await?; + for edge in &sg.edges { + // Only follow edges originating from the current + // node — `graph_traverse` may include adjacent + // edges that don't extend the BFS frontier. + if &edge.from != node { + continue; + } + let dst = &edge.to; + if dst == from || parent.contains_key(dst) { + continue; + } + parent.insert(dst.clone(), node.clone()); + if dst == to { + let mut path = vec![to.clone()]; + let mut cur = to.clone(); + while &cur != from { + let p = parent + .get(&cur) + .ok_or_else(|| { + NodeDbError::storage(format!( + "graph shortest path parent missing for '{}': retry traversal", + cur.as_str() + )) + })? + .clone(); + path.push(p.clone()); + cur = p; + } + path.reverse(); + return Ok(Some(path)); + } + next_frontier.push(dst.clone()); + } + } + if next_frontier.is_empty() { + return Ok(None); + } + frontier = next_frontier; + } + Ok(None) +} diff --git a/nodedb-client/src/traits/core/default_impls/mod.rs b/nodedb-client/src/traits/core/default_impls/mod.rs new file mode 100644 index 000000000..ac901aa3a --- /dev/null +++ b/nodedb-client/src/traits/core/default_impls/mod.rs @@ -0,0 +1,21 @@ +// SPDX-License-Identifier: Apache-2.0 + +//! Default-provided method bodies for [`crate::traits::core::NodeDb`]. +//! +//! Each function implements one default-provided `NodeDb` method. + +mod batch; +mod collections; +mod dispatch; +mod graph; + +pub(super) use batch::{batch_graph_insert_edges_default, batch_vector_insert_default}; +pub(super) use collections::{ + drop_collection_purge_default, list_dropped_collections_default, on_collection_purged_default, + undrop_collection_default, +}; +pub(super) use dispatch::{ + document_get_as_of_default, document_put_with_valid_time_default, + document_put_with_vector_default, vector_insert_field_default, vector_search_field_default, +}; +pub(super) use graph::{graph_pagerank_default, graph_shortest_path_default}; diff --git a/nodedb-client/src/traits/core/trait_def.rs b/nodedb-client/src/traits/core/trait_def.rs index 5f4aec3c5..b82e912e7 100644 --- a/nodedb-client/src/traits/core/trait_def.rs +++ b/nodedb-client/src/traits/core/trait_def.rs @@ -1,18 +1,19 @@ // SPDX-License-Identifier: Apache-2.0 -//! The `NodeDb` trait: unified query interface for both Origin and Lite. +//! The `NodeDb` trait: one query interface for Origin and Lite. //! -//! Application code writes against this trait once. The runtime determines -//! whether queries execute locally (in-memory engines on Lite) or remotely -//! (pgwire to Origin). +//! Application code writes against this trait once. Lite runs each call on +//! in-process engines. The native and pgwire clients send it to Origin. All +//! methods are `async`: Tokio on native targets, `wasm-bindgen-futures` on +//! WASM. //! -//! All methods are `async` — on native this runs on Tokio, on WASM this -//! runs on `wasm-bindgen-futures`. +//! `NodeDb` is one trait, so an implementation is one `impl` block and one +//! import. Long provided-method bodies live in `default_impls`. //! -//! This file must remain a single block for object-safety: Rust does not -//! permit a `trait` body to be split across files. Splitting `NodeDb` into -//! supertraits would break the `Arc` pattern all callers depend -//! on and is therefore out of scope for any mechanical refactor. +//! Each method states its contract first. The `On Lite` / `On Remote` / +//! `On Native` lines name what each backend runs. `Remote` is the pgwire +//! client. `Native` is the MessagePack client, the canonical typed +//! transport. "Key column" is a collection's declared primary key, else `id`. use std::collections::HashSet; @@ -35,30 +36,29 @@ use crate::traits::document::CollectionPurgedHandler; /// Unified database interface for NodeDB. /// -/// Two implementations: -/// - `NodeDbLite`: executes queries against in-memory HNSW/CSR/Loro engines -/// on the edge device. Writes produce CRDT deltas synced to Origin in background. -/// - `NodeDbRemote`: translates trait calls into parameterized SQL and sends -/// them over pgwire to the Origin cluster. -/// -/// The developer writes agent logic once. Switching between local and cloud -/// is a one-line configuration change. +/// Implementations: +/// - `NodeDbLite`: in-process HNSW/CSR/Loro engines on the edge device. +/// Writes produce CRDT deltas synced to Origin in the background. +/// - `NativeClient`: the MessagePack protocol to Origin. +/// - `NodeDbRemote`: SQL over pgwire to Origin. #[cfg_attr(not(target_arch = "wasm32"), async_trait)] #[cfg_attr(target_arch = "wasm32", async_trait(?Send))] pub trait NodeDb: NodeDbMarker { - // ─── Vector Operations ─────────────────────────────────────────── - /// Search for the `k` nearest vectors to `query` in `collection`. /// - /// Returns results ordered by ascending distance. Optional metadata - /// filter constrains which vectors are considered. When `allowed_ids` - /// is `Some`, only documents whose string ID appears in the set are - /// eligible — the filter is pushed into HNSW traversal on Lite so - /// the returned top-k is drawn exclusively from the allowed set. - /// - /// On Lite: direct in-memory HNSW search with optional ID prefilter. Sub-millisecond. - /// On Remote: translated to `SELECT ... ORDER BY embedding <-> $1 LIMIT $2` - /// (allowed_ids is ignored on the remote path — pass `None`). + /// Results are ordered by ascending distance. `filter` constrains which + /// vectors are considered. `Some(allowed_ids)` applies before ranking on + /// every implementation: the top-k comes from the allowed set even when + /// nearer vectors lie outside it. An empty set returns no result. + /// + /// On Lite: in-memory HNSW search with the id prefilter pushed into + /// traversal. + /// On Remote: `SELECT id, vector_distance(..) AS distance FROM + /// [WHERE ] [AND IN (..)] ORDER BY vector_distance(..) + /// LIMIT `. The server lowers the key restriction to a candidate bitmap + /// the index search honors. + /// On Native: the ids travel in the request. The server lowers them to + /// the same candidate bitmap. async fn vector_search( &self, collection: &str, @@ -70,8 +70,8 @@ pub trait NodeDb: NodeDbMarker { /// Insert a vector with optional metadata into `collection`. /// - /// On Lite: inserts into in-memory HNSW + emits CRDT delta + persists to SQLite. - /// On Remote: translated to `INSERT INTO collection (id, embedding, metadata) VALUES (...)`. + /// On Lite: HNSW insert + CRDT delta + persist to local storage. + /// On Remote: `INSERT INTO collection (id, embedding, metadata) VALUES (...)`. async fn vector_insert( &self, collection: &str, @@ -86,34 +86,33 @@ pub trait NodeDb: NodeDbMarker { /// On Remote: `DELETE FROM collection WHERE id = $1`. async fn vector_delete(&self, collection: &str, id: &str) -> NodeDbResult<()>; - // ─── Graph Operations ──────────────────────────────────────────── - - /// Traverse the graph from `start` up to `depth` hops within - /// `collection`. + /// Traverse the graph from `start` up to `depth` hops within `collection`. + /// + /// Edges are scoped per collection, so the caller picks which graph to + /// walk. Returns the discovered subgraph (nodes + edges). `direction` + /// selects outgoing, incoming, or both orientations. /// - /// `collection` names the graph collection holding the adjacency - /// data. NodeDB's graph overlay scopes edges per collection, so the - /// caller picks which graph to walk. Returns the discovered subgraph - /// (nodes + edges). Optional edge filter constrains which edges are - /// followed. + /// Every label of the filter is followed. Every property filter must + /// admit an edge before it is crossed. Each returned edge carries its + /// properties. /// - /// On Lite: direct CSR pointer-chasing in contiguous memory. Microseconds. - /// On Remote: `GRAPH TRAVERSE IN '' FROM '' DEPTH - /// [LABEL '']`. + /// On Lite: CSR pointer-chasing in contiguous memory. + /// On Remote and Native: `GRAPH TRAVERSE IN '' FROM '' + /// DEPTH DIRECTION out|in|both [LABEL '', '' …] + /// [EDGE WHERE ]`, one round trip. async fn graph_traverse( &self, collection: &str, start: &NodeId, depth: u8, + direction: nodedb_types::graph::Direction, edge_filter: Option<&EdgeFilter>, ) -> NodeDbResult; - /// Insert a directed edge from `from` to `to` with the given label - /// into `collection`. + /// Insert a directed edge `from` → `to` with label `edge_type` into + /// `collection`. Returns the generated edge ID. /// - /// Returns the generated edge ID. - /// - /// On Lite: appends to mutable adjacency buffer + CRDT delta + SQLite. + /// On Lite: mutable adjacency buffer + CRDT delta + local storage. /// On Remote: `GRAPH INSERT EDGE IN '' FROM '' TO '' TYPE '', '' …] + /// [EDGE WHERE ]`, one round trip. async fn graph_shortest_path( &self, collection: &str, @@ -404,25 +373,26 @@ pub trait NodeDb: NodeDbMarker { .await } - // ─── Text Search ──────────────────────────────────────────────── - - /// Full-text search with BM25 scoring against the FTS-indexed - /// `field` on `collection`. - /// - /// NodeDB's FTS is per-field — every BM25 index is scoped to one - /// declared field, so the caller names which field to search. - /// Returns document IDs with relevance scores, ordered by - /// descending score. Pass [`TextSearchParams::default()`] for - /// standard OR-mode non-fuzzy search. - /// - /// Default returns `Err` — `Ok(Vec::new())` is indistinguishable - /// from a real "no matches" answer and would silently mask the - /// missing implementation. Implementations must override (e.g., a - /// `SEARCH IN '' FIELD '' QUERY ''` round-trip - /// via `execute_sql`). - /// Full-text BM25 search. When `allowed_ids` is `Some`, only documents - /// whose ID is in the set are returned. On Lite, the filter is applied - /// after an over-fetch; on Remote, `allowed_ids` is ignored (pass `None`). + /// Full-text BM25 search of the indexed `field` of `collection`. + /// + /// Every BM25 index covers one declared field, so the caller names it. + /// An empty `field` searches the whole-document index. Returns ids with + /// scores, best first, at most `top_k`. `Some(allowed_ids)` returns only + /// those ids. + /// + /// `params` means the same on every implementation. `mode` combines the + /// terms (`Or`: any, `And`: all). `fuzzy`, or the collection's fuzzy + /// default, lets a term with no exact match fall back to a Levenshtein + /// match. [`TextSearchParams::default`] is `{ mode: Or, fuzzy: false }`. + /// + /// On Lite: honors every `params` value, and applies `allowed_ids` after + /// an over-fetch. + /// On Remote and Native: one shared statement, `SELECT id, bm25_score(.., + /// opts) AS score FROM WHERE text_match(.., opts) [AND IN (..)] ORDER BY score DESC LIMIT `, with `opts` the + /// non-default `mode => '..'` and `fuzzy => ..`. Both return the same ids + /// and scores. A double-quoted query is a phrase query, which takes only + /// the default `params`. async fn text_search( &self, collection: &str, @@ -431,11 +401,7 @@ pub trait NodeDb: NodeDbMarker { top_k: usize, params: TextSearchParams, allowed_ids: Option<&HashSet>, - ) -> NodeDbResult> { - default_impls::text_search_default(collection, field, query, top_k, params, allowed_ids) - } - - // ─── Batch Operations ─────────────────────────────────────────── + ) -> NodeDbResult>; /// Batch insert vectors — amortizes CRDT delta export to O(1) per batch. async fn batch_vector_insert( @@ -456,407 +422,73 @@ pub trait NodeDb: NodeDbMarker { default_impls::batch_graph_insert_edges_default(self, collection, edges).await } - // ─── Connection Metadata ───────────────────────────────────────────── - /// The protocol version negotiated during the connection handshake. - /// - /// Returns `0` for implementations that do not maintain a persistent - /// connection and therefore never perform a handshake. + /// `0` for an implementation that performs no handshake. fn proto_version(&self) -> u16 { 0 } - /// The raw capability bitfield advertised by the server. - /// - /// Returns `0` when no handshake was performed. Use - /// `Capabilities::from_raw(self.capabilities())` for named predicates. + /// The raw capability bitfield advertised by the server. `0` when no + /// handshake ran. `Capabilities::from_raw` gives named predicates. fn capabilities(&self) -> u64 { 0 } /// The server version string from `HelloAckFrame` (e.g. `"0.1.0-dev"`). - /// - /// Returns an empty string when no handshake was performed. + /// Empty when no handshake ran. fn server_version(&self) -> String { String::new() } - /// Per-operation limits announced by the server. - /// - /// All fields are `None` when no handshake was performed — the caller - /// should treat `None` as "no server-side cap" for that dimension. + /// Per-operation limits announced by the server. Every field is `None` + /// when no handshake ran. `None` means no server-side cap. fn limits(&self) -> Limits { Limits::default() } - // ─── SQL Escape Hatch ──────────────────────────────────────────── - /// Execute a raw SQL query with parameters. /// - /// On Lite: requires the `sql` feature flag (compiles in DataFusion parser). - /// Returns `NodeDbError::SqlNotEnabled` if the feature is not compiled in. + /// On Lite: requires the `sql` feature. Without it, returns + /// `NodeDbError::SqlNotEnabled`. /// On Remote: pass-through to Origin via pgwire. /// - /// For most AI agent workloads, the typed methods above are sufficient - /// and faster. Use this for BI tools, existing ORMs, or ad-hoc queries. + /// The typed methods above cover most agent workloads. Use this for BI + /// tools, ORMs, and ad-hoc queries. async fn execute_sql(&self, query: &str, params: &[Value]) -> NodeDbResult; - // ─── Collection Lifecycle (soft-delete / undrop / hard-delete) ─── - /// Restore a soft-deleted collection within its retention window. /// - /// Equivalent to `UNDROP COLLECTION `. Fails with 42P01 if - /// the retention window has elapsed and the row is gone, or with - /// 42501 if the caller is neither preserved owner nor admin. - /// - /// Default impl routes through `execute_sql` so any implementation - /// that can execute SQL inherits the correct behavior for free. + /// Runs `UNDROP COLLECTION ` through `execute_sql`. Fails with + /// 42P01 once the retention window has elapsed, or with 42501 when the + /// caller is neither the preserved owner nor an admin. async fn undrop_collection(&self, name: &str) -> NodeDbResult<()> { default_impls::undrop_collection_default(self, name).await } /// Hard-delete a collection, skipping soft-delete and retention. /// - /// Equivalent to `DROP COLLECTION PURGE`. Admin-only on the - /// server; the server rejects non-admin callers with 42501. - /// Bypasses the retention safety net — data is unrecoverable. + /// Runs `DROP COLLECTION PURGE`. Admin-only: the server rejects + /// other callers with 42501. The data is unrecoverable. async fn drop_collection_purge(&self, name: &str) -> NodeDbResult<()> { default_impls::drop_collection_purge_default(self, name).await } - /// List every soft-deleted collection in the current tenant that - /// is still within its retention window. + /// List the current tenant's soft-deleted collections still within their + /// retention window. Empty when there are none. /// - /// Equivalent to `SELECT tenant_id, name, owner, deactivated_at_ns, - /// retention_expires_at_ns FROM _system.dropped_collections`. - /// Returns `Vec` — empty if no soft-deleted rows - /// exist for the caller's tenant. + /// Reads `tenant_id, name, owner, deactivated_at_ns, + /// retention_expires_at_ns` from `_system.dropped_collections`. async fn list_dropped_collections(&self) -> NodeDbResult> { default_impls::list_dropped_collections_default(self).await } - /// Register a handler fired when a collection the caller has - /// synced is purged on Origin and the local copy is removed. + /// Register a handler fired when a collection the caller synced is purged + /// on Origin and the local copy is removed. /// - /// Default impl returns `NodeDbError::storage` with a - /// `"not supported"` detail — implementations that maintain a - /// sync client (Lite, any future push-capable remote client) - /// override with registration into their internal handler list. - /// Stateless clients (pgwire-only `NodeDbRemote`) have nothing - /// to push, so the default rejection is the correct behavior. + /// The default returns `NodeDbError::storage` ("not supported"). + /// Implementations with a sync client (Lite) override it. A stateless + /// pgwire client has nothing to push, so the default is correct there. async fn on_collection_purged(&self, _handler: CollectionPurgedHandler) -> NodeDbResult<()> { default_impls::on_collection_purged_default() } } - -#[cfg(test)] -mod tests { - use super::*; - use crate::capabilities::Capabilities; - use async_trait::async_trait; - use nodedb_types::document::Document; - use nodedb_types::error::{NodeDbError, NodeDbResult}; - use nodedb_types::filter::{EdgeFilter, MetadataFilter}; - use nodedb_types::graph::GraphStats; - use nodedb_types::id::{EdgeId, NodeId}; - use nodedb_types::result::{QueryResult, SearchResult, SubGraph}; - use nodedb_types::value::Value; - use std::collections::HashMap; - - /// Mock implementation to verify the trait is object-safe and - /// can be used as `Arc`. - struct MockDb; - - #[cfg_attr(not(target_arch = "wasm32"), async_trait)] - #[cfg_attr(target_arch = "wasm32", async_trait(?Send))] - impl NodeDb for MockDb { - async fn vector_search( - &self, - _collection: &str, - _query: &[f32], - _k: usize, - _filter: Option<&MetadataFilter>, - _allowed_ids: Option<&HashSet>, - ) -> NodeDbResult> { - Ok(vec![SearchResult { - id: "vec-1".into(), - node_id: None, - distance: 0.1, - metadata: HashMap::new(), - }]) - } - - async fn vector_insert( - &self, - _collection: &str, - _id: &str, - _embedding: &[f32], - _metadata: Option, - ) -> NodeDbResult<()> { - Ok(()) - } - - async fn vector_delete(&self, _collection: &str, _id: &str) -> NodeDbResult<()> { - Ok(()) - } - - async fn graph_traverse( - &self, - _collection: &str, - _start: &NodeId, - _depth: u8, - _edge_filter: Option<&EdgeFilter>, - ) -> NodeDbResult { - Ok(SubGraph::empty()) - } - - async fn graph_insert_edge( - &self, - _collection: &str, - from: &NodeId, - to: &NodeId, - edge_type: &str, - _properties: Option, - ) -> NodeDbResult { - EdgeId::try_first(from.clone(), to.clone(), edge_type) - .map_err(|e| NodeDbError::storage(format!("invalid edge label: {e}"))) - } - - async fn graph_delete_edge( - &self, - _collection: &str, - _edge_id: &EdgeId, - ) -> NodeDbResult<()> { - Ok(()) - } - - async fn graph_stats( - &self, - collection: Option<&str>, - _as_of: Option, - ) -> NodeDbResult> { - Ok(vec![GraphStats::zero(collection.unwrap_or("mock"))]) - } - - async fn document_get( - &self, - _collection: &str, - id: &str, - ) -> NodeDbResult> { - let mut doc = Document::new(id); - doc.set("title", Value::String("test".into())); - Ok(Some(doc)) - } - - async fn document_put(&self, _collection: &str, _doc: Document) -> NodeDbResult<()> { - Ok(()) - } - - async fn document_delete(&self, _collection: &str, _id: &str) -> NodeDbResult<()> { - Ok(()) - } - - async fn list_insert( - &self, - _collection: &str, - _document_id: &str, - _list_path: &str, - _index: usize, - _fields: &Value, - ) -> NodeDbResult<()> { - Ok(()) - } - - async fn list_delete( - &self, - _collection: &str, - _document_id: &str, - _list_path: &str, - _index: usize, - ) -> NodeDbResult<()> { - Ok(()) - } - - async fn list_move( - &self, - _collection: &str, - _document_id: &str, - _list_path: &str, - _from_index: usize, - _to_index: usize, - ) -> NodeDbResult<()> { - Ok(()) - } - - async fn execute_sql(&self, _query: &str, _params: &[Value]) -> NodeDbResult { - Ok(QueryResult::empty()) - } - } - - #[test] - fn trait_is_object_safe() { - fn _accepts_dyn(_db: &dyn NodeDb) {} - let db = MockDb; - _accepts_dyn(&db); - } - - #[test] - fn trait_works_with_arc() { - use std::sync::Arc; - let db: Arc = Arc::new(MockDb); - let _ = db; - } - - #[tokio::test] - async fn mock_vector_search() { - let db = MockDb; - let results = db - .vector_search("embeddings", &[0.1, 0.2, 0.3], 5, None, None) - .await - .unwrap(); - assert_eq!(results.len(), 1); - assert_eq!(results[0].id, "vec-1"); - assert!(results[0].distance < 1.0); - } - - #[tokio::test] - async fn mock_vector_insert_and_delete() { - let db = MockDb; - db.vector_insert("coll", "v1", &[1.0, 2.0], None) - .await - .unwrap(); - db.vector_delete("coll", "v1").await.unwrap(); - } - - #[tokio::test] - async fn mock_graph_stats_returns_zero() { - let db = MockDb; - let result = db.graph_stats(Some("social"), None).await.unwrap(); - assert_eq!(result.len(), 1); - let stats = &result[0]; - assert_eq!(stats.collection, "social"); - assert_eq!(stats.node_count, 0); - assert_eq!(stats.edge_count, 0); - assert_eq!(stats.distinct_label_count, 0); - assert!(stats.labels.is_empty()); - } - - #[tokio::test] - async fn mock_graph_stats_tenant_wide_uses_mock_key() { - let db = MockDb; - let result = db.graph_stats(None, None).await.unwrap(); - assert_eq!(result.len(), 1); - assert_eq!(result[0].collection, "mock"); - } - - #[tokio::test] - async fn mock_graph_operations() { - let db = MockDb; - let start = NodeId::try_new("alice").expect("test fixture"); - let subgraph = db.graph_traverse("social", &start, 2, None).await.unwrap(); - assert_eq!(subgraph.node_count(), 0); - - let from = NodeId::try_new("alice").expect("test fixture"); - let to = NodeId::try_new("bob").expect("test fixture"); - let edge_id = db - .graph_insert_edge("social", &from, &to, "KNOWS", None) - .await - .unwrap(); - assert_eq!(edge_id.src.as_str(), "alice"); - assert_eq!(edge_id.dst.as_str(), "bob"); - assert_eq!(edge_id.label, "KNOWS"); - assert_eq!(edge_id.seq, 0); - - db.graph_delete_edge("social", &edge_id).await.unwrap(); - } - - #[tokio::test] - async fn mock_document_operations() { - let db = MockDb; - let doc = db.document_get("notes", "n1").await.unwrap().unwrap(); - assert_eq!(doc.id, "n1"); - assert_eq!(doc.get_str("title"), Some("test")); - - let mut new_doc = Document::new("n2"); - new_doc.set("body", Value::String("hello".into())); - db.document_put("notes", new_doc).await.unwrap(); - - db.document_delete("notes", "n1").await.unwrap(); - } - - #[tokio::test] - async fn mock_execute_sql() { - let db = MockDb; - let result = db.execute_sql("SELECT 1", &[]).await.unwrap(); - assert_eq!(result.row_count(), 0); - } - - /// Verify the full "one API, any runtime" pattern: application - /// code switches between `NodeDbLite` and `NodeDbRemote` only at - /// the construction site. - #[tokio::test] - async fn unified_api_pattern() { - use std::sync::Arc; - - let db: Arc = Arc::new(MockDb); - - let results = db - .vector_search("knowledge_base", &[0.1, 0.2], 5, None, None) - .await - .unwrap(); - assert!(!results.is_empty()); - - let start = NodeId::from_validated(results[0].id.clone()); - let _subgraph = db - .graph_traverse("knowledge_base", &start, 2, None) - .await - .unwrap(); - - let doc = Document::new("note-1"); - db.document_put("notes", doc).await.unwrap(); - } - - #[test] - fn default_proto_version_is_zero() { - let db = MockDb; - assert_eq!(db.proto_version(), 0); - } - - #[test] - fn default_capabilities_is_zero() { - let db = MockDb; - assert_eq!(db.capabilities(), 0); - let caps = Capabilities::from_raw(db.capabilities()); - assert!(!caps.supports_streaming()); - assert!(!caps.supports_graphrag()); - } - - #[test] - fn default_server_version_is_empty() { - let db = MockDb; - assert!(db.server_version().is_empty()); - } - - #[test] - fn default_limits_all_none() { - let db = MockDb; - let limits = db.limits(); - assert!(limits.max_vector_dim.is_none()); - assert!(limits.max_top_k.is_none()); - assert!(limits.max_scan_limit.is_none()); - assert!(limits.max_batch_size.is_none()); - assert!(limits.max_crdt_delta_bytes.is_none()); - assert!(limits.max_query_text_bytes.is_none()); - assert!(limits.max_graph_depth.is_none()); - } - - #[test] - fn capabilities_newtype_smoke() { - use nodedb_types::protocol::{CAP_FTS, CAP_STREAMING}; - let caps = Capabilities::from_raw(CAP_STREAMING | CAP_FTS); - assert!(caps.supports_streaming()); - assert!(caps.supports_fts()); - assert!(!caps.supports_graphrag()); - assert!(!caps.supports_crdt()); - } -} diff --git a/nodedb-cluster-tests/tests/dml_suite/cases/sql_cluster_cross_node_dml.rs b/nodedb-cluster-tests/tests/dml_suite/cases/sql_cluster_cross_node_dml.rs index f4470ae60..d415d499a 100644 --- a/nodedb-cluster-tests/tests/dml_suite/cases/sql_cluster_cross_node_dml.rs +++ b/nodedb-cluster-tests/tests/dml_suite/cases/sql_cluster_cross_node_dml.rs @@ -29,6 +29,8 @@ mod graph_algo_pagerank_personalized_cross_node; mod graph_algo_wcc_cross_node; #[path = "../../sql_cluster_cross_node_dml_tests/graph_delete_reverse_cross_node.rs"] mod graph_delete_reverse_cross_node; +#[path = "../../sql_cluster_cross_node_dml_tests/graph_edge_predicate_cross_node.rs"] +mod graph_edge_predicate_cross_node; #[path = "../../sql_cluster_cross_node_dml_tests/graph_homed_read_failover.rs"] mod graph_homed_read_failover; #[path = "../../sql_cluster_cross_node_dml_tests/graph_homed_read_per_core.rs"] diff --git a/nodedb-cluster-tests/tests/misc_suite/cases/cluster_triggers.rs b/nodedb-cluster-tests/tests/misc_suite/cases/cluster_triggers.rs index 213dc9d7e..c98fb8e65 100644 --- a/nodedb-cluster-tests/tests/misc_suite/cases/cluster_triggers.rs +++ b/nodedb-cluster-tests/tests/misc_suite/cases/cluster_triggers.rs @@ -203,6 +203,7 @@ fn event_source_preserved_through_write_event() { .duration_since(std::time::UNIX_EPOCH) .ok() .and_then(|elapsed| u64::try_from(elapsed.as_nanos()).ok()), + image_fault: None, }; // After leader failover, new leader's Event Plane replays from WAL. // The replayed events have source: User → triggers fire. diff --git a/nodedb-cluster-tests/tests/sql_cluster_cross_node_dml_tests/graph_edge_predicate_cross_node.rs b/nodedb-cluster-tests/tests/sql_cluster_cross_node_dml_tests/graph_edge_predicate_cross_node.rs new file mode 100644 index 000000000..01ff0cc2d --- /dev/null +++ b/nodedb-cluster-tests/tests/sql_cluster_cross_node_dml_tests/graph_edge_predicate_cross_node.rs @@ -0,0 +1,143 @@ +// SPDX-License-Identifier: BUSL-1.1 + +//! Cross-node `EDGE WHERE` and edge properties on `GRAPH TRAVERSE` and +//! `GRAPH PATH`. +//! +//! Each frontier node expands at the node that owns its key vShard, and the +//! predicate runs on the core that stores the edge. A coordinator that +//! evaluated the predicate against its own partitions, or dropped it on the +//! remote `NeighborsMulti`, reaches a node the predicate rejects or misses +//! one it admits. Destination names are distinct, so the edges spread +//! across vShards spanning nodes. + +use std::collections::BTreeSet; +use std::time::Duration; + +use crate::common::cluster_harness::{TestCluster, wait_for}; + +async fn result_json(client: &tokio_postgres::Client, sql: &str) -> serde_json::Value { + let msgs = client + .simple_query(sql) + .await + .unwrap_or_else(|e| panic!("{sql}: {e}")); + let row = msgs + .iter() + .find_map(|m| match m { + tokio_postgres::SimpleQueryMessage::Row(r) => Some(r), + _ => None, + }) + .expect("graph query returned no result row"); + let raw = row.get("result").expect("result column present"); + serde_json::from_str(raw).expect("result column is valid JSON") +} + +fn node_ids(v: &serde_json::Value) -> BTreeSet { + v["nodes"] + .as_array() + .expect("nodes array") + .iter() + .map(|n| n["id"].as_str().expect("node id").to_string()) + .collect() +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 6)] +async fn edge_predicate_and_properties_hold_from_any_coordinator() { + let cluster = TestCluster::spawn_three().await.expect("3-node cluster"); + + cluster + .exec_ddl_on_any_leader("CREATE COLLECTION ep_xnode") + .await + .expect("CREATE COLLECTION ep_xnode"); + wait_for( + "all 3 nodes see ep_xnode", + Duration::from_secs(10), + Duration::from_millis(50), + || { + cluster + .nodes + .iter() + .all(|n| n.cached_collection_count() >= 1) + }, + ) + .await; + + // root -> hot_i (score 9) and root -> cold_i (score 1); each hot_i -> + // deep_i (score 9) and each cold_i -> lost_i (score 9). + const FAN: usize = 6; + let mut edges = Vec::new(); + for i in 0..FAN { + edges.push(("root".to_string(), format!("hot_{i}"), 9)); + edges.push(("root".to_string(), format!("cold_{i}"), 1)); + edges.push((format!("hot_{i}"), format!("deep_{i}"), 9)); + edges.push((format!("cold_{i}"), format!("lost_{i}"), 9)); + } + for (src, dst, score) in &edges { + cluster.nodes[0] + .client + .simple_query(&format!( + "GRAPH INSERT EDGE IN 'ep_xnode' FROM '{src}' TO '{dst}' TYPE 'L' \ + PROPERTIES {{ score: {score} }}" + )) + .await + .unwrap_or_else(|e| panic!("insert edge {src}->{dst}: {e}")); + } + + let unfiltered = "GRAPH TRAVERSE IN 'ep_xnode' FROM 'root' DEPTH 2 DIRECTION out"; + let total = 1 + 4 * FAN; + for idx in 0..cluster.nodes.len() { + wait_for( + &format!("node {idx} reaches every seeded node"), + Duration::from_secs(20), + Duration::from_millis(100), + || { + tokio::task::block_in_place(|| { + tokio::runtime::Handle::current().block_on(async { + node_ids(&result_json(&cluster.nodes[idx].client, unfiltered).await).len() + == total + }) + }) + }, + ) + .await; + } + + let expected: BTreeSet = std::iter::once("root".to_string()) + .chain((0..FAN).flat_map(|i| [format!("hot_{i}"), format!("deep_{i}")])) + .collect(); + for idx in 0..cluster.nodes.len() { + let client = &cluster.nodes[idx].client; + let got = result_json( + client, + "GRAPH TRAVERSE IN 'ep_xnode' FROM 'root' DEPTH 2 DIRECTION out EDGE WHERE score > 5", + ) + .await; + assert_eq!(node_ids(&got), expected, "node {idx}: {got}"); + for edge in got["edges"].as_array().expect("edges array") { + assert_eq!( + edge["properties"], + serde_json::json!({ "score": 9 }), + "node {idx}: every crossed edge carries score 9: {got}" + ); + } + + // Backward expansion tests the stored edge `hot_0 -> deep_0`. + let path = result_json( + client, + "GRAPH PATH IN 'ep_xnode' FROM 'root' TO 'deep_0' MAX_DEPTH 4 EDGE WHERE score > 5", + ) + .await; + assert_eq!( + path, + serde_json::json!(["root", "hot_0", "deep_0"]), + "node {idx}" + ); + let blocked = result_json( + client, + "GRAPH PATH IN 'ep_xnode' FROM 'root' TO 'lost_0' MAX_DEPTH 4 EDGE WHERE score > 5", + ) + .await; + assert_eq!(blocked, serde_json::json!([]), "node {idx}"); + } + + cluster.shutdown().await; +} diff --git a/nodedb-cluster-tests/tests/sql_cluster_cross_node_dml_tests/graph_rag_fusion_cross_node.rs b/nodedb-cluster-tests/tests/sql_cluster_cross_node_dml_tests/graph_rag_fusion_cross_node.rs index 8cc92d9cd..07e8e5875 100644 --- a/nodedb-cluster-tests/tests/sql_cluster_cross_node_dml_tests/graph_rag_fusion_cross_node.rs +++ b/nodedb-cluster-tests/tests/sql_cluster_cross_node_dml_tests/graph_rag_fusion_cross_node.rs @@ -195,7 +195,7 @@ fn opcode_fields() -> TextFields { collection: Some(COLL.to_string()), query_vector: Some(QUERY.iter().map(|v| *v as f32).collect()), vector_top_k: Some(3), - edge_label: Some("hop".to_string()), + edge_labels: Some(vec!["hop".to_string()]), expansion_depth: Some(2), final_top_k: Some(10), vector_k: Some(60.0), diff --git a/nodedb-cluster-tests/tests/sql_cluster_cross_node_dml_tests/native_implicit_edge_delete_cross_node.rs b/nodedb-cluster-tests/tests/sql_cluster_cross_node_dml_tests/native_implicit_edge_delete_cross_node.rs index 8020255a3..9dab58c30 100644 --- a/nodedb-cluster-tests/tests/sql_cluster_cross_node_dml_tests/native_implicit_edge_delete_cross_node.rs +++ b/nodedb-cluster-tests/tests/sql_cluster_cross_node_dml_tests/native_implicit_edge_delete_cross_node.rs @@ -176,8 +176,8 @@ async fn native_implicit_edge_delete_cleans_reverse_cross_node() { || { let ids = cluster.nodes[idx] .traversed_node_ids("GRAPH TRAVERSE IN 'nat_impl_edge_del' FROM 'hub' DEPTH 1 LABEL 'l' DIRECTION in"); - // Only the start node `hub` should remain reachable (the - // traversal includes the start node in its result). + // At most the start node `hub` remains. With every edge gone + // `hub` is absent from the graph and the result is empty. ids.iter().all(|id| id == "hub") }, ) diff --git a/nodedb-cluster/src/rpc_codec/data_plane_error.rs b/nodedb-cluster/src/rpc_codec/data_plane_error.rs index 5a99333e8..72ecc56ec 100644 --- a/nodedb-cluster/src/rpc_codec/data_plane_error.rs +++ b/nodedb-cluster/src/rpc_codec/data_plane_error.rs @@ -12,7 +12,20 @@ /// One variant per `nodedb::bridge::envelope::ErrorCode` variant; the `nodedb` /// side converts both ways with exhaustive matches, so a new code fails to /// compile there until it is mirrored here. +/// +/// The enum is recursive: `RollbackFailed` carries the code of its cause. +/// The recursive field omits its derived bounds, and the bounds below are +/// the ones its `Box` needs. #[derive(Debug, Clone, PartialEq, Eq, rkyv::Archive, rkyv::Serialize, rkyv::Deserialize)] +#[rkyv(serialize_bounds( + __S: rkyv::ser::Writer + rkyv::ser::Allocator, + __S::Error: rkyv::rancor::Source, +))] +#[rkyv(deserialize_bounds(__D::Error: rkyv::rancor::Source))] +#[rkyv(bytecheck(bounds( + __C: rkyv::validation::ArchiveContext, + __C::Error: rkyv::rancor::Source, +)))] pub enum DataPlaneErrorCode { DeadlineExceeded, RejectedConstraint { @@ -102,9 +115,13 @@ pub enum DataPlaneErrorCode { detail: String, }, /// `entry_index` is `u64` on the wire, for the same reason as `max_depth`. + /// `cause` is the code of the failed reverse write, `None` when no + /// reverse write ran. RollbackFailed { entry_index: u64, detail: String, + #[rkyv(omit_bounds)] + cause: Option>, }, OllpRetryRequired, /// `limit` is `u64` on the wire, for the same reason as `max_depth`. @@ -175,6 +192,51 @@ pub enum DataPlaneErrorCode { object: String, detail: String, }, + /// A full-text search named a field that cannot serve it (SQLSTATE + /// `42703` or `42804`, by fault). + TextColumn { + collection: String, + column: String, + fault: DataPlaneTextColumnFault, + }, + /// A computed value lies outside the range of its result type (SQLSTATE + /// `22003`). + NumericValueOutOfRange { + detail: String, + }, + /// A label write needs a node label past the partition's node-label cap + /// (SQLSTATE `54000`). `limit` is `u64` on the wire, for the same reason + /// as `max_depth`. + NodeLabelLimit { + node: String, + label: String, + limit: u64, + }, + /// Text does not parse as the column's type (SQLSTATE `22P02`). + InvalidTextRepresentation { + detail: String, + }, + /// A value of the wrong kind for the column's type (SQLSTATE `42804`). + DatatypeMismatch { + detail: String, + }, + /// Text does not parse as a timestamp (SQLSTATE `22007`). + InvalidDatetimeFormat { + detail: String, + }, + /// An instant outside the timestamp range (SQLSTATE `22008`). + DatetimeFieldOverflow { + detail: String, + }, +} + +/// Wire mirror of `nodedb_types::text_search::TextColumnFault`. +#[derive(Debug, Clone, PartialEq, Eq, rkyv::Archive, rkyv::Serialize, rkyv::Deserialize)] +pub enum DataPlaneTextColumnFault { + Undeclared, + NotText { data_type: String }, + NotAColumn, + NotIndexed, } /// Wire mirror of `nodedb::bridge::envelope::SyncHold`. @@ -193,3 +255,30 @@ pub enum DataPlaneCounterFault { IntegerOverflow, NonFinite, } + +#[cfg(test)] +mod tests { + use super::*; + + /// A nested rollback cause survives an rkyv encode and a checked decode. + #[test] + fn nested_rollback_cause_roundtrips_through_rkyv() { + let code = DataPlaneErrorCode::RollbackFailed { + entry_index: 2, + detail: "outer".into(), + cause: Some(Box::new(DataPlaneErrorCode::RollbackFailed { + entry_index: 1, + detail: "inner".into(), + cause: Some(Box::new(DataPlaneErrorCode::Internal { + detail: "commit".into(), + })), + })), + }; + let bytes = rkyv::to_bytes::(&code).expect("encode"); + let mut aligned = rkyv::util::AlignedVec::<16>::with_capacity(bytes.len()); + aligned.extend_from_slice(&bytes); + let decoded = + rkyv::from_bytes::(&aligned).expect("decode"); + assert_eq!(decoded, code); + } +} diff --git a/nodedb-cluster/src/rpc_codec/mod.rs b/nodedb-cluster/src/rpc_codec/mod.rs index 9e72c7647..83ade303b 100644 --- a/nodedb-cluster/src/rpc_codec/mod.rs +++ b/nodedb-cluster/src/rpc_codec/mod.rs @@ -49,7 +49,9 @@ pub use cluster_mgmt::{ JoinGroupInfo, JoinNodeInfo, JoinRequest, JoinResponse, LEADER_REDIRECT_PREFIX, PingRequest, PongResponse, TopologyAck, TopologyUpdate, }; -pub use data_plane_error::{DataPlaneCounterFault, DataPlaneErrorCode, DataPlaneSyncHold}; +pub use data_plane_error::{ + DataPlaneCounterFault, DataPlaneErrorCode, DataPlaneSyncHold, DataPlaneTextColumnFault, +}; pub use data_propose::{ DataProposeRequest, DataProposeResponse, ForwardedProposeRefusal, ProposeTarget, }; diff --git a/nodedb-columnar/Cargo.toml b/nodedb-columnar/Cargo.toml index 438e50f29..69d2727ca 100644 --- a/nodedb-columnar/Cargo.toml +++ b/nodedb-columnar/Cargo.toml @@ -12,6 +12,13 @@ documentation = "https://docs.rs/nodedb-columnar" keywords = ["columnar", "segment", "compression", "olap", "analytics"] categories = ["database-implementations", "compression"] +[features] +default = [] +# Corrupt memtable cells report through the `faultbox` black-box recorder. +# Inert until the host application calls `faultbox::init`. A no-op on wasm32. +# Off by default so Lite and WASM embedders do not pull the recorder. +diagnostics = ["dep:faultbox"] + [dependencies] nodedb-types = { workspace = true } nodedb-codec = { workspace = true } @@ -23,6 +30,8 @@ sonic-rs = { workspace = true } zerompk = { workspace = true } crc32c = { workspace = true } uuid = { workspace = true } +ulid = { workspace = true } roaring = { workspace = true } rust_decimal = { workspace = true } nodedb-wal = { workspace = true } +faultbox = { workspace = true, optional = true } diff --git a/nodedb-columnar/src/compaction/segment.rs b/nodedb-columnar/src/compaction/segment.rs index 94f9560aa..febcc6f90 100644 --- a/nodedb-columnar/src/compaction/segment.rs +++ b/nodedb-columnar/src/compaction/segment.rs @@ -3,13 +3,13 @@ //! Single-segment compaction: drop deleted rows from one segment, write a new one. use nodedb_mem::ScopedMemory; -use nodedb_types::columnar::ColumnarSchema; +use nodedb_types::columnar::{ColumnType, ColumnarSchema}; +use nodedb_types::value::Value; use crate::delete_bitmap::DeleteBitmap; use crate::error::ColumnarError; -use crate::materialize_rows::extract::extract_row_value; use crate::memtable::ColumnarMemtable; -use crate::reader::SegmentReader; +use crate::reader::{DecodedColumn, SegmentReader, decoded_cell_value}; use crate::writer::SegmentWriter; /// Default compaction threshold: compact when >20% of rows are deleted. @@ -81,8 +81,7 @@ pub fn compact_segment( row_values.clear(); for (col_idx, decoded) in decoded_cols.iter().enumerate() { let col = &schema.columns[col_idx]; - let value = extract_row_value(decoded, row_idx, &col.column_type, &col.name)?; - row_values.push(value); + row_values.push(carried_cell(decoded, row_idx, &col.column_type)?); } memtable.append_row(&row_values)?; @@ -99,6 +98,30 @@ pub fn compact_segment( }) } +/// One cell as the value that re-appends it to a memtable unchanged. +/// +/// A JSON cell carries its stored MessagePack bytes, which a JSON column +/// appends as is. Its decoded value does not always re-append: a JSON string +/// scalar reads as `Value::String`, which a JSON column parses as JSON text. +/// Every other cell carries its read value. +fn carried_cell( + decoded: &DecodedColumn, + row_idx: usize, + declared: &ColumnType, +) -> Result { + let value = decoded_cell_value(decoded, row_idx, declared)?; + if !matches!(declared, ColumnType::Json) || value == Value::Null { + return Ok(value); + } + let DecodedColumn::Binary { data, offsets, .. } = decoded else { + return Ok(value); + }; + // `decoded_cell_value` read this cell, so its byte range is in bounds. + let start = offsets[row_idx] as usize; + let end = offsets[row_idx + 1] as usize; + Ok(Value::Bytes(data[start..end].to_vec())) +} + #[cfg(test)] mod tests { use nodedb_types::columnar::{ColumnDef, ColumnType, ColumnarSchema}; @@ -236,4 +259,63 @@ mod tests { _ => panic!("expected Binary"), } } + + /// Compaction carries every surviving cell unchanged: identifiers, + /// decimals, booleans, geometry text and JSON, a JSON string scalar + /// included. + #[test] + fn compact_preserves_typed_cells() { + let schema = ColumnarSchema::new(vec![ + ColumnDef::required("id", ColumnType::Int64).with_primary_key(), + ColumnDef::nullable("u", ColumnType::Uuid), + ColumnDef::nullable("ul", ColumnType::Ulid), + ColumnDef::nullable("dec", ColumnType::Decimal(None)), + ColumnDef::nullable("b", ColumnType::Bool), + ColumnDef::nullable("geo", ColumnType::Geometry), + ColumnDef::nullable("j", ColumnType::Json), + ]) + .expect("valid"); + let mut mt = ColumnarMemtable::new(&schema); + for i in 0..6u128 { + let json = if i % 2 == 0 { + Value::String("\"scalar\"".into()) + } else { + Value::Array(vec![Value::Integer(i as i64)]) + }; + mt.append_row(&[ + Value::Integer(i as i64), + Value::Uuid(uuid::Uuid::from_u128(i + 1).to_string()), + Value::Ulid(ulid::Ulid(i + 1).to_string()), + Value::Decimal(rust_decimal::Decimal::new(i as i64 * 5, 1)), + Value::Bool(i % 3 == 0), + Value::String(format!("POINT({i} 2)")), + json, + ]) + .expect("append"); + } + let expected: Vec> = (1..6) + .map(|i| mt.get_row(i).expect("read").expect("row")) + .collect(); + let (schema, columns, row_count) = mt.drain(); + let segment = SegmentWriter::new(0, test_memory()) + .write_segment(&schema, &columns, row_count, None) + .expect("write"); + + let mut deletes = DeleteBitmap::new(); + deletes.mark_deleted(0); + let result = + compact_segment(&segment, &deletes, &schema, 0, &test_memory(), None).expect("compact"); + let new_seg = result.segment.as_ref().expect("segment"); + let reader = SegmentReader::open(new_seg).expect("open"); + let indices: Vec = (0..schema.columns.len()).collect(); + let decoded = reader.read_columns(&indices, &[]).expect("read"); + + for (row, want_row) in expected.iter().enumerate() { + for ((col, def), want) in decoded.iter().zip(&schema.columns).zip(want_row) { + let got = crate::reader::decoded_cell_value(col, row, &def.column_type) + .expect("decodable cell"); + assert_eq!(&got, want, "row {row} column '{}'", def.name); + } + } + } } diff --git a/nodedb-columnar/src/delete_bitmap.rs b/nodedb-columnar/src/delete_bitmap.rs index 124e2b70e..c5e8a4773 100644 --- a/nodedb-columnar/src/delete_bitmap.rs +++ b/nodedb-columnar/src/delete_bitmap.rs @@ -131,6 +131,14 @@ impl DeleteBitmap { self.inner.remove(row_idx) } + /// Remove every row index at or above `start` from the deleted set. + /// + /// Used when the rows from `start` on are cut from a memtable: an index + /// that later holds a new row must not inherit a tombstone. + pub fn unmark_from(&mut self, start: u32) { + self.inner.remove_range(start..); + } + /// Merge another bitmap into this one (union). /// /// Used when two views of the same segment's tombstones must be combined @@ -160,6 +168,18 @@ mod tests { assert_eq!(bm.deleted_count(), 1); } + #[test] + fn unmark_from_clears_the_tail_only() { + let mut bm = DeleteBitmap::new(); + bm.mark_deleted_batch(&[1, 4, 5, 9]); + bm.unmark_from(4); + assert!(bm.is_deleted(1)); + assert!(!bm.is_deleted(4)); + assert!(!bm.is_deleted(5)); + assert!(!bm.is_deleted(9)); + assert_eq!(bm.deleted_count(), 1); + } + #[test] fn batch_delete() { let mut bm = DeleteBitmap::new(); diff --git a/nodedb-columnar/src/diag/context.rs b/nodedb-columnar/src/diag/context.rs new file mode 100644 index 000000000..93aadc884 --- /dev/null +++ b/nodedb-columnar/src/diag/context.rs @@ -0,0 +1,96 @@ +// SPDX-License-Identifier: Apache-2.0 + +//! Forensic payloads carried by columnar reports. +//! +//! Grouping keys carry no row index. The row identifies the occurrence, so a +//! scan over many bad rows of one column files one report with a count. + +use faultbox::DomainContext; +use faultbox::serde_json::{Value, json}; + +/// A memtable cell whose bytes do not hold a value of its declared type. +pub(super) struct MemtableCellCorrupt<'a> { + /// Collection whose memtable holds the cell. + pub collection: &'a str, + /// Column of the cell. + pub column: &'a str, + /// The error text: column, row, and why the bytes do not decode. + pub detail: String, +} + +/// A columnar WAL row whose bytes do not decode. +pub(super) struct WalRowCorrupt { + /// The error text: the byte offset and why the bytes do not decode. + pub detail: String, +} + +impl DomainContext for WalRowCorrupt { + fn domain_kind(&self) -> &'static str { + "nodedb_columnar.wal_row_corrupt" + } + + fn grouping_key(&self) -> String { + "wal_row".to_string() + } + + fn to_json(&self) -> Value { + json!({ + "detail": self.detail, + "why_fatal": "the row the record carries cannot be rebuilt, so its decode is \ + refused", + "operator_action": "restore the collection from a snapshot taken before the \ + damaged record", + }) + } +} + +/// A string column cell whose bytes are not UTF-8 when a segment is written. +pub(super) struct StringCellNotUtf8<'a> { + /// Column of the cell. + pub column: &'a str, + /// The error text: column and row. + pub detail: String, +} + +impl DomainContext for StringCellNotUtf8<'_> { + fn domain_kind(&self) -> &'static str { + "nodedb_columnar.string_cell_not_utf8" + } + + fn grouping_key(&self) -> String { + format!("column={}", self.column) + } + + fn to_json(&self) -> Value { + json!({ + "column": self.column, + "detail": self.detail, + "why_fatal": "the segment write is refused: block statistics built from these \ + bytes would prune rows a predicate matches", + "operator_action": "rewrite or delete the named row, then let the flush retry", + }) + } +} + +impl DomainContext for MemtableCellCorrupt<'_> { + fn domain_kind(&self) -> &'static str { + "nodedb_columnar.memtable_cell_corrupt" + } + + fn grouping_key(&self) -> String { + format!("collection={} column={}", self.collection, self.column) + } + + fn to_json(&self) -> Value { + json!({ + "collection": self.collection, + "column": self.column, + "detail": self.detail, + "why_fatal": "every read of the row is refused until the row is rewritten or \ + removed. A flush of the memtable would copy the bad bytes into a \ + segment", + "operator_action": "rewrite or delete the named row. If many rows fail, \ + restore the collection from a snapshot", + }) + } +} diff --git a/nodedb-columnar/src/diag/inert.rs b/nodedb-columnar/src/diag/inert.rs new file mode 100644 index 000000000..22b8d02d7 --- /dev/null +++ b/nodedb-columnar/src/diag/inert.rs @@ -0,0 +1,17 @@ +// SPDX-License-Identifier: Apache-2.0 + +//! The non-recording implementation of the columnar report sites. +//! +//! Compiled when the `diagnostics` feature is off, and on wasm32. Every entry +//! point keeps the signature of its recording counterpart and does nothing. + +use crate::error::ColumnarError; + +#[inline] +pub fn memtable_cell_corrupt(_err: &ColumnarError, _collection: &str) {} + +#[inline] +pub fn wal_row_corrupt(_err: &ColumnarError) {} + +#[inline] +pub fn string_cell_not_utf8(_err: &ColumnarError) {} diff --git a/nodedb-columnar/src/diag/mod.rs b/nodedb-columnar/src/diag/mod.rs new file mode 100644 index 000000000..c3c6bded4 --- /dev/null +++ b/nodedb-columnar/src/diag/mod.rs @@ -0,0 +1,26 @@ +// SPDX-License-Identifier: Apache-2.0 + +//! Black-box recorder wiring for corrupt columnar state. +//! +//! A corrupt memtable cell, WAL row, or string cell is refused with an +//! error. The report filed here records the corruption at the site that +//! detects it, so it survives a restart that clears the memtable. The real implementation compiles only +//! under the `diagnostics` feature and off wasm32. Otherwise every entry +//! point is a no-op with the same signature, so call sites need no `cfg`. +//! +//! This crate never calls `faultbox::init`. The host binary owns the +//! recorder, so everything here is inert until the host initializes it. + +#[cfg(all(feature = "diagnostics", not(target_arch = "wasm32")))] +mod context; +#[cfg(all(feature = "diagnostics", not(target_arch = "wasm32")))] +mod recording; + +#[cfg(not(all(feature = "diagnostics", not(target_arch = "wasm32"))))] +mod inert; + +#[cfg(all(feature = "diagnostics", not(target_arch = "wasm32")))] +pub use recording::{memtable_cell_corrupt, string_cell_not_utf8, wal_row_corrupt}; + +#[cfg(not(all(feature = "diagnostics", not(target_arch = "wasm32"))))] +pub use inert::{memtable_cell_corrupt, string_cell_not_utf8, wal_row_corrupt}; diff --git a/nodedb-columnar/src/diag/recording.rs b/nodedb-columnar/src/diag/recording.rs new file mode 100644 index 000000000..ccce2f047 --- /dev/null +++ b/nodedb-columnar/src/diag/recording.rs @@ -0,0 +1,78 @@ +// SPDX-License-Identifier: Apache-2.0 + +//! The recording implementation of the columnar report sites. +//! +//! Each function runs only on a path that already returns the error it +//! reports. `Capture::emit` never panics and returns `None` when the host +//! never initialized the recorder, so its result is discarded. + +use faultbox::{Capture, EventKind, error_chain_of}; + +use super::context; +use crate::error::ColumnarError; + +/// Report a memtable cell whose bytes do not hold a value of its type. +/// +/// Called only from the `MutationEngine` memtable row reader, the one site +/// that detects it. The memtable writer encodes every cell, so the bytes +/// were damaged in memory or restored damaged from a checkpoint. +pub fn memtable_cell_corrupt(err: &ColumnarError, collection: &str) { + let column = match err { + ColumnarError::MemtableCellCorrupt { column, .. } => column.as_str(), + _ => "", + }; + let ctx = context::MemtableCellCorrupt { + collection, + column, + detail: err.to_string(), + }; + let _ = Capture::new( + EventKind::Corruption, + "columnar memtable cell does not decode, so the read is refused", + ) + .error_chain(error_chain_of(err)) + .domain(&ctx) + .with_backtrace() + .emit(); +} + +/// Report a columnar WAL row whose bytes do not decode. +/// +/// Called only from `decode_row_from_wal`, the one site that detects it. +pub fn wal_row_corrupt(err: &ColumnarError) { + let ctx = context::WalRowCorrupt { + detail: err.to_string(), + }; + let _ = Capture::new( + EventKind::Corruption, + "columnar WAL row does not decode, so the decode is refused", + ) + .error_chain(error_chain_of(err)) + .domain(&ctx) + .with_backtrace() + .emit(); +} + +/// Report a string column cell whose bytes are not UTF-8 at segment write. +/// +/// Called only from the block statistics builder, the one site that +/// detects it. The memtable push stores only UTF-8, so the bytes were +/// damaged in memory. +pub fn string_cell_not_utf8(err: &ColumnarError) { + let column = match err { + ColumnarError::StringCellNotUtf8 { column, .. } => column.as_str(), + _ => "", + }; + let ctx = context::StringCellNotUtf8 { + column, + detail: err.to_string(), + }; + let _ = Capture::new( + EventKind::Corruption, + "string column cell is not UTF-8, so the segment write is refused", + ) + .error_chain(error_chain_of(err)) + .domain(&ctx) + .with_backtrace() + .emit(); +} diff --git a/nodedb-columnar/src/error.rs b/nodedb-columnar/src/error.rs index 617dd02cf..56f4e926c 100644 --- a/nodedb-columnar/src/error.rs +++ b/nodedb-columnar/src/error.rs @@ -87,6 +87,33 @@ pub enum ColumnarError { offset: Option, }, + /// A memtable cell whose bytes do not hold a value of its declared type. + /// + /// The memtable writer encodes every cell it stores, so this is in-memory + /// corruption. The read is refused rather than answered with `NULL`. + #[error("memtable cell corrupt: column '{column}' row {row}: {reason}")] + MemtableCellCorrupt { + column: String, + row: usize, + reason: String, + }, + + /// A columnar WAL row whose bytes do not decode. + /// + /// `offset` is the byte position in the row where decoding stopped. The + /// decode is refused rather than answered with `NULL` or replacement + /// characters. + #[error("columnar WAL row corrupt at byte {offset}: {reason}")] + WalRowCorrupt { offset: usize, reason: String }, + + /// A string column cell whose bytes are not UTF-8 at segment write time. + /// + /// The memtable push stores only valid UTF-8, so these bytes were + /// damaged in memory. The segment write is refused rather than written + /// with wrong block statistics. + #[error("string column '{column}' row {row} holds bytes that are not UTF-8")] + StringCellNotUtf8 { column: String, row: usize }, + /// Segment is encrypted (starts with `SEGV`) but no KEK was supplied. #[error( "columnar segment is encrypted but no encryption key was provided; \ diff --git a/nodedb-columnar/src/format.rs b/nodedb-columnar/src/format.rs index 223fc5eac..a15d1f382 100644 --- a/nodedb-columnar/src/format.rs +++ b/nodedb-columnar/src/format.rs @@ -229,6 +229,32 @@ impl BlockStats { } } +/// Physical layout of every block of one column. +/// +/// The writer records the layout it encodes. The reader decodes by it, so a +/// column's codec never has to stand in for its layout. +#[derive( + Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize, ToMessagePack, FromMessagePack, +)] +#[serde(rename_all = "snake_case")] +#[repr(u8)] +#[msgpack(c_enum)] +pub enum BlockLayout { + /// `i64` values through the integer codec pipeline. + Int64 = 1, + /// `f64` values through the float codec pipeline. + Float64 = 2, + /// One bit per row: row `i` is bit `i % 8` of byte `i / 8`. + PackedBool = 3, + /// A compressed `i64` offset table, then the variable-length bytes. + VarLen = 4, + /// Cells of one width: 16-byte decimals and identifiers, or packed `f32` + /// vectors. Every row holds a full cell, a null row included. + FixedWidth = 5, + /// `i64` dictionary IDs. The strings are in [`ColumnMeta::dictionary`]. + DictIds = 6, +} + /// Metadata for a single column within the segment footer. #[derive(Debug, Clone, Serialize, Deserialize, ToMessagePack, FromMessagePack)] pub struct ColumnMeta { @@ -243,6 +269,8 @@ pub struct ColumnMeta { /// /// Always a concrete, resolved codec — never `Auto`. pub codec: nodedb_codec::ResolvedColumnCodec, + /// Physical layout of this column's blocks. + pub layout: BlockLayout, /// Number of blocks for this column. pub block_count: u32, /// Per-block statistics (one entry per block). @@ -401,6 +429,7 @@ mod tests { offset: 7, length: 512, codec: nodedb_codec::ResolvedColumnCodec::DeltaFastLanesLz4, + layout: BlockLayout::Int64, block_count: 2, block_stats: vec![ BlockStats::numeric(1.0, 1024.0, 0, 1024), @@ -413,6 +442,7 @@ mod tests { offset: 519, length: 256, codec: nodedb_codec::ResolvedColumnCodec::FsstLz4, + layout: BlockLayout::VarLen, block_count: 2, block_stats: vec![ BlockStats::non_numeric(0, 1024), @@ -425,6 +455,7 @@ mod tests { offset: 775, length: 128, codec: nodedb_codec::ResolvedColumnCodec::AlpFastLanesLz4, + layout: BlockLayout::Float64, block_count: 2, block_stats: vec![ BlockStats::numeric(0.0, 100.0, 10, 1024), @@ -452,6 +483,7 @@ mod tests { assert_eq!(parsed.columns[0].name, "id"); assert_eq!(parsed.columns[1].name, "name"); assert_eq!(parsed.columns[2].name, "score"); + assert_eq!(parsed.columns[1].layout, BlockLayout::VarLen); } #[test] @@ -520,6 +552,7 @@ mod tests { offset: 0, length: 64, codec: nodedb_codec::ResolvedColumnCodec::Lz4, + layout: BlockLayout::FixedWidth, block_count: 1, block_stats: vec![BlockStats::non_numeric(0, 128)], dictionary: None, diff --git a/nodedb-columnar/src/lib.rs b/nodedb-columnar/src/lib.rs index 07e205e3b..e7c5171de 100644 --- a/nodedb-columnar/src/lib.rs +++ b/nodedb-columnar/src/lib.rs @@ -14,6 +14,7 @@ //! [SegmentFooter: schema_hash, column metadata, block stats, CRC32C] //! ``` +pub(crate) mod diag; pub(crate) mod encrypt; pub mod compaction; @@ -41,12 +42,15 @@ pub use filter::{ dict_eval_ne, words_for, }; pub use format::{ - BLOCK_SIZE, BlockStats, BloomFilter, ColumnMeta, MAGIC, SegmentFooter, SegmentHeader, - VERSION_MAJOR, VERSION_MINOR, + BLOCK_SIZE, BlockLayout, BlockStats, BloomFilter, ColumnMeta, MAGIC, SegmentFooter, + SegmentHeader, VERSION_MAJOR, VERSION_MINOR, }; pub use materialize_rows::materialize_segment_live_rows; pub use memtable::{ColumnarMemtable, IngestValue, MemtableRowIter}; -pub use mutation::{ColumnDataSnapshot, ColumnarEngineSnapshot, MutationEngine, TruncatedRows}; +pub use mutation::{ + BatchConflict, BatchRow, ColumnDataSnapshot, ColumnarEngineSnapshot, MutationEngine, + TruncatedRows, +}; pub use pk_index::PkIndex; pub use predicate::{ BLOOM_BITS_DEFAULT, BLOOM_BYTES, BLOOM_K_DEFAULT, PredicateOp, PredicateValue, ScanPredicate, diff --git a/nodedb-columnar/src/materialize_rows/extract.rs b/nodedb-columnar/src/materialize_rows/extract.rs deleted file mode 100644 index faab06c22..000000000 --- a/nodedb-columnar/src/materialize_rows/extract.rs +++ /dev/null @@ -1,105 +0,0 @@ -// SPDX-License-Identifier: Apache-2.0 - -//! Shared helper: materialize a row value from a DecodedColumn. - -use nodedb_types::value_from_msgpack; - -use crate::error::ColumnarError; -use crate::reader::DecodedColumn; - -/// Extract a single row value from a `DecodedColumn`. -/// -/// Returns `Err(ColumnarError::MsgpackDeserialize)` if a `Json` column -/// contains bytes that cannot be decoded as MessagePack — this indicates -/// segment corruption rather than a missing value, so `Value::Null` would -/// silently hide the problem. -pub(crate) fn extract_row_value( - col: &DecodedColumn, - row_idx: usize, - col_type: &nodedb_types::columnar::ColumnType, - col_name: &str, -) -> Result { - use nodedb_types::value::Value; - - let v = match col { - // The reader infers a column's physical kind from its codec, so a - // time column decodes as `Int64`. The declared type decides the cell: - // an instant column yields the variant its declared type names; - // every other type yields the integer stored. - DecodedColumn::Int64 { values, valid } | DecodedColumn::Timestamp { values, valid } => { - if !valid[row_idx] { - Value::Null - } else { - col_type.time_cell(values[row_idx]) - } - } - DecodedColumn::Float64 { values, valid } => { - if !valid[row_idx] { - Value::Null - } else { - Value::Float(values[row_idx]) - } - } - DecodedColumn::Bool { values, valid } => { - if !valid[row_idx] { - Value::Null - } else { - Value::Bool(values[row_idx]) - } - } - DecodedColumn::Binary { - data, - offsets, - valid, - } => { - if !valid[row_idx] { - Value::Null - } else { - let start = offsets[row_idx] as usize; - let end = offsets[row_idx + 1] as usize; - let bytes = &data[start..end]; - - match col_type { - nodedb_types::columnar::ColumnType::String - | nodedb_types::columnar::ColumnType::SparseVector => { - Value::String(String::from_utf8_lossy(bytes).into_owned()) - } - nodedb_types::columnar::ColumnType::Json => { - // MessagePack-encoded JSON; decode back to structured Value. - // An empty byte slice means the JSON value was NULL at write time. - if bytes.is_empty() { - Value::Null - } else { - value_from_msgpack(bytes).map_err(|e| { - ColumnarError::MsgpackDeserialize { - column: col_name.to_string(), - source: e, - } - })? - } - } - nodedb_types::columnar::ColumnType::Bytes - | nodedb_types::columnar::ColumnType::Geometry => Value::Bytes(bytes.to_vec()), - _ => Value::Bytes(bytes.to_vec()), - } - } - } - DecodedColumn::DictEncoded { - ids, - dictionary, - valid, - } => { - if !valid[row_idx] { - Value::Null - } else { - let id = ids[row_idx] as usize; - if let Some(s) = dictionary.get(id) { - Value::String(s.clone()) - } else { - Value::Null - } - } - } - }; - Ok(v) -} diff --git a/nodedb-columnar/src/materialize_rows/mod.rs b/nodedb-columnar/src/materialize_rows/mod.rs index f91a42223..2df80fdae 100644 --- a/nodedb-columnar/src/materialize_rows/mod.rs +++ b/nodedb-columnar/src/materialize_rows/mod.rs @@ -4,10 +4,10 @@ //! //! On Origin, segments are write-once: once encoded they are only ever read, //! scanned, or (for RESTORE) decoded back into rows and re-issued through the -//! normal durable write path. The per-row decode lives here; the in-place -//! rewrite embedded deployments use is in [`crate::compaction`]. +//! normal durable write path. Each cell decodes through +//! [`crate::reader::decoded_cell_value`]. The in-place rewrite embedded +//! deployments use is in [`crate::compaction`]. -pub mod extract; pub mod rows; pub use rows::materialize_segment_live_rows; diff --git a/nodedb-columnar/src/materialize_rows/rows.rs b/nodedb-columnar/src/materialize_rows/rows.rs index f07f98484..7f785ad0d 100644 --- a/nodedb-columnar/src/materialize_rows/rows.rs +++ b/nodedb-columnar/src/materialize_rows/rows.rs @@ -14,9 +14,7 @@ use nodedb_types::value::Value; use crate::delete_bitmap::DeleteBitmap; use crate::error::ColumnarError; -use crate::reader::OwnedSegmentReader; - -use super::extract::extract_row_value; +use crate::reader::{OwnedSegmentReader, decoded_cell_value}; /// Decode the live rows of one flushed segment blob into `Value::Object`s, /// paired with their per-row cross-engine surrogate. @@ -62,7 +60,7 @@ pub fn materialize_segment_live_rows( let mut map = std::collections::HashMap::with_capacity(col_count); for (col_idx, decoded) in decoded_cols.iter().enumerate() { let col = &schema.columns[col_idx]; - let value = extract_row_value(decoded, row_idx, &col.column_type, &col.name)?; + let value = decoded_cell_value(decoded, row_idx, &col.column_type)?; map.insert(col.name.clone(), value); } let surrogate = row_surrogates.get(row_idx).copied().flatten(); diff --git a/nodedb-columnar/src/memtable/column_data/access.rs b/nodedb-columnar/src/memtable/column_data/access.rs index 2d6a769ad..bf0b632a8 100644 --- a/nodedb-columnar/src/memtable/column_data/access.rs +++ b/nodedb-columnar/src/memtable/column_data/access.rs @@ -7,6 +7,7 @@ use nodedb_types::value::Value; use nodedb_types::value_from_msgpack; use super::types::ColumnData; +use crate::error::ColumnarError; impl ColumnData { /// Get the validity bitmap, or generate an all-true one for non-nullable columns. @@ -64,11 +65,26 @@ impl ColumnData { /// (`Timestamptz`) from the stored epoch microseconds, and every other /// declared type backed by time storage (`SystemTimestamp`, `Duration`) /// yields the integer stored. - pub(crate) fn get_value(&self, row: usize, declared: &ColumnType) -> Value { + /// + /// Returns `MemtableCellCorrupt` for a cell whose bytes do not hold a + /// value of its type: a JSON cell that is not MessagePack, a text cell + /// that is not UTF-8, or a dictionary ID outside the dictionary. + /// `column` names the column in that error. + pub(crate) fn get_value( + &self, + row: usize, + declared: &ColumnType, + column: &str, + ) -> Result { if self.is_null(row) { - return Value::Null; + return Ok(Value::Null); } - match self { + let corrupt = |reason: String| ColumnarError::MemtableCellCorrupt { + column: column.to_string(), + row, + reason, + }; + let value = match self { Self::Int64 { values, .. } => Value::Integer(values[row]), Self::Float64 { values, .. } => Value::Float(values[row]), Self::Bool { values, .. } => Value::Bool(values[row]), @@ -76,16 +92,17 @@ impl ColumnData { Self::Decimal { values, .. } => { Value::Decimal(rust_decimal::Decimal::deserialize(values[row])) } - Self::Uuid { values, .. } => { - Value::Uuid(uuid::Uuid::from_bytes(values[row]).to_string()) - } + Self::Uuid { values, .. } => match declared { + ColumnType::Ulid => Value::Ulid(ulid::Ulid::from_bytes(values[row]).to_string()), + _ => Value::Uuid(uuid::Uuid::from_bytes(values[row]).to_string()), + }, Self::String { data, offsets, .. } => { let start = offsets[row] as usize; let end = offsets[row + 1] as usize; - let s = std::str::from_utf8(&data[start..end]) - .unwrap_or("") - .to_string(); - Value::String(s) + Value::String( + utf8_cell(&data[start..end]) + .map_err(|e| corrupt(format!("text cell is not UTF-8: {e}")))?, + ) } Self::Bytes { data, offsets, .. } => { let start = offsets[row] as usize; @@ -99,16 +116,17 @@ impl ColumnData { if slice.is_empty() { Value::Null } else { - value_from_msgpack(slice).unwrap_or(Value::Null) + value_from_msgpack(slice) + .map_err(|e| corrupt(format!("JSON cell is not MessagePack: {e}")))? } } Self::Geometry { data, offsets, .. } => { let start = offsets[row] as usize; let end = offsets[row + 1] as usize; - let s = std::str::from_utf8(&data[start..end]) - .unwrap_or("") - .to_string(); - Value::String(s) + Value::String( + utf8_cell(&data[start..end]) + .map_err(|e| corrupt(format!("text cell is not UTF-8: {e}")))?, + ) } Self::Vector { data, dim, .. } => { let d = *dim as usize; @@ -122,13 +140,97 @@ impl ColumnData { Self::DictEncoded { ids, dictionary, .. } => { - let id = ids[row] as usize; - if id < dictionary.len() { - Value::String(dictionary[id].clone()) - } else { - Value::Null - } + let id = ids[row]; + let text = usize::try_from(id) + .ok() + .and_then(|i| dictionary.get(i)) + .ok_or_else(|| { + corrupt(format!( + "dictionary ID {id} is outside a dictionary of {} entries", + dictionary.len() + )) + })?; + Value::String(text.clone()) } + }; + Ok(value) + } +} + +/// The text of a string or geometry cell, or the UTF-8 error of its bytes. +fn utf8_cell(bytes: &[u8]) -> Result { + std::str::from_utf8(bytes).map(str::to_owned) +} + +#[cfg(test)] +mod tests { + use nodedb_types::columnar::ColumnType; + use nodedb_types::value::Value; + + use super::super::types::ColumnData; + use crate::error::ColumnarError; + + fn corrupt_reason(result: Result) -> String { + match result { + Err(ColumnarError::MemtableCellCorrupt { + column, + row, + reason, + }) => { + assert_eq!(column, "c"); + assert_eq!(row, 0); + reason + } + other => panic!("expected MemtableCellCorrupt, got {other:?}"), } } + + #[test] + fn a_json_cell_that_is_not_msgpack_is_refused_not_null() { + // 0xC1 is the one MessagePack marker that is never valid. + let col = ColumnData::Json { + data: vec![0xC1], + offsets: vec![0, 1], + valid: None, + }; + let reason = corrupt_reason(col.get_value(0, &ColumnType::Json, "c")); + assert!(reason.contains("MessagePack"), "{reason}"); + } + + #[test] + fn a_valid_json_cell_reads_back() { + let bytes = nodedb_types::value_to_msgpack(&Value::Integer(7)).expect("encode"); + let col = ColumnData::Json { + offsets: vec![0, bytes.len() as u32], + data: bytes, + valid: None, + }; + assert_eq!( + col.get_value(0, &ColumnType::Json, "c").expect("read"), + Value::Integer(7) + ); + } + + #[test] + fn a_text_cell_that_is_not_utf8_is_refused() { + let col = ColumnData::String { + data: vec![0xFF, 0xFE], + offsets: vec![0, 2], + valid: None, + }; + let reason = corrupt_reason(col.get_value(0, &ColumnType::String, "c")); + assert!(reason.contains("UTF-8"), "{reason}"); + } + + #[test] + fn a_dictionary_id_outside_the_dictionary_is_refused() { + let col = ColumnData::DictEncoded { + ids: vec![3], + dictionary: vec!["a".into()], + reverse: std::collections::HashMap::new(), + valid: None, + }; + let reason = corrupt_reason(col.get_value(0, &ColumnType::String, "c")); + assert!(reason.contains("dictionary ID 3"), "{reason}"); + } } diff --git a/nodedb-columnar/src/memtable/column_data/mod.rs b/nodedb-columnar/src/memtable/column_data/mod.rs index e3ed1f903..47301bb69 100644 --- a/nodedb-columnar/src/memtable/column_data/mod.rs +++ b/nodedb-columnar/src/memtable/column_data/mod.rs @@ -4,6 +4,7 @@ mod access; mod backfill; mod dict_encode; mod push; +mod push_ref; mod truncate; mod types; diff --git a/nodedb-columnar/src/memtable/column_data/push.rs b/nodedb-columnar/src/memtable/column_data/push.rs index ab7885f19..0d962a474 100644 --- a/nodedb-columnar/src/memtable/column_data/push.rs +++ b/nodedb-columnar/src/memtable/column_data/push.rs @@ -1,16 +1,37 @@ // SPDX-License-Identifier: Apache-2.0 -//! Append operations on `ColumnData`: push owned values and push borrowed values. +//! Append owned values on `ColumnData`. +//! +//! A column accepts every shape the strict document coercion yields for its +//! declared type, so columnar and strict collections hold the same values. use nodedb_types::columnar::ColumnType; use nodedb_types::value::Value; -use nodedb_types::value_to_msgpack; +use nodedb_types::{value_from_msgpack, value_to_msgpack}; use crate::error::ColumnarError; -use super::super::IngestValue; use super::types::ColumnData; +/// The 16 stored bytes of an identifier column cell. +/// +/// A `Ulid` column parses ULID text. Every other column backed by 16-byte +/// identifier storage parses UUID text. Text that does not parse is refused. +fn parse_id_bytes( + s: &str, + col_name: &str, + col_type: &ColumnType, +) -> Result<[u8; 16], ColumnarError> { + let parsed = match col_type { + ColumnType::Ulid => ulid::Ulid::from_string(s).ok().map(|u| u.to_bytes()), + _ => uuid::Uuid::parse_str(s).ok().map(|u| *u.as_bytes()), + }; + parsed.ok_or_else(|| ColumnarError::TypeMismatch { + column: col_name.to_string(), + expected: col_type.to_string(), + }) +} + /// Encode a `Value` as MessagePack bytes for JSON/Array/Set/Record storage. /// /// For `Value::String` input, the string is first parsed as JSON so that @@ -170,11 +191,12 @@ impl ColumnData { values.push(d.serialize()); Self::push_valid(valid, true); } - (Self::Uuid { values, valid }, Value::Uuid(s)) => { - let bytes = uuid::Uuid::parse_str(s) - .map(|u| *u.as_bytes()) - .unwrap_or([0u8; 16]); - values.push(bytes); + (Self::Uuid { values, valid }, Value::Uuid(s) | Value::Ulid(s) | Value::String(s)) => { + values.push(parse_id_bytes(s, col_name, col_type)?); + Self::push_valid(valid, true); + } + (Self::Timestamp { values, valid }, Value::Duration(d)) => { + values.push(d.micros); Self::push_valid(valid, true); } ( @@ -183,7 +205,7 @@ impl ColumnData { offsets, valid, }, - Value::String(s), + Value::String(s) | Value::Uuid(s) | Value::Ulid(s) | Value::Regex(s), ) => { data.extend_from_slice(s.as_bytes()); offsets.push(data.len() as u32); @@ -245,7 +267,12 @@ impl ColumnData { }, Value::Bytes(b), ) => { - // Assume already MessagePack-encoded bytes. + // Bytes are stored as an encoded MessagePack cell. Bytes that + // do not decode are refused, so every stored JSON cell reads. + value_from_msgpack(b).map_err(|e| ColumnarError::MsgpackDeserialize { + column: col_name.to_string(), + source: e, + })?; data.extend_from_slice(b); offsets.push(data.len() as u32); Self::push_valid(valid, true); @@ -307,6 +334,18 @@ impl ColumnData { } Self::push_valid(valid, true); } + // Packed little-endian `f32`s, the form strict coercion yields. + (Self::Vector { data, dim, valid }, Value::Bytes(b)) + if b.len() == *dim as usize * 4 => + { + data.extend( + b.as_chunks::<4>() + .0 + .iter() + .map(|chunk| f32::from_le_bytes(*chunk)), + ); + Self::push_valid(valid, true); + } (Self::DictEncoded { ids, valid, .. }, Value::Null) => { ids.push(0); Self::push_valid(valid, false); @@ -318,7 +357,7 @@ impl ColumnData { reverse, valid, }, - Value::String(s), + Value::String(s) | Value::Uuid(s) | Value::Ulid(s) | Value::Regex(s), ) => { let id = if let Some(&existing) = reverse.get(s.as_str()) { existing @@ -331,147 +370,118 @@ impl ColumnData { ids.push(id); Self::push_valid(valid, true); } - (other, val) => { - let type_name = match other { - Self::Int64 { .. } => "Int64", - Self::Float64 { .. } => "Float64", - Self::Bool { .. } => "Bool", - Self::Timestamp { .. } => "Timestamp", - Self::Decimal { .. } => "Decimal", - Self::Uuid { .. } => "Uuid", - Self::String { .. } => "String", - Self::Bytes { .. } => "Bytes", - Self::Json { .. } => "Json", - Self::Geometry { .. } => "Geometry", - Self::Vector { .. } => "Vector", - Self::DictEncoded { .. } => "DictEncoded", - }; - let _ = val; + (other, _) => { return Err(ColumnarError::TypeMismatch { column: col_name.to_string(), - expected: type_name.to_string(), + expected: other.type_name().to_string(), }); } } Ok(()) } +} - /// Append a borrowed value (zero-copy for strings). Used by `ingest_row_refs`. - pub(crate) fn push_ref( - &mut self, - value: &IngestValue<'_>, - col_name: &str, - ) -> Result<(), ColumnarError> { - match (self, value) { - (Self::Int64 { values, valid }, IngestValue::Null) => { - values.push(0); - Self::push_valid(valid, false); - } - (Self::Float64 { values, valid }, IngestValue::Null) => { - values.push(0.0); - Self::push_valid(valid, false); - } - (Self::Bool { values, valid }, IngestValue::Null) => { - values.push(false); - Self::push_valid(valid, false); - } - (Self::Timestamp { values, valid }, IngestValue::Null) => { - values.push(0); - Self::push_valid(valid, false); - } - (Self::String { offsets, valid, .. }, IngestValue::Null) => { - offsets.push(*offsets.last().unwrap_or(&0)); - Self::push_valid(valid, false); - } - (Self::Bytes { offsets, valid, .. }, IngestValue::Null) => { - offsets.push(*offsets.last().unwrap_or(&0)); - Self::push_valid(valid, false); - } - (Self::Json { offsets, valid, .. }, IngestValue::Null) => { - offsets.push(*offsets.last().unwrap_or(&0)); - Self::push_valid(valid, false); - } - (Self::DictEncoded { ids, valid, .. }, IngestValue::Null) => { - ids.push(0); - Self::push_valid(valid, false); - } - (Self::Int64 { values, valid }, IngestValue::Int64(v)) => { - values.push(*v); - Self::push_valid(valid, true); - } - (Self::Float64 { values, valid }, IngestValue::Float64(v)) => { - values.push(*v); - Self::push_valid(valid, true); - } - (Self::Float64 { values, valid }, IngestValue::Int64(v)) => { - values.push(*v as f64); - Self::push_valid(valid, true); - } - (Self::Bool { values, valid }, IngestValue::Bool(v)) => { - values.push(*v); - Self::push_valid(valid, true); - } - (Self::Timestamp { values, valid }, IngestValue::Timestamp(v)) => { - values.push(*v); - Self::push_valid(valid, true); - } - (Self::Timestamp { values, valid }, IngestValue::Int64(v)) => { - values.push(*v); - Self::push_valid(valid, true); - } - ( - Self::String { - data, - offsets, - valid, - }, - IngestValue::Str(s), - ) => { - data.extend_from_slice(s.as_bytes()); - offsets.push(data.len() as u32); - Self::push_valid(valid, true); - } - ( - Self::DictEncoded { - ids, - dictionary, - reverse, - valid, - }, - IngestValue::Str(s), - ) => { - let id = if let Some(&existing) = reverse.get(*s) { - existing - } else { - let new_id = dictionary.len() as u32; - dictionary.push((*s).to_string()); - reverse.insert((*s).to_string(), new_id); - new_id - }; - ids.push(id); - Self::push_valid(valid, true); - } - (other, _) => { - let type_name = match other { - Self::Int64 { .. } => "Int64", - Self::Float64 { .. } => "Float64", - Self::Bool { .. } => "Bool", - Self::Timestamp { .. } => "Timestamp", - Self::Decimal { .. } => "Decimal", - Self::Uuid { .. } => "Uuid", - Self::String { .. } => "String", - Self::Bytes { .. } => "Bytes", - Self::Json { .. } => "Json", - Self::Geometry { .. } => "Geometry", - Self::Vector { .. } => "Vector", - Self::DictEncoded { .. } => "DictEncoded", - }; - return Err(ColumnarError::TypeMismatch { - column: col_name.to_string(), - expected: type_name.to_string(), - }); - } - } - Ok(()) +#[cfg(test)] +mod tests { + use nodedb_types::NdbDuration; + use nodedb_types::columnar::{ColumnDef, ColumnarSchema}; + + use super::*; + use crate::memtable::ColumnarMemtable; + + fn memtable(col_type: ColumnType) -> ColumnarMemtable { + let schema = + ColumnarSchema::new(vec![ColumnDef::required("c", col_type)]).expect("valid schema"); + ColumnarMemtable::new(&schema) + } + + #[test] + fn uuid_column_stores_uuid_text() { + let mut mt = memtable(ColumnType::Uuid); + let text = "67e55044-10b1-426f-9247-bb680e5fe0c8"; + mt.append_row(&[Value::String(text.into())]) + .expect("append"); + assert_eq!( + mt.get_row(0).expect("read"), + Some(vec![Value::Uuid(text.into())]) + ); + } + + #[test] + fn uuid_column_refuses_text_that_is_not_a_uuid() { + let mut mt = memtable(ColumnType::Uuid); + let err = mt + .append_row(&[Value::Uuid("not-a-uuid".into())]) + .unwrap_err(); + assert!(matches!(err, ColumnarError::TypeMismatch { ref column, .. } if column == "c")); + assert_eq!(mt.row_count(), 0); + } + + #[test] + fn ulid_column_round_trips_ulid_text() { + let mut mt = memtable(ColumnType::Ulid); + let text = "01ARZ3NDEKTSV4RRFFQ69G5FAV"; + mt.append_row(&[Value::String(text.into())]) + .expect("append"); + assert_eq!( + mt.get_row(0).expect("read"), + Some(vec![Value::Ulid(text.into())]) + ); + } + + #[test] + fn string_column_accepts_identifier_text() { + let mut mt = memtable(ColumnType::String); + mt.append_row(&[Value::Uuid("u".into())]).expect("append"); + assert_eq!( + mt.get_row(0).expect("read"), + Some(vec![Value::String("u".into())]) + ); + } + + #[test] + fn json_column_refuses_bytes_that_are_not_msgpack() { + let mut mt = memtable(ColumnType::Json); + let err = mt.append_row(&[Value::Bytes(vec![0xC1])]).unwrap_err(); + assert!( + matches!(err, ColumnarError::MsgpackDeserialize { ref column, .. } if column == "c") + ); + assert_eq!(mt.row_count(), 0); + + let encoded = nodedb_types::value_to_msgpack(&Value::Integer(9)).expect("encode"); + mt.append_row(&[Value::Bytes(encoded)]).expect("append"); + assert_eq!(mt.get_row(0).expect("read"), Some(vec![Value::Integer(9)])); + } + + #[test] + fn duration_column_stores_micros() { + let mut mt = memtable(ColumnType::Duration); + mt.append_row(&[Value::Duration(NdbDuration::from_micros(1_500))]) + .expect("append"); + assert_eq!( + mt.get_row(0).expect("read"), + Some(vec![Value::Integer(1_500)]) + ); + } + + #[test] + fn vector_column_accepts_packed_floats() { + let mut mt = memtable(ColumnType::Vector(2)); + let packed: Vec = [0.5f32, 1.25f32] + .iter() + .flat_map(|f| f.to_le_bytes()) + .collect(); + mt.append_row(&[Value::Bytes(packed)]).expect("append"); + assert_eq!( + mt.get_row(0).expect("read"), + Some(vec![Value::Array(vec![ + Value::Float(0.5), + Value::Float(1.25) + ])]) + ); + + let err = mt.append_row(&[Value::Bytes(vec![0; 3])]).unwrap_err(); + assert!(matches!(err, ColumnarError::TypeMismatch { .. })); + assert_eq!(mt.row_count(), 1); } } diff --git a/nodedb-columnar/src/memtable/column_data/push_ref.rs b/nodedb-columnar/src/memtable/column_data/push_ref.rs new file mode 100644 index 000000000..341a1049a --- /dev/null +++ b/nodedb-columnar/src/memtable/column_data/push_ref.rs @@ -0,0 +1,115 @@ +// SPDX-License-Identifier: Apache-2.0 + +//! Append borrowed values on `ColumnData` for zero-copy ingest. + +use crate::error::ColumnarError; + +use super::super::IngestValue; +use super::types::ColumnData; + +impl ColumnData { + /// Append a borrowed value (zero-copy for strings). Used by `ingest_row_refs`. + pub(crate) fn push_ref( + &mut self, + value: &IngestValue<'_>, + col_name: &str, + ) -> Result<(), ColumnarError> { + match (self, value) { + (Self::Int64 { values, valid }, IngestValue::Null) => { + values.push(0); + Self::push_valid(valid, false); + } + (Self::Float64 { values, valid }, IngestValue::Null) => { + values.push(0.0); + Self::push_valid(valid, false); + } + (Self::Bool { values, valid }, IngestValue::Null) => { + values.push(false); + Self::push_valid(valid, false); + } + (Self::Timestamp { values, valid }, IngestValue::Null) => { + values.push(0); + Self::push_valid(valid, false); + } + (Self::String { offsets, valid, .. }, IngestValue::Null) => { + offsets.push(*offsets.last().unwrap_or(&0)); + Self::push_valid(valid, false); + } + (Self::Bytes { offsets, valid, .. }, IngestValue::Null) => { + offsets.push(*offsets.last().unwrap_or(&0)); + Self::push_valid(valid, false); + } + (Self::Json { offsets, valid, .. }, IngestValue::Null) => { + offsets.push(*offsets.last().unwrap_or(&0)); + Self::push_valid(valid, false); + } + (Self::DictEncoded { ids, valid, .. }, IngestValue::Null) => { + ids.push(0); + Self::push_valid(valid, false); + } + (Self::Int64 { values, valid }, IngestValue::Int64(v)) => { + values.push(*v); + Self::push_valid(valid, true); + } + (Self::Float64 { values, valid }, IngestValue::Float64(v)) => { + values.push(*v); + Self::push_valid(valid, true); + } + (Self::Float64 { values, valid }, IngestValue::Int64(v)) => { + values.push(*v as f64); + Self::push_valid(valid, true); + } + (Self::Bool { values, valid }, IngestValue::Bool(v)) => { + values.push(*v); + Self::push_valid(valid, true); + } + (Self::Timestamp { values, valid }, IngestValue::Timestamp(v)) => { + values.push(*v); + Self::push_valid(valid, true); + } + (Self::Timestamp { values, valid }, IngestValue::Int64(v)) => { + values.push(*v); + Self::push_valid(valid, true); + } + ( + Self::String { + data, + offsets, + valid, + }, + IngestValue::Str(s), + ) => { + data.extend_from_slice(s.as_bytes()); + offsets.push(data.len() as u32); + Self::push_valid(valid, true); + } + ( + Self::DictEncoded { + ids, + dictionary, + reverse, + valid, + }, + IngestValue::Str(s), + ) => { + let id = if let Some(&existing) = reverse.get(*s) { + existing + } else { + let new_id = dictionary.len() as u32; + dictionary.push((*s).to_string()); + reverse.insert((*s).to_string(), new_id); + new_id + }; + ids.push(id); + Self::push_valid(valid, true); + } + (other, _) => { + return Err(ColumnarError::TypeMismatch { + column: col_name.to_string(), + expected: other.type_name().to_string(), + }); + } + } + Ok(()) + } +} diff --git a/nodedb-columnar/src/memtable/column_data/types.rs b/nodedb-columnar/src/memtable/column_data/types.rs index d0f0e5f03..23571dbd7 100644 --- a/nodedb-columnar/src/memtable/column_data/types.rs +++ b/nodedb-columnar/src/memtable/column_data/types.rs @@ -214,4 +214,22 @@ impl ColumnData { Self::DictEncoded { ids, .. } => ids.len(), } } + + /// Name of the storage kind, for type-mismatch errors. + pub(crate) fn type_name(&self) -> &'static str { + match self { + Self::Int64 { .. } => "Int64", + Self::Float64 { .. } => "Float64", + Self::Bool { .. } => "Bool", + Self::Timestamp { .. } => "Timestamp", + Self::Decimal { .. } => "Decimal", + Self::Uuid { .. } => "Uuid", + Self::String { .. } => "String", + Self::Bytes { .. } => "Bytes", + Self::Json { .. } => "Json", + Self::Geometry { .. } => "Geometry", + Self::Vector { .. } => "Vector", + Self::DictEncoded { .. } => "DictEncoded", + } + } } diff --git a/nodedb-columnar/src/memtable/iter.rs b/nodedb-columnar/src/memtable/iter.rs index d411a6b6b..99b15d574 100644 --- a/nodedb-columnar/src/memtable/iter.rs +++ b/nodedb-columnar/src/memtable/iter.rs @@ -7,6 +7,7 @@ use nodedb_types::value::Value; use super::column_data::ColumnData; use super::core::ColumnarMemtable; +use crate::error::ColumnarError; impl ColumnarMemtable { /// Iterate rows as `Vec`. For scan/read operations. @@ -23,18 +24,30 @@ impl ColumnarMemtable { } /// Get a single row by index as `Vec`. - pub fn get_row(&self, row_idx: usize) -> Option> { + /// + /// `Ok(None)` when `row_idx` is past the last row. `Err` when a cell of + /// the row is corrupt. + pub fn get_row(&self, row_idx: usize) -> Result>, ColumnarError> { if row_idx >= self.row_count { - return None; + return Ok(None); } - let mut row = Vec::with_capacity(self.columns.len()); - for (col, def) in self.columns.iter().zip(&self.schema.columns) { - row.push(col.get_value(row_idx, &def.column_type)); - } - Some(row) + read_row(&self.columns, &self.schema.columns, row_idx).map(Some) } } +/// Read row `row_idx` of `columns`, typing each cell by its column def. +fn read_row( + columns: &[ColumnData], + column_defs: &[ColumnDef], + row_idx: usize, +) -> Result, ColumnarError> { + columns + .iter() + .zip(column_defs) + .map(|(col, def)| col.get_value(row_idx, &def.column_type, &def.name)) + .collect() +} + /// Row iterator over a columnar memtable. pub struct MemtableRowIter<'a> { columns: &'a [ColumnData], @@ -43,17 +56,16 @@ pub struct MemtableRowIter<'a> { current: usize, } +/// Yields `Err` for a row with a corrupt cell. The iterator still advances +/// past that row. impl Iterator for MemtableRowIter<'_> { - type Item = Vec; + type Item = Result, ColumnarError>; fn next(&mut self) -> Option { if self.current >= self.row_count { return None; } - let mut row = Vec::with_capacity(self.columns.len()); - for (col, def) in self.columns.iter().zip(self.column_defs) { - row.push(col.get_value(self.current, &def.column_type)); - } + let row = read_row(self.columns, self.column_defs, self.current); self.current += 1; Some(row) } @@ -104,8 +116,9 @@ mod tests { Value::DateTime(dt), Value::Integer(7), ]; - assert_eq!(mt.get_row(0), Some(expected.clone())); - assert_eq!(mt.iter_rows().collect::>(), vec![expected]); - assert_eq!(mt.get_row(1), None); + assert_eq!(mt.get_row(0).expect("read"), Some(expected.clone())); + let rows: Vec> = mt.iter_rows().collect::>().expect("read"); + assert_eq!(rows, vec![expected]); + assert_eq!(mt.get_row(1).expect("read"), None); } } diff --git a/nodedb-columnar/src/memtable/mutation.rs b/nodedb-columnar/src/memtable/mutation.rs index d4189016a..1a7c176f2 100644 --- a/nodedb-columnar/src/memtable/mutation.rs +++ b/nodedb-columnar/src/memtable/mutation.rs @@ -13,6 +13,11 @@ use super::ingest_value::IngestValue; impl ColumnarMemtable { /// Append a row of values. Validates types and nullability. + /// + /// The append is all-or-nothing. On an error, every column is cut back + /// to `row_count`, so the memtable holds exactly the rows it held before + /// the call. A rolled-back push to a `DictEncoded` column can leave one + /// dictionary entry that no row references. Every id stays valid. pub fn append_row(&mut self, values: &[Value]) -> Result<(), ColumnarError> { if values.len() != self.schema.columns.len() { return Err(ColumnarError::SchemaMismatch { @@ -21,11 +26,21 @@ impl ColumnarMemtable { }); } - for (i, (col_def, value)) in self.schema.columns.iter().zip(values.iter()).enumerate() { - if matches!(value, Value::Null) && !col_def.nullable { - return Err(ColumnarError::NullViolation(col_def.name.clone())); - } - self.columns[i].push(value, &col_def.name, &col_def.column_type)?; + let pushed = self + .schema + .columns + .iter() + .zip(values.iter()) + .zip(self.columns.iter_mut()) + .try_for_each(|((col_def, value), column)| { + if matches!(value, Value::Null) && !col_def.nullable { + return Err(ColumnarError::NullViolation(col_def.name.clone())); + } + column.push(value, &col_def.name, &col_def.column_type) + }); + if let Err(e) = pushed { + self.rollback_partial_row(); + return Err(e); } self.row_count += 1; @@ -36,6 +51,21 @@ impl ColumnarMemtable { Ok(()) } + /// Cut every column back to `row_count` after a failed row push. + /// + /// The columns before the failed one hold one extra value. The failed + /// column and the columns after it hold `row_count` values already. + fn rollback_partial_row(&mut self) { + let n = self.row_count; + for col in &mut self.columns { + col.truncate(n); + } + debug_assert!( + self.columns.iter().all(|c| c.len() == self.row_count), + "column lengths must stay aligned with row_count after rollback" + ); + } + /// Convert low-cardinality `String` columns to `DictEncoded` in-place. pub fn try_dict_encode_columns(&mut self, max_cardinality: u32) { for col in &mut self.columns { @@ -90,6 +120,7 @@ impl ColumnarMemtable { /// /// Accepts borrowed values via `IngestValue<'_>`, avoiding string cloning /// for tag columns that are already interned in the `DictEncoded` dictionary. + /// The ingest is all-or-nothing, as in [`Self::append_row`]. pub fn ingest_row_refs(&mut self, values: &[IngestValue<'_>]) -> Result<(), ColumnarError> { if values.len() != self.schema.columns.len() { return Err(ColumnarError::SchemaMismatch { @@ -98,11 +129,21 @@ impl ColumnarMemtable { }); } - for (i, (col_def, value)) in self.schema.columns.iter().zip(values.iter()).enumerate() { - if matches!(value, IngestValue::Null) && !col_def.nullable { - return Err(ColumnarError::NullViolation(col_def.name.clone())); - } - self.columns[i].push_ref(value, &col_def.name)?; + let pushed = self + .schema + .columns + .iter() + .zip(values.iter()) + .zip(self.columns.iter_mut()) + .try_for_each(|((col_def, value), column)| { + if matches!(value, IngestValue::Null) && !col_def.nullable { + return Err(ColumnarError::NullViolation(col_def.name.clone())); + } + column.push_ref(value, &col_def.name) + }); + if let Err(e) = pushed { + self.rollback_partial_row(); + return Err(e); } self.row_count += 1; @@ -155,6 +196,87 @@ mod tests { assert!(matches!(err, ColumnarError::NullViolation(ref s) if s == "id")); } + #[test] + fn null_violation_mid_row_leaves_columns_aligned() { + let schema = test_schema(); + let mut mt = ColumnarMemtable::new(&schema); + + let err = mt + .append_row(&[Value::Integer(1), Value::Null, Value::Null]) + .unwrap_err(); + assert!(matches!(err, ColumnarError::NullViolation(ref s) if s == "name")); + assert_eq!(mt.row_count(), 0); + assert!(mt.columns().iter().all(|c| c.len() == 0)); + + mt.append_row(&[Value::Integer(2), Value::String("b".into()), Value::Null]) + .expect("append after rejected row"); + assert_eq!(mt.row_count(), 1); + assert!(mt.columns().iter().all(|c| c.len() == 1)); + assert_eq!( + mt.get_row(0).expect("read"), + Some(vec![ + Value::Integer(2), + Value::String("b".into()), + Value::Null + ]) + ); + } + + #[test] + fn type_mismatch_mid_row_leaves_columns_aligned() { + let schema = test_schema(); + let mut mt = ColumnarMemtable::new(&schema); + mt.append_row(&[ + Value::Integer(1), + Value::String("a".into()), + Value::Float(0.5), + ]) + .expect("append"); + + let err = mt + .append_row(&[ + Value::Integer(2), + Value::String("b".into()), + Value::Bool(true), + ]) + .unwrap_err(); + assert!(matches!(err, ColumnarError::TypeMismatch { ref column, .. } if column == "score")); + assert_eq!(mt.row_count(), 1); + assert!(mt.columns().iter().all(|c| c.len() == 1)); + + mt.append_row(&[ + Value::Integer(3), + Value::String("c".into()), + Value::Float(0.25), + ]) + .expect("append after rejected row"); + assert_eq!( + mt.get_row(1).expect("read"), + Some(vec![ + Value::Integer(3), + Value::String("c".into()), + Value::Float(0.25) + ]) + ); + } + + #[test] + fn rejected_ref_ingest_leaves_columns_aligned() { + let schema = test_schema(); + let mut mt = ColumnarMemtable::new(&schema); + + let err = mt + .ingest_row_refs(&[ + IngestValue::Int64(1), + IngestValue::Str("a"), + IngestValue::Bool(true), + ]) + .unwrap_err(); + assert!(matches!(err, ColumnarError::TypeMismatch { .. })); + assert_eq!(mt.row_count(), 0); + assert!(mt.columns().iter().all(|c| c.len() == 0)); + } + #[test] fn schema_mismatch_rejected() { let schema = test_schema(); diff --git a/nodedb-columnar/src/mutation/batch.rs b/nodedb-columnar/src/mutation/batch.rs new file mode 100644 index 000000000..7180ce88d --- /dev/null +++ b/nodedb-columnar/src/mutation/batch.rs @@ -0,0 +1,220 @@ +// SPDX-License-Identifier: Apache-2.0 + +//! All-or-nothing multi-row insert. + +use std::collections::HashMap; + +use nodedb_types::surrogate::Surrogate; +use nodedb_types::value::Value; + +use crate::error::ColumnarError; +use crate::pk_index::RowLocation; + +use super::engine::{MutationEngine, MutationResult}; + +/// One row of a batch insert. +pub struct BatchRow<'a> { + pub values: &'a [Value], + pub surrogate: Option, +} + +/// What a batch insert does with a row whose primary key is already bound. +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum BatchConflict { + /// Replace the bound row, as [`MutationEngine::insert`] does. + Upsert, + /// Skip the row, as [`MutationEngine::insert_if_absent`] does. + Skip, +} + +/// The PK binding a batch found before it first wrote that PK. +struct PriorBinding { + location: Option, + /// Whether the batch tombstoned `location`. + tombstoned: bool, +} + +impl MutationEngine { + /// Insert every row of `rows`, or none of them. + /// + /// Each row follows the single-row rules for `conflict`. When a row + /// fails, the rows this call already wrote are undone: the memtable is + /// cut back, their PK bindings are restored, the rows they replaced are + /// live again, and the error is returned. + /// + /// Returns one `MutationResult` per row, in order. A skipped row has an + /// empty `wal_records`. + pub fn insert_batch<'a>( + &mut self, + rows: impl IntoIterator>, + conflict: BatchConflict, + ) -> Result, ColumnarError> { + let start = self.memtable.row_count(); + let mut priors: HashMap, PriorBinding> = HashMap::new(); + let mut results = Vec::new(); + for row in rows { + match self.insert_batch_row(&row, conflict, &mut priors) { + Ok(result) => results.push(result), + Err(e) => { + self.undo_batch(start, priors); + return Err(e); + } + } + } + Ok(results) + } + + /// Write one batch row and record the PK binding it replaced. + fn insert_batch_row( + &mut self, + row: &BatchRow<'_>, + conflict: BatchConflict, + priors: &mut HashMap, PriorBinding>, + ) -> Result { + let pk_bytes = self.extract_pk_bytes(row.values)?; + let prior = self.pk_index.get(&pk_bytes).copied(); + if conflict == BatchConflict::Skip && prior.is_some() { + return Ok(MutationResult { + wal_records: Vec::new(), + }); + } + let replaced = if self.schema.is_bitemporal() { + None + } else { + prior + }; + + let mut wal_records = Vec::with_capacity(2); + let key = pk_bytes.clone(); + self.commit_row( + row.values, + pk_bytes, + row.surrogate, + replaced.as_slice(), + &mut wal_records, + )?; + priors.entry(key).or_insert(PriorBinding { + location: prior, + tombstoned: replaced.is_some(), + }); + Ok(MutationResult { wal_records }) + } + + /// Undo the rows a failed batch wrote from memtable row `start` on. + fn undo_batch(&mut self, start: usize, priors: HashMap, PriorBinding>) { + for (pk_bytes, prior) in priors { + match prior.location { + Some(location) => { + if prior.tombstoned + && let Some(bm) = self.delete_bitmaps.get_mut(&location.segment_id) + { + bm.unmark_deleted(location.row_index); + } + self.pk_index.upsert(pk_bytes, location); + } + None => { + self.pk_index.remove(&pk_bytes); + } + } + } + self.cut_memtable_to(start); + } +} + +#[cfg(test)] +mod tests { + use nodedb_types::columnar::{ColumnDef, ColumnType, ColumnarSchema}; + + use super::*; + use crate::pk_index::encode_pk; + + fn schema() -> ColumnarSchema { + ColumnarSchema::new(vec![ + ColumnDef::required("id", ColumnType::Int64).with_primary_key(), + ColumnDef::required("name", ColumnType::String), + ]) + .expect("valid") + } + + fn row(id: i64, name: &str) -> Vec { + vec![Value::Integer(id), Value::String(name.into())] + } + + fn batch(rows: &[Vec]) -> Vec> { + rows.iter() + .map(|values| BatchRow { + values, + surrogate: None, + }) + .collect() + } + + #[test] + fn failed_batch_applies_nothing() { + let mut engine = MutationEngine::new("t".into(), schema()); + engine.insert(&row(1, "a")).expect("insert"); + + let rows = vec![ + row(1, "a2"), + row(2, "b"), + row(2, "b2"), + vec![Value::Integer(3), Value::Null], + ]; + let err = engine + .insert_batch(batch(&rows), BatchConflict::Upsert) + .unwrap_err(); + assert!(matches!(err, ColumnarError::NullViolation(_))); + + let live: Vec> = engine + .scan_memtable_rows() + .collect::>() + .expect("read"); + assert_eq!(live, vec![row(1, "a")]); + assert_eq!(engine.memtable().row_count(), 1); + assert_eq!(engine.memtable_surrogates().len(), 1); + assert_eq!(engine.pk_index().len(), 1); + let loc = engine + .pk_index() + .get(&encode_pk(&Value::Integer(1))) + .copied(); + assert_eq!(loc.map(|l| l.row_index), Some(0)); + + // An index cut by the undo starts live when a new row lands on it. + engine.insert(&row(4, "d")).expect("insert after undo"); + engine.insert(&row(5, "e")).expect("insert after undo"); + let live: Vec> = engine + .scan_memtable_rows() + .collect::>() + .expect("read"); + assert_eq!(live, vec![row(1, "a"), row(4, "d"), row(5, "e")]); + } + + #[test] + fn batch_upserts_and_skips_like_single_rows() { + let mut engine = MutationEngine::new("t".into(), schema()); + engine.insert(&row(1, "a")).expect("insert"); + + let rows = vec![row(1, "a2"), row(2, "b")]; + let results = engine + .insert_batch(batch(&rows), BatchConflict::Upsert) + .expect("batch"); + assert_eq!(results.len(), 2); + let live: Vec> = engine + .scan_memtable_rows() + .collect::>() + .expect("read"); + assert_eq!(live, vec![row(1, "a2"), row(2, "b")]); + + let rows = vec![row(2, "b2"), row(3, "c")]; + let results = engine + .insert_batch(batch(&rows), BatchConflict::Skip) + .expect("batch"); + assert!(results[0].wal_records.is_empty()); + assert_eq!(results[1].wal_records.len(), 1); + let live: Vec> = engine + .scan_memtable_rows() + .collect::>() + .expect("read"); + assert_eq!(live, vec![row(1, "a2"), row(2, "b"), row(3, "c")]); + } +} diff --git a/nodedb-columnar/src/mutation/engine.rs b/nodedb-columnar/src/mutation/engine.rs index 29534c539..3882fb40e 100644 --- a/nodedb-columnar/src/mutation/engine.rs +++ b/nodedb-columnar/src/mutation/engine.rs @@ -168,7 +168,10 @@ impl MutationEngine { /// /// Skips rows marked as deleted in the memtable's virtual segment /// delete bitmap. For rows in flushed segments, use `SegmentReader`. - pub fn scan_memtable_rows(&self) -> impl Iterator> + '_ { + /// Yields `Err` for a live row with a corrupt cell. + pub fn scan_memtable_rows( + &self, + ) -> impl Iterator, ColumnarError>> + '_ { let deletes = self.delete_bitmaps.get(&self.memtable_segment_id); self.memtable .iter_rows() @@ -177,7 +180,7 @@ impl MutationEngine { if deletes.is_some_and(|bm| bm.is_deleted(row_idx as u32)) { None } else { - Some(row) + Some(row.map_err(|e| self.memtable_read_fault(e))) } }) } @@ -187,9 +190,10 @@ impl MutationEngine { /// Yields `(Option, Vec)`. The surrogate is `None` /// for rows inserted without one (test fixtures, legacy paths). Deleted /// rows are filtered out exactly as in [`Self::scan_memtable_rows`]. + /// Yields `Err` for a live row with a corrupt cell. pub fn scan_memtable_rows_with_surrogates( &self, - ) -> impl Iterator, Vec)> + '_ { + ) -> impl Iterator, Vec), ColumnarError>> + '_ { let deletes = self.delete_bitmaps.get(&self.memtable_segment_id); let surrogates = &self.memtable_surrogates; self.memtable @@ -200,20 +204,37 @@ impl MutationEngine { return None; } let surrogate = surrogates.get(row_idx).copied().flatten(); - Some((surrogate, row)) + Some( + row.map(|row| (surrogate, row)) + .map_err(|e| self.memtable_read_fault(e)), + ) }) } - /// Get a single row from the memtable by index (None if deleted). - pub fn get_memtable_row(&self, row_idx: usize) -> Option> { + /// Get a single row from the memtable by index. + /// + /// `Ok(None)` if the row is deleted or past the end. `Err` if a cell of + /// the row is corrupt. + pub fn get_memtable_row(&self, row_idx: usize) -> Result>, ColumnarError> { if self .delete_bitmaps .get(&self.memtable_segment_id) .is_some_and(|bm| bm.is_deleted(row_idx as u32)) { - return None; + return Ok(None); } - self.memtable.get_row(row_idx) + self.memtable + .get_row(row_idx) + .map_err(|e| self.memtable_read_fault(e)) + } + + /// Record a corrupt memtable cell and hand the error back. + /// + /// Every memtable row read of this engine routes its error through here, + /// so the corruption is reported once, where it is detected. + pub(super) fn memtable_read_fault(&self, err: ColumnarError) -> ColumnarError { + crate::diag::memtable_cell_corrupt(&err, &self.collection); + err } /// Roll back in-memory inserts to `row_count_before`. @@ -245,9 +266,22 @@ impl MutationEngine { } } // 3. Truncate memtable and surrogate list. - self.memtable.truncate_to(row_count_before); - self.memtable_surrogates.truncate(row_count_before); - self.memtable_row_counter = row_count_before as u32; + self.cut_memtable_to(row_count_before); + } + + /// Cut the memtable, its surrogate table and its row counter back to + /// `row_count` rows. + /// + /// A cut row can carry a tombstone: a later row of the same write + /// replaced it. Those tombstones are cleared too, so a row appended + /// later at that index starts live. + pub(super) fn cut_memtable_to(&mut self, row_count: usize) { + if let Some(bm) = self.delete_bitmaps.get_mut(&self.memtable_segment_id) { + bm.unmark_from(row_count as u32); + } + self.memtable.truncate_to(row_count); + self.memtable_surrogates.truncate(row_count); + self.memtable_row_counter = row_count as u32; } /// Reverse one or more positional deletes: for each `(pk_bytes, location)` @@ -562,6 +596,7 @@ mod tests { // Row-boundary: count rows that pass the bitmap membership check. let passing: Vec<_> = engine .scan_memtable_rows_with_surrogates() + .map(|r| r.expect("read")) .filter(|(sur, _)| sur.is_some_and(|s| bitmap.contains(s))) .collect(); assert_eq!(passing.len(), 2); @@ -635,7 +670,7 @@ mod tests { .expect("update"); let live: Vec> = engine .scan_memtable_rows_with_surrogates() - .map(|(surrogate, _)| surrogate) + .map(|r| r.expect("read").0) .collect(); assert_eq!( live, @@ -665,7 +700,7 @@ mod tests { .expect("update"); let live: Vec> = engine .scan_memtable_rows_with_surrogates() - .map(|(surrogate, _)| surrogate) + .map(|r| r.expect("read").0) .collect(); assert_eq!( live, diff --git a/nodedb-columnar/src/mutation/mod.rs b/nodedb-columnar/src/mutation/mod.rs index 9ab451f01..e9df26865 100644 --- a/nodedb-columnar/src/mutation/mod.rs +++ b/nodedb-columnar/src/mutation/mod.rs @@ -7,12 +7,14 @@ //! columnar write operations. It produces WAL records that must be //! persisted before the mutation is considered durable. +pub mod batch; pub mod engine; pub mod flush; pub mod snapshot; pub mod truncate; pub mod write; +pub use batch::{BatchConflict, BatchRow}; pub use engine::{MutationEngine, MutationResult}; pub use snapshot::{ColumnDataSnapshot, ColumnarEngineSnapshot}; pub use truncate::TruncatedRows; diff --git a/nodedb-columnar/src/mutation/snapshot.rs b/nodedb-columnar/src/mutation/snapshot.rs index 830dcba7c..c95d6f744 100644 --- a/nodedb-columnar/src/mutation/snapshot.rs +++ b/nodedb-columnar/src/mutation/snapshot.rs @@ -502,7 +502,10 @@ mod tests { assert_eq!(restored.pk_index().len(), 3); // Scan rows — all 3 should be present. - let rows: Vec> = restored.scan_memtable_rows().collect(); + let rows: Vec> = restored + .scan_memtable_rows() + .collect::>() + .expect("read"); assert_eq!(rows.len(), 3); assert_eq!(rows[0][0], Value::Integer(1)); assert_eq!(rows[1][1], Value::String("Bob".into())); @@ -533,7 +536,10 @@ mod tests { let (restored, _, _) = MutationEngine::from_snapshot(snap2).expect("from_snapshot"); // scan_memtable_rows skips deleted row 1 (id=20). - let rows: Vec> = restored.scan_memtable_rows().collect(); + let rows: Vec> = restored + .scan_memtable_rows() + .collect::>() + .expect("read"); assert_eq!(rows.len(), 1); assert_eq!(rows[0][0], Value::Integer(10)); @@ -691,4 +697,45 @@ mod tests { assert_eq!(flushed.len(), 1); assert!(flushed_surrogates.is_empty()); } + + /// A memtable JSON cell restored with bytes that are not MessagePack is + /// refused on every read path, never read as NULL or as an absent row. + #[test] + fn a_corrupt_json_cell_is_refused_not_read_as_null() { + let schema = ColumnarSchema { + columns: vec![ + ColumnDef::required("id", ColumnType::Int64).with_primary_key(), + ColumnDef::nullable("doc", ColumnType::Json), + ], + version: 1, + }; + let mut engine = MutationEngine::new("json_col".to_string(), schema); + engine + .insert(&[Value::Integer(1), Value::Integer(5)]) + .expect("insert"); + let mut snap = engine.export_snapshot(&[], &[]).expect("export"); + let Some(ColumnDataSnapshot::Json { data, offsets, .. }) = snap.memtable_columns.get_mut(1) + else { + panic!("doc column snapshots as JSON"); + }; + // 0xC1 is the one MessagePack marker that is never valid. + *data = vec![0xC1]; + *offsets = vec![0, 1]; + let (restored, _, _) = MutationEngine::from_snapshot(snap).expect("from_snapshot"); + + let is_corrupt = |e: &ColumnarError| matches!(e, ColumnarError::MemtableCellCorrupt { column, row: 0, .. } if column == "doc"); + let scanned: Vec<_> = restored.scan_memtable_rows().collect(); + assert_eq!(scanned.len(), 1); + assert!(scanned[0].as_ref().is_err_and(is_corrupt)); + let with_surrogates: Vec<_> = restored.scan_memtable_rows_with_surrogates().collect(); + assert!(with_surrogates[0].as_ref().is_err_and(is_corrupt)); + assert!(restored.get_memtable_row(0).as_ref().is_err_and(is_corrupt)); + let pk = crate::pk_index::encode_pk(&Value::Integer(1)); + assert!( + restored + .lookup_memtable_row_by_pk(&pk) + .as_ref() + .is_err_and(is_corrupt) + ); + } } diff --git a/nodedb-columnar/src/mutation/truncate.rs b/nodedb-columnar/src/mutation/truncate.rs index 30584b981..6849416d8 100644 --- a/nodedb-columnar/src/mutation/truncate.rs +++ b/nodedb-columnar/src/mutation/truncate.rs @@ -124,7 +124,10 @@ mod tests { #[test] fn restore_truncated_brings_back_rows_and_tombstones_exactly() { let mut engine = seeded(); - let before: Vec> = engine.scan_memtable_rows().collect(); + let before: Vec> = engine + .scan_memtable_rows() + .collect::>() + .expect("read"); let pre = engine.truncate(); engine @@ -132,7 +135,10 @@ mod tests { .expect("insert after truncate"); engine.restore_truncated(pre); - let after: Vec> = engine.scan_memtable_rows().collect(); + let after: Vec> = engine + .scan_memtable_rows() + .collect::>() + .expect("read"); assert_eq!( after, before, "restore must reproduce the pre-truncate rows" diff --git a/nodedb-columnar/src/mutation/write.rs b/nodedb-columnar/src/mutation/write.rs index b77341bf4..915a257ef 100644 --- a/nodedb-columnar/src/mutation/write.rs +++ b/nodedb-columnar/src/mutation/write.rs @@ -17,7 +17,8 @@ impl MutationEngine { /// /// Validates schema. If the PK already exists, the prior row is /// tombstoned via the segment's delete bitmap (a single positional - /// delete) before the new row is appended to the memtable. The PK + /// delete) after the new row is appended to the memtable. A row that + /// fails validation returns an error and leaves the prior row live. The PK /// index is rebound to the new row location. This matches the /// ClickHouse / Iceberg "sparse PK + positional delete" model and /// keeps `SELECT WHERE pk = X` linearizable on one row without a @@ -28,47 +29,7 @@ impl MutationEngine { /// want `ON CONFLICT DO NOTHING` semantics should use /// [`Self::insert_if_absent`]. pub fn insert(&mut self, values: &[Value]) -> Result { - let pk_bytes = self.extract_pk_bytes(values)?; - let mut wal_records = Vec::with_capacity(2); - - // Bitemporal collections preserve every version of a PK: each - // write appends a new row with a distinct `_ts_system` stamp and - // the prior row stays visible to `AS OF` queries. Skipping the - // upsert-tombstone here keeps compaction lossless without - // needing a separate "version-aware" delete bitmap. The PK - // index is still rebound below so current-state reads see the - // latest version. - let bitemporal = self.schema.is_bitemporal(); - - // If a prior row exists for this PK, tombstone it in place so - // subsequent scans skip the stale row. The PK index is rebound - // below to the freshly-appended row. - if !bitemporal && let Some(prior) = self.pk_index.get(&pk_bytes).copied() { - let bitmap = self.delete_bitmaps.entry(prior.segment_id).or_default(); - bitmap.mark_deleted(prior.row_index); - wal_records.push(ColumnarWalRecord::DeleteRows { - collection: self.collection.clone(), - segment_id: prior.segment_id, - row_indices: vec![prior.row_index], - }); - } - - let row_data = encode_row_for_wal(values)?; - wal_records.push(ColumnarWalRecord::InsertRow { - collection: self.collection.clone(), - row_data, - }); - - self.memtable.append_row(values)?; - let location = RowLocation { - segment_id: self.memtable_segment_id, - row_index: self.memtable_row_counter, - }; - self.pk_index.upsert(pk_bytes, location); - self.memtable_surrogates.push(None); - self.memtable_row_counter += 1; - - Ok(MutationResult { wal_records }) + self.upsert_row(values, None) } /// Insert with a stable cross-engine surrogate identity. @@ -80,37 +41,81 @@ impl MutationEngine { &mut self, values: &[Value], surrogate: Surrogate, + ) -> Result { + self.upsert_row(values, Some(surrogate)) + } + + /// Shared body of [`Self::insert`] and [`Self::insert_with_surrogate`]. + /// + /// A non-bitemporal write tombstones the prior row for the PK. A + /// bitemporal write keeps every version of a PK: the prior row stays + /// visible to `AS OF` queries, and the PK index still rebinds so + /// current-state reads see the latest version. + fn upsert_row( + &mut self, + values: &[Value], + surrogate: Option, ) -> Result { let pk_bytes = self.extract_pk_bytes(values)?; + let prior = if self.schema.is_bitemporal() { + None + } else { + self.pk_index.get(&pk_bytes).copied() + }; let mut wal_records = Vec::with_capacity(2); + self.commit_row( + values, + pk_bytes, + surrogate, + prior.as_slice(), + &mut wal_records, + )?; + Ok(MutationResult { wal_records }) + } + + /// Append `values` as a memtable row bound to `pk_bytes`, then + /// tombstone each row in `replaced`. + /// + /// Every fallible step runs before any state changes: + /// `encode_row_for_wal` runs first, and `append_row` is all-or-nothing. + /// An error leaves the engine and `wal_records` unchanged. Tombstone + /// records go into `wal_records` before the insert record, so replay + /// applies them in the same order. + pub(super) fn commit_row( + &mut self, + values: &[Value], + pk_bytes: Vec, + surrogate: Option, + replaced: &[RowLocation], + wal_records: &mut Vec, + ) -> Result<(), ColumnarError> { + let row_data = encode_row_for_wal(values)?; + self.memtable.append_row(values)?; - let bitemporal = self.schema.is_bitemporal(); - if !bitemporal && let Some(prior) = self.pk_index.get(&pk_bytes).copied() { - let bitmap = self.delete_bitmaps.entry(prior.segment_id).or_default(); - bitmap.mark_deleted(prior.row_index); + for prior in replaced { + self.delete_bitmaps + .entry(prior.segment_id) + .or_default() + .mark_deleted(prior.row_index); wal_records.push(ColumnarWalRecord::DeleteRows { collection: self.collection.clone(), segment_id: prior.segment_id, row_indices: vec![prior.row_index], }); } - - let row_data = encode_row_for_wal(values)?; wal_records.push(ColumnarWalRecord::InsertRow { collection: self.collection.clone(), row_data, }); - self.memtable.append_row(values)?; let location = RowLocation { segment_id: self.memtable_segment_id, row_index: self.memtable_row_counter, }; self.pk_index.upsert(pk_bytes, location); - self.memtable_surrogates.push(Some(surrogate)); + self.memtable_surrogates.push(surrogate); self.memtable_row_counter += 1; - - Ok(MutationResult { wal_records }) + Ok(()) } /// `INSERT ... ON CONFLICT DO NOTHING` semantics: append only if the @@ -126,24 +131,9 @@ impl MutationEngine { wal_records: Vec::new(), }); } - - let row_data = encode_row_for_wal(values)?; - let wal = ColumnarWalRecord::InsertRow { - collection: self.collection.clone(), - row_data, - }; - self.memtable.append_row(values)?; - let location = RowLocation { - segment_id: self.memtable_segment_id, - row_index: self.memtable_row_counter, - }; - self.pk_index.upsert(pk_bytes, location); - self.memtable_surrogates.push(None); - self.memtable_row_counter += 1; - - Ok(MutationResult { - wal_records: vec![wal], - }) + let mut wal_records = Vec::with_capacity(1); + self.commit_row(values, pk_bytes, None, &[], &mut wal_records)?; + Ok(MutationResult { wal_records }) } /// Look up the current row for a PK in the memtable, if present. @@ -154,12 +144,22 @@ impl MutationEngine { /// used by `ON CONFLICT DO UPDATE` to read the would-be-merged row /// when the duplicate hits the memtable — the common case under /// back-to-back inserts. - pub fn lookup_memtable_row_by_pk(&self, pk_bytes: &[u8]) -> Option> { - let loc = self.pk_index.get(pk_bytes).copied()?; + /// + /// `Err` when a cell of the bound row is corrupt: a corrupt row is not an + /// absent row. + pub fn lookup_memtable_row_by_pk( + &self, + pk_bytes: &[u8], + ) -> Result>, ColumnarError> { + let Some(loc) = self.pk_index.get(pk_bytes).copied() else { + return Ok(None); + }; if loc.segment_id != self.memtable_segment_id { - return None; + return Ok(None); } - self.memtable.get_row(loc.row_index as usize) + self.memtable + .get_row(loc.row_index as usize) + .map_err(|e| self.memtable_read_fault(e)) } /// Delete a row by PK value. Returns WAL record to persist. @@ -199,7 +199,9 @@ impl MutationEngine { /// `updates` maps column names to new values. Columns not in the map /// retain their existing values from the old row. /// - /// Returns WAL records for both the delete and the insert. + /// Returns WAL records for both the delete and the insert. The old row + /// is tombstoned only after the new row is appended, so an invalid new + /// row returns an error and leaves the old row live. /// /// NOTE: The caller must provide the full old row values for the re-insert. /// This method takes the complete new row (already merged with old values). @@ -215,29 +217,175 @@ impl MutationEngine { new_values: &[Value], flushed_surrogate: Option, ) -> Result { - let surrogate = match self.pk_index.get(&encode_pk(old_pk)) { - Some(loc) if loc.segment_id == self.memtable_segment_id => self - .memtable_surrogates - .get(loc.row_index as usize) + let old_pk_bytes = encode_pk(old_pk); + let old_location = self + .pk_index + .get(&old_pk_bytes) + .copied() + .ok_or(ColumnarError::PrimaryKeyNotFound)?; + let surrogate = if old_location.segment_id == self.memtable_segment_id { + self.memtable_surrogates + .get(old_location.row_index as usize) .copied() - .flatten(), - Some(_) => flushed_surrogate, - None => None, + .flatten() + } else { + flushed_surrogate }; - // Delete the old row. - let delete_result = self.delete(old_pk)?; + let new_pk_bytes = self.extract_pk_bytes(new_values)?; + let pk_changed = new_pk_bytes != old_pk_bytes; - // Insert the new row under the old row's surrogate. - let insert_result = match surrogate { - Some(surrogate) => self.insert_with_surrogate(new_values, surrogate)?, - None => self.insert(new_values)?, - }; + // The old row is always tombstoned. When the PK changes, a + // non-bitemporal row already bound to the new PK is replaced too, + // as an upsert replaces it. + let mut replaced = Vec::with_capacity(2); + replaced.push(old_location); + if pk_changed + && !self.schema.is_bitemporal() + && let Some(prior) = self.pk_index.get(&new_pk_bytes).copied() + { + replaced.push(prior); + } - // Combine WAL records. - let mut wal_records = delete_result.wal_records; - wal_records.extend(insert_result.wal_records); + let mut wal_records = Vec::with_capacity(replaced.len() + 1); + self.commit_row( + new_values, + new_pk_bytes, + surrogate, + &replaced, + &mut wal_records, + )?; + if pk_changed { + self.pk_index.remove(&old_pk_bytes); + } Ok(MutationResult { wal_records }) } } + +#[cfg(test)] +mod tests { + use nodedb_types::columnar::{ColumnDef, ColumnType, ColumnarSchema}; + + use super::*; + + fn schema() -> ColumnarSchema { + ColumnarSchema::new(vec![ + ColumnDef::required("id", ColumnType::Int64).with_primary_key(), + ColumnDef::required("name", ColumnType::String), + ColumnDef::nullable("score", ColumnType::Float64), + ]) + .expect("valid") + } + + fn row(id: i64, name: &str, score: f64) -> Vec { + vec![ + Value::Integer(id), + Value::String(name.into()), + Value::Float(score), + ] + } + + fn live_rows(engine: &MutationEngine) -> Vec> { + engine + .scan_memtable_rows() + .collect::>() + .expect("read") + } + + #[test] + fn rejected_upsert_keeps_prior_row() { + let mut engine = MutationEngine::new("t".into(), schema()); + engine.insert(&row(1, "a", 0.5)).expect("insert"); + + let err = engine + .insert(&[Value::Integer(1), Value::Null, Value::Null]) + .unwrap_err(); + assert!(matches!(err, ColumnarError::NullViolation(_))); + + assert_eq!(live_rows(&engine), vec![row(1, "a", 0.5)]); + assert_eq!(engine.memtable().row_count(), 1); + let loc = engine + .pk_index() + .get(&encode_pk(&Value::Integer(1))) + .copied(); + assert_eq!(loc.map(|l| l.row_index), Some(0)); + assert!( + engine + .delete_bitmap(engine.memtable_segment_id()) + .is_none_or(|bm| !bm.is_deleted(0)) + ); + } + + #[test] + fn rejected_upsert_with_surrogate_keeps_prior_row() { + let mut engine = MutationEngine::new("t".into(), schema()); + engine + .insert_with_surrogate(&row(1, "a", 0.5), Surrogate(7)) + .expect("insert"); + + let err = engine + .insert_with_surrogate( + &[ + Value::Integer(1), + Value::String("b".into()), + Value::Bool(true), + ], + Surrogate(8), + ) + .unwrap_err(); + assert!(matches!(err, ColumnarError::TypeMismatch { .. })); + + assert_eq!(live_rows(&engine), vec![row(1, "a", 0.5)]); + assert_eq!(engine.memtable_surrogates(), &[Some(Surrogate(7))]); + } + + #[test] + fn rejected_insert_if_absent_leaves_engine_unchanged() { + let mut engine = MutationEngine::new("t".into(), schema()); + let err = engine + .insert_if_absent(&[Value::Integer(1), Value::Null, Value::Null]) + .unwrap_err(); + assert!(matches!(err, ColumnarError::NullViolation(_))); + assert!(engine.pk_index().is_empty()); + assert_eq!(engine.memtable().row_count(), 0); + + engine + .insert_if_absent(&row(1, "a", 0.5)) + .expect("insert after rejected row"); + assert_eq!(live_rows(&engine), vec![row(1, "a", 0.5)]); + } + + #[test] + fn rejected_update_keeps_old_row() { + let mut engine = MutationEngine::new("t".into(), schema()); + engine.insert(&row(1, "a", 0.5)).expect("insert"); + + let err = engine + .update( + &Value::Integer(1), + &[Value::Integer(1), Value::Null, Value::Null], + None, + ) + .unwrap_err(); + assert!(matches!(err, ColumnarError::NullViolation(_))); + + assert_eq!(live_rows(&engine), vec![row(1, "a", 0.5)]); + assert!(engine.pk_index().contains(&encode_pk(&Value::Integer(1)))); + } + + #[test] + fn update_to_new_pk_rebinds_index() { + let mut engine = MutationEngine::new("t".into(), schema()); + engine.insert(&row(1, "a", 0.5)).expect("insert"); + + let result = engine + .update(&Value::Integer(1), &row(2, "b", 0.25), None) + .expect("update"); + assert_eq!(result.wal_records.len(), 2); + + assert_eq!(live_rows(&engine), vec![row(2, "b", 0.25)]); + assert!(!engine.pk_index().contains(&encode_pk(&Value::Integer(1)))); + assert!(engine.pk_index().contains(&encode_pk(&Value::Integer(2)))); + } +} diff --git a/nodedb-columnar/src/reader/block_decode.rs b/nodedb-columnar/src/reader/block_decode.rs index 0d326e6e8..8ef344d09 100644 --- a/nodedb-columnar/src/reader/block_decode.rs +++ b/nodedb-columnar/src/reader/block_decode.rs @@ -1,79 +1,36 @@ // SPDX-License-Identifier: Apache-2.0 -//! Block-level decode helpers: type inference, null-fill, and compressed-block decoding. +//! Block-level decode helpers: null-fill and compressed-block decoding by the +//! column's recorded [`BlockLayout`]. use nodedb_codec::{ColumnCodec, ResolvedColumnCodec}; use crate::error::ColumnarError; -use crate::format::ColumnMeta; +use crate::format::BlockLayout; use super::types::DecodedColumn; -/// Simplified column kind for decode dispatch. -#[derive(Debug, Clone, Copy)] -pub(super) enum ColumnKind { - Int64, - Float64, - VarLen, - Binary, - DictEncoded, -} - -/// Infer a simplified column type from ColumnMeta for decode dispatch. -/// -/// We use the codec as a strong signal: DeltaFastLanesLz4 = numeric, -/// FsstLz4 = string, etc. The name is a fallback heuristic. -pub(super) fn infer_column_type(meta: &ColumnMeta) -> ColumnKind { - // Dict-encoded columns store IDs as DeltaFastLanesLz4 but must be decoded - // differently — the presence of a dictionary distinguishes them. - if meta.dictionary.is_some() { - return ColumnKind::DictEncoded; - } - - match meta.codec { - ResolvedColumnCodec::DeltaFastLanesLz4 - | ResolvedColumnCodec::DeltaFastLanesRans - | ResolvedColumnCodec::FastLanesLz4 - | ResolvedColumnCodec::Delta - | ResolvedColumnCodec::DoubleDelta => ColumnKind::Int64, - - ResolvedColumnCodec::AlpFastLanesLz4 - | ResolvedColumnCodec::AlpFastLanesRans - | ResolvedColumnCodec::AlpRdLz4 - | ResolvedColumnCodec::PcodecLz4 - | ResolvedColumnCodec::Gorilla => ColumnKind::Float64, - - ResolvedColumnCodec::FsstLz4 | ResolvedColumnCodec::FsstRans => ColumnKind::VarLen, - - // LZ4/Raw/Zstd could be bool, binary, decimal, uuid, vector — use - // block_stats to distinguish: if min/max are NaN → binary-like. - ResolvedColumnCodec::Lz4 | ResolvedColumnCodec::Raw | ResolvedColumnCodec::Zstd => { - if meta.block_stats.first().is_some_and(|s| !s.min.is_nan()) { - ColumnKind::Int64 // Numeric fallback. - } else { - ColumnKind::Binary - } - } - } -} - -/// Create an empty DecodedColumn for the given kind. -pub(super) fn empty_decoded(kind: &ColumnKind) -> DecodedColumn { - match kind { - ColumnKind::Int64 => DecodedColumn::Int64 { +/// Create an empty DecodedColumn for the given layout. +pub(super) fn empty_decoded(layout: BlockLayout) -> DecodedColumn { + match layout { + BlockLayout::Int64 => DecodedColumn::Int64 { values: Vec::new(), valid: Vec::new(), }, - ColumnKind::Float64 => DecodedColumn::Float64 { + BlockLayout::Float64 => DecodedColumn::Float64 { values: Vec::new(), valid: Vec::new(), }, - ColumnKind::VarLen | ColumnKind::Binary => DecodedColumn::Binary { + BlockLayout::PackedBool => DecodedColumn::Bool { + values: Vec::new(), + valid: Vec::new(), + }, + BlockLayout::VarLen | BlockLayout::FixedWidth => DecodedColumn::Binary { data: Vec::new(), offsets: Vec::new(), valid: Vec::new(), }, - ColumnKind::DictEncoded => DecodedColumn::DictEncoded { + BlockLayout::DictIds => DecodedColumn::DictEncoded { ids: Vec::new(), dictionary: Vec::new(), // Populated during decode_block. valid: Vec::new(), @@ -149,7 +106,7 @@ pub(super) fn result_valid_slice_mut(result: &mut DecodedColumn, offset: usize) pub(super) fn decode_block( result: &mut DecodedColumn, block_data: &[u8], - kind: &ColumnKind, + layout: BlockLayout, codec: ResolvedColumnCodec, row_count: usize, dictionary: Option<&[String]>, @@ -171,8 +128,8 @@ pub(super) fn decode_block( .map(|i| bitmap[i / 8] & (1 << (i % 8)) != 0) .collect(); - match kind { - ColumnKind::Int64 => { + match layout { + BlockLayout::Int64 => { let DecodedColumn::Int64 { values, valid: v } = result else { append_null_fill(result, row_count); return Ok(()); @@ -184,7 +141,7 @@ pub(super) fn decode_block( } v.extend_from_slice(&valid); } - ColumnKind::Float64 => { + BlockLayout::Float64 => { let DecodedColumn::Float64 { values, valid: v } = result else { append_null_fill(result, row_count); return Ok(()); @@ -196,7 +153,22 @@ pub(super) fn decode_block( } v.extend_from_slice(&valid); } - ColumnKind::VarLen => { + BlockLayout::PackedBool => { + let DecodedColumn::Bool { values, valid: v } = result else { + append_null_fill(result, row_count); + return Ok(()); + }; + let packed = nodedb_codec::decode_bytes_pipeline(payload, codec.into_column_codec())?; + if packed.len() != bitmap_size { + return Err(block_corruption(format!( + "packed bool block holds {} bytes for {row_count} rows", + packed.len() + ))); + } + values.extend((0..row_count).map(|i| packed[i / 8] & (1 << (i % 8)) != 0)); + v.extend_from_slice(&valid); + } + BlockLayout::VarLen => { let DecodedColumn::Binary { data, offsets, @@ -262,18 +234,27 @@ pub(super) fn decode_block( let decoded_bytes = nodedb_codec::decode_bytes_pipeline(string_data, codec.into_column_codec())?; - // decoded_offsets has row_count + 1 entries (including sentinel). - // Map them to absolute positions in the output data buffer. - let base = data.len() as u32; - let n_offsets = (row_count + 1).min(decoded_offsets.len()); - for &off in &decoded_offsets[..n_offsets] { - offsets.push(base + off as u32); + // The block's offsets are relative to its first byte: row_count + 1 + // entries, the first 0, the last the block's byte length. + if decoded_offsets.len() != row_count + 1 + || decoded_offsets.first() != Some(&0) + || decoded_offsets.windows(2).any(|w| w[0] > w[1]) + || decoded_offsets + .last() + .is_some_and(|&end| usize::try_from(end) != Ok(decoded_bytes.len())) + { + return Err(block_corruption(format!( + "variable-length offset table of {} entries does not span {} bytes \ + for {row_count} rows", + decoded_offsets.len(), + decoded_bytes.len() + ))); } - + append_block_offsets(data, offsets, &decoded_offsets[1..])?; data.extend_from_slice(&decoded_bytes); v.extend_from_slice(&valid); } - ColumnKind::Binary => { + BlockLayout::FixedWidth => { let DecodedColumn::Binary { data, offsets, @@ -285,23 +266,23 @@ pub(super) fn decode_block( }; let decoded_bytes = nodedb_codec::decode_bytes_pipeline(payload, codec.into_column_codec())?; - let base = data.len() as u32; - - if row_count > 0 && !decoded_bytes.is_empty() { - let chunk_size = decoded_bytes.len() / row_count; - for i in 0..row_count { - offsets.push(base + (i * chunk_size) as u32); - } - offsets.push(base + decoded_bytes.len() as u32); - } else { - let last = *offsets.last().unwrap_or(&0); - offsets.extend(std::iter::repeat_n(last, row_count + 1)); + // Every row holds a full cell, so the block divides evenly. + let width = decoded_bytes.len().checked_div(row_count).unwrap_or(0); + if width * row_count != decoded_bytes.len() { + return Err(block_corruption(format!( + "fixed-width block of {} bytes does not divide into {row_count} rows", + decoded_bytes.len() + ))); } - + let ends: Vec = (1..=row_count) + .map(|i| i64::try_from(i * width)) + .collect::>() + .map_err(|_| block_corruption("fixed-width cell end exceeds i64".into()))?; + append_block_offsets(data, offsets, &ends)?; data.extend_from_slice(&decoded_bytes); v.extend_from_slice(&valid); } - ColumnKind::DictEncoded => { + BlockLayout::DictIds => { let DecodedColumn::DictEncoded { ids, dictionary: col_dict, @@ -335,6 +316,43 @@ pub(super) fn decode_block( Ok(()) } +/// Append one block's row end offsets, rebased onto the end of `data`. +/// +/// `ends` holds each row's end relative to the block's first byte. The start +/// sentinel is pushed once, before the first row of the column, so `offsets` +/// always holds one more entry than the rows decoded so far. +fn append_block_offsets( + data: &[u8], + offsets: &mut Vec, + ends: &[i64], +) -> Result<(), ColumnarError> { + let base = data.len(); + if offsets.is_empty() { + offsets.push(absolute_offset(base, 0)?); + } + for &end in ends { + let relative = usize::try_from(end) + .map_err(|_| block_corruption(format!("negative row end offset {end}")))?; + offsets.push(absolute_offset(base, relative)?); + } + Ok(()) +} + +/// `base + relative` as a `u32` column offset. +fn absolute_offset(base: usize, relative: usize) -> Result { + base.checked_add(relative) + .and_then(|offset| u32::try_from(offset).ok()) + .ok_or_else(|| block_corruption(format!("column offset {base} + {relative} exceeds u32"))) +} + +fn block_corruption(reason: String) -> ColumnarError { + ColumnarError::Corruption { + segment_id: None, + reason, + offset: None, + } +} + #[cfg(test)] mod tests { use nodedb_types::columnar::{ColumnDef, ColumnType, ColumnarSchema}; @@ -361,7 +379,7 @@ mod tests { super::decode_block( &mut result, &block, - &super::ColumnKind::VarLen, + crate::format::BlockLayout::VarLen, nodedb_codec::ResolvedColumnCodec::FsstLz4, 1, None, diff --git a/nodedb-columnar/src/reader/cell.rs b/nodedb-columnar/src/reader/cell.rs new file mode 100644 index 000000000..fab8d5f6e --- /dev/null +++ b/nodedb-columnar/src/reader/cell.rs @@ -0,0 +1,355 @@ +// SPDX-License-Identifier: Apache-2.0 + +//! One decoded segment cell as a `Value`. +//! +//! The rule is the memtable's read rule (`ColumnData::get_value`): a cell +//! reads as the same `Value` before and after its memtable flushes. The +//! declared column type decides the variant, because a segment block records +//! only its physical layout. + +use nodedb_types::columnar::ColumnType; +use nodedb_types::value::Value; +use nodedb_types::value_from_msgpack; + +use crate::error::ColumnarError; + +use super::types::DecodedColumn; + +/// Width of a decimal, UUID or ULID cell. +const ID_WIDTH: usize = 16; + +/// Width of one vector element: a little-endian `f32`. +const F32_WIDTH: usize = 4; + +/// Row `row` of `col` as the `Value` the memtable reads for a cell of +/// `declared` type. +/// +/// A null cell or a row past the column is `Value::Null`. A cell whose bytes +/// do not hold a value of `declared` type is a corrupt segment and an error. +pub fn decoded_cell_value( + col: &DecodedColumn, + row: usize, + declared: &ColumnType, +) -> Result { + let is_valid = |valid: &[bool]| valid.get(row).copied().unwrap_or(false); + match col { + // A time column decodes as `Int64`. The declared type types the cell. + DecodedColumn::Int64 { values, valid } | DecodedColumn::Timestamp { values, valid } => { + Ok(match values.get(row) { + Some(&v) if is_valid(valid) => declared.time_cell(v), + _ => Value::Null, + }) + } + DecodedColumn::Float64 { values, valid } => Ok(match values.get(row) { + Some(&v) if is_valid(valid) => Value::Float(v), + _ => Value::Null, + }), + DecodedColumn::Bool { values, valid } => Ok(match values.get(row) { + Some(&v) if is_valid(valid) => Value::Bool(v), + _ => Value::Null, + }), + DecodedColumn::DictEncoded { + ids, + dictionary, + valid, + } => { + let Some(&id) = ids.get(row).filter(|_| is_valid(valid)) else { + return Ok(Value::Null); + }; + let text = usize::try_from(id) + .ok() + .and_then(|id| dictionary.get(id)) + .ok_or_else(|| { + corruption(format!( + "dictionary ID {id} is outside a dictionary of {} entries", + dictionary.len() + )) + })?; + Ok(Value::String(text.clone())) + } + DecodedColumn::Binary { + data, + offsets, + valid, + } => { + if !is_valid(valid) { + return Ok(Value::Null); + } + let bytes = cell_bytes(data, offsets, row)?; + binary_cell_value(bytes, declared) + } + } +} + +/// The bytes of row `row` of a binary column. +fn cell_bytes<'a>(data: &'a [u8], offsets: &[u32], row: usize) -> Result<&'a [u8], ColumnarError> { + let bound = |i: usize| offsets.get(i).and_then(|&o| usize::try_from(o).ok()); + bound(row) + .zip(row.checked_add(1).and_then(bound)) + .and_then(|(start, end)| data.get(start..end)) + .ok_or_else(|| { + corruption(format!( + "row {row} has no byte range in a binary column of {} offsets and {} bytes", + offsets.len(), + data.len() + )) + }) +} + +/// A binary cell as the `Value` the memtable reads for `declared`. +fn binary_cell_value(bytes: &[u8], declared: &ColumnType) -> Result { + Ok(match declared { + ColumnType::String + | ColumnType::Regex + | ColumnType::SparseVector + | ColumnType::Geometry => Value::String(utf8(bytes, declared)?.to_owned()), + ColumnType::Bytes + | ColumnType::Array + | ColumnType::Set + | ColumnType::Range + | ColumnType::Record => Value::Bytes(bytes.to_vec()), + ColumnType::Json if bytes.is_empty() => Value::Null, + ColumnType::Json => value_from_msgpack(bytes) + .map_err(|e| corruption(format!("JSON cell is not MessagePack: {e}")))?, + ColumnType::Decimal { .. } => Value::Decimal(rust_decimal::Decimal::deserialize(id_cell( + bytes, declared, + )?)), + ColumnType::Uuid => { + Value::Uuid(uuid::Uuid::from_bytes(id_cell(bytes, declared)?).to_string()) + } + ColumnType::Ulid => { + Value::Ulid(ulid::Ulid::from_bytes(id_cell(bytes, declared)?).to_string()) + } + ColumnType::Vector(dim) => vector_cell(bytes, *dim)?, + ColumnType::Int64 + | ColumnType::Float64 + | ColumnType::Bool + | ColumnType::Timestamp + | ColumnType::Timestamptz + | ColumnType::SystemTimestamp + | ColumnType::Duration => { + return Err(corruption(format!( + "a {declared} column has no binary cells" + ))); + } + // `ColumnType` is `#[non_exhaustive]`. The memtable stores a type it + // does not know as raw bytes. + _ => Value::Bytes(bytes.to_vec()), + }) +} + +fn utf8<'a>(bytes: &'a [u8], declared: &ColumnType) -> Result<&'a str, ColumnarError> { + std::str::from_utf8(bytes).map_err(|e| corruption(format!("{declared} cell is not UTF-8: {e}"))) +} + +fn id_cell(bytes: &[u8], declared: &ColumnType) -> Result<[u8; ID_WIDTH], ColumnarError> { + <[u8; ID_WIDTH]>::try_from(bytes).map_err(|_| { + corruption(format!( + "{declared} cell holds {} bytes, not {ID_WIDTH}", + bytes.len() + )) + }) +} + +/// A packed little-endian `f32` vector cell as an array of floats. +fn vector_cell(bytes: &[u8], dim: u32) -> Result { + let expected = usize::try_from(dim) + .ok() + .and_then(|d| d.checked_mul(F32_WIDTH)); + if expected != Some(bytes.len()) { + return Err(corruption(format!( + "VECTOR({dim}) cell holds {} bytes", + bytes.len() + ))); + } + let (chunks, _) = bytes.as_chunks::(); + Ok(Value::Array( + chunks + .iter() + .map(|c| Value::Float(f64::from(f32::from_le_bytes(*c)))) + .collect(), + )) +} + +fn corruption(reason: String) -> ColumnarError { + ColumnarError::Corruption { + segment_id: None, + reason, + offset: None, + } +} + +#[cfg(test)] +mod tests { + use std::collections::HashMap; + + use nodedb_types::columnar::{ColumnDef, ColumnType, ColumnarSchema}; + use nodedb_types::value::Value; + use nodedb_types::{NdbDateTime, NdbDuration}; + + use super::decoded_cell_value; + use crate::memtable::ColumnarMemtable; + use crate::reader::SegmentReader; + use crate::test_support::test_memory; + use crate::writer::{PROFILE_PLAIN, SegmentWriter}; + + /// Rows past one block, so the second block's offsets are covered. + const ROWS: usize = 1500; + + const BASE_MICROS: i64 = 1_700_000_000_000_000; + + fn every_type_schema() -> ColumnarSchema { + ColumnarSchema::new(vec![ + ColumnDef::required("id", ColumnType::Int64).with_primary_key(), + ColumnDef::nullable("f", ColumnType::Float64), + ColumnDef::nullable("b", ColumnType::Bool), + ColumnDef::nullable("s", ColumnType::String), + ColumnDef::nullable("tag", ColumnType::String), + ColumnDef::nullable("raw", ColumnType::Bytes), + ColumnDef::nullable("ts", ColumnType::Timestamp), + ColumnDef::nullable("tz", ColumnType::Timestamptz), + ColumnDef::nullable("sys", ColumnType::SystemTimestamp), + ColumnDef::nullable("dec", ColumnType::Decimal(None)), + ColumnDef::nullable("u", ColumnType::Uuid), + ColumnDef::nullable("ul", ColumnType::Ulid), + ColumnDef::nullable("geo", ColumnType::Geometry), + ColumnDef::nullable("vec", ColumnType::Vector(3)), + ColumnDef::nullable("j", ColumnType::Json), + ColumnDef::nullable("dur", ColumnType::Duration), + ColumnDef::nullable("re", ColumnType::Regex), + ColumnDef::nullable("arr", ColumnType::Array), + ColumnDef::nullable("sparse", ColumnType::SparseVector), + ]) + .expect("valid schema") + } + + fn ulid_text(i: usize) -> String { + let mut bytes = [0u8; 16]; + bytes[..8].copy_from_slice(&(i as u64 + 1).to_be_bytes()); + bytes[15] = 7; + ulid::Ulid::from_bytes(bytes).to_string() + } + + fn json_cell(i: usize) -> Value { + if i.is_multiple_of(5) { + // JSON text of a string scalar. + Value::String("\"scalar\"".into()) + } else { + Value::Object(HashMap::from([ + ("n".to_string(), Value::Integer(i as i64)), + ("tag".to_string(), Value::String("x".into())), + ])) + } + } + + fn row(i: usize) -> Vec { + let mut values = vec![Value::Integer(i as i64)]; + if i % 11 == 4 { + values.extend(std::iter::repeat_n(Value::Null, 18)); + return values; + } + values.extend([ + Value::Float(i as f64 * 0.5), + Value::Bool(i.is_multiple_of(2)), + Value::String(format!("row-{i}")), + Value::String(["alpha", "beta", "gamma"][i % 3].into()), + Value::Bytes(vec![i as u8, 0xFF, 0x00]), + Value::NaiveDateTime(NdbDateTime::from_micros(BASE_MICROS + i as i64)), + Value::DateTime(NdbDateTime::from_micros(BASE_MICROS - i as i64)), + Value::Integer(i as i64 * 1_000), + Value::Decimal(rust_decimal::Decimal::new(i as i64 * 25, 2)), + Value::Uuid(uuid::Uuid::from_u128(i as u128 + 1).to_string()), + Value::Ulid(ulid_text(i)), + Value::String(format!("POINT({i} 1)")), + Value::Array(vec![ + Value::Float(i as f64), + Value::Float(0.5), + Value::Float(-1.25), + ]), + json_cell(i), + Value::Duration(NdbDuration::from_micros(i as i64 * 3)), + Value::Regex(format!("^r{i}$")), + Value::Array(vec![Value::Integer(i as i64), Value::String("a".into())]), + Value::String(format!("{{{i}: 0.5}}")), + ]); + values + } + + /// Every column type reads from a flushed segment as the same `Value` the + /// live memtable reads for the same cell, nulls and the second block + /// included. + #[test] + fn a_flushed_cell_reads_as_its_memtable_value() { + let schema = every_type_schema(); + let mut mt = ColumnarMemtable::new(&schema); + for i in 0..ROWS { + mt.append_row(&row(i)).expect("append"); + } + let expected: Vec> = (0..ROWS) + .map(|i| mt.get_row(i).expect("read").expect("memtable row")) + .collect(); + + let (schema, columns, row_count) = mt.drain_optimized(); + let segment = SegmentWriter::new(PROFILE_PLAIN, test_memory()) + .write_segment(&schema, &columns, row_count, None) + .expect("write segment"); + let reader = SegmentReader::open(&segment).expect("open segment"); + let indices: Vec = (0..schema.columns.len()).collect(); + let decoded = reader.read_columns(&indices, &[]).expect("read columns"); + + for (i, memtable_row) in expected.iter().enumerate() { + for ((col, def), want) in decoded.iter().zip(&schema.columns).zip(memtable_row) { + let got = decoded_cell_value(col, i, &def.column_type).expect("decodable cell"); + assert_eq!(&got, want, "row {i} column '{}'", def.name); + } + } + + // The memtable shapes the comparison relies on. + let first = &expected[1]; + assert_eq!(first[10], Value::Uuid(uuid::Uuid::from_u128(2).to_string())); + assert_eq!(first[11], Value::Ulid(ulid_text(1))); + assert_eq!(first[12], Value::String("POINT(1 1)".into())); + assert_eq!(first[15], Value::Integer(3)); + assert_eq!(first[16], Value::String("^r1$".into())); + assert_eq!(expected[4][10], Value::Null); + } + + /// A cell whose bytes do not hold its declared type is an error, not a + /// default value. + #[test] + fn a_corrupt_binary_cell_is_an_error() { + let col = crate::reader::DecodedColumn::Binary { + data: vec![1, 2, 3], + offsets: vec![0, 3], + valid: vec![true], + }; + for ty in [ + ColumnType::Uuid, + ColumnType::Ulid, + ColumnType::Decimal(None), + ColumnType::Vector(2), + ] { + assert!( + decoded_cell_value(&col, 0, &ty).is_err(), + "{ty} must reject a 3-byte cell" + ); + } + // 0xC1 is the one byte MessagePack never uses. + let not_msgpack = crate::reader::DecodedColumn::Binary { + data: vec![0xC1], + offsets: vec![0, 1], + valid: vec![true], + }; + assert!(decoded_cell_value(¬_msgpack, 0, &ColumnType::Json).is_err()); + let not_utf8 = crate::reader::DecodedColumn::Binary { + data: vec![0xFF], + offsets: vec![0, 1], + valid: vec![true], + }; + assert!(decoded_cell_value(¬_utf8, 0, &ColumnType::String).is_err()); + assert_eq!( + decoded_cell_value(¬_utf8, 0, &ColumnType::Bytes).expect("bytes"), + Value::Bytes(vec![0xFF]) + ); + } +} diff --git a/nodedb-columnar/src/reader/mod.rs b/nodedb-columnar/src/reader/mod.rs index b23e1d3ed..340d0e081 100644 --- a/nodedb-columnar/src/reader/mod.rs +++ b/nodedb-columnar/src/reader/mod.rs @@ -1,9 +1,11 @@ // SPDX-License-Identifier: Apache-2.0 mod block_decode; +mod cell; mod segment_reader; mod types; +pub use cell::decoded_cell_value; pub use segment_reader::OwnedSegmentReader; pub use segment_reader::SegmentReader; pub use types::DecodedColumn; diff --git a/nodedb-columnar/src/reader/segment_reader.rs b/nodedb-columnar/src/reader/segment_reader.rs index 9573661e4..81ba43453 100644 --- a/nodedb-columnar/src/reader/segment_reader.rs +++ b/nodedb-columnar/src/reader/segment_reader.rs @@ -18,8 +18,7 @@ use crate::format::{HEADER_SIZE, SegmentFooter, SegmentHeader}; use crate::predicate::ScanPredicate; use super::block_decode::{ - append_null_fill, decode_block, empty_decoded, infer_column_type, result_valid_len, - result_valid_slice_mut, + append_null_fill, decode_block, empty_decoded, result_valid_len, result_valid_slice_mut, }; use super::types::DecodedColumn; @@ -241,8 +240,7 @@ impl<'a> SegmentReader<'a> { }); } let mut cursor = col_start; - let col_type = infer_column_type(col_meta); - let mut result = empty_decoded(&col_type); + let mut result = empty_decoded(col_meta.layout); let mut global_row: u32 = 0; for block_stat in &col_meta.block_stats { @@ -298,7 +296,7 @@ impl<'a> SegmentReader<'a> { decode_block( &mut result, block_data, - &col_type, + col_meta.layout, col_meta.codec, block_row_count as usize, col_meta.dictionary.as_deref(), diff --git a/nodedb-columnar/src/wal_record.rs b/nodedb-columnar/src/wal_record.rs index a02099f98..272803a4e 100644 --- a/nodedb-columnar/src/wal_record.rs +++ b/nodedb-columnar/src/wal_record.rs @@ -134,21 +134,20 @@ pub fn encode_row_for_wal( buf.extend_from_slice(&(bytes.len() as u32).to_le_bytes()); buf.extend_from_slice(bytes); } - Value::Array(arr) => { - // Vectors stored as: tag(9) + count(u32) + f32 values. + Value::Array(arr) if arr.iter().all(|v| matches!(v, Value::Float(_))) => { + // An all-float array: tag(9) + count(u32) + f64 values, so + // every element decodes to the value it was. buf.push(9); buf.extend_from_slice(&(arr.len() as u32).to_le_bytes()); for v in arr { - let f = match v { - Value::Float(f) => *f as f32, - Value::Integer(n) => *n as f32, - _ => 0.0, - }; - buf.extend_from_slice(&f.to_le_bytes()); + if let Value::Float(f) = v { + buf.extend_from_slice(&f.to_le_bytes()); + } } } _ => { - // Geometry and other complex types: serialize as JSON bytes. + // Any other array, geometry and other complex types: JSON + // bytes. buf.push(10); let json = sonic_rs::to_vec(value).map_err(|e| { crate::error::ColumnarError::Serialization(format!( @@ -168,6 +167,28 @@ pub fn encode_row_for_wal( /// Prevents OOM from crafted/corrupt records with bogus length prefixes. const MAX_FIELD_LEN: usize = 256 * 1024 * 1024; +/// The error for a WAL row that stops decoding at byte `offset`. +fn corrupt(offset: usize, reason: impl Into) -> crate::error::ColumnarError { + crate::error::ColumnarError::WalRowCorrupt { + offset, + reason: reason.into(), + } +} + +/// Read exactly `N` bytes from `data` at `cursor` as an array, advancing +/// cursor. Returns `Err` if not enough bytes remain. +fn read_array( + data: &[u8], + cursor: &mut usize, + context: &str, +) -> Result<[u8; N], crate::error::ColumnarError> { + let at = *cursor; + let slice = read_slice(data, cursor, N, context)?; + slice + .try_into() + .map_err(|_| corrupt(at, format!("truncated {context}"))) +} + /// Read exactly `n` bytes from `data` at `cursor`, advancing cursor. /// Returns `Err` if not enough bytes remain. fn read_slice<'a>( @@ -176,14 +197,17 @@ fn read_slice<'a>( n: usize, context: &str, ) -> Result<&'a [u8], crate::error::ColumnarError> { - let end = cursor.checked_add(n).ok_or_else(|| { - crate::error::ColumnarError::Serialization(format!("overflow in {context}")) - })?; + let end = cursor + .checked_add(n) + .ok_or_else(|| corrupt(*cursor, format!("overflow in {context}")))?; if end > data.len() { - return Err(crate::error::ColumnarError::Serialization(format!( - "truncated {context}: need {n} bytes at offset {cursor}, have {}", - data.len().saturating_sub(*cursor) - ))); + return Err(corrupt( + *cursor, + format!( + "truncated {context}: need {n} bytes, have {}", + data.len().saturating_sub(*cursor) + ), + )); } let slice = &data[*cursor..end]; *cursor = end; @@ -197,126 +221,108 @@ fn read_length_prefixed<'a>( cursor: &mut usize, context: &str, ) -> Result<&'a [u8], crate::error::ColumnarError> { - let len_bytes = read_slice(data, cursor, 4, context)?; - let len = u32::from_le_bytes(len_bytes.try_into().map_err(|_| { - crate::error::ColumnarError::Serialization(format!("truncated {context} len")) - })?) as usize; + let at = *cursor; + let len = u32::from_le_bytes(read_array::<4>(data, cursor, context)?) as usize; if len > MAX_FIELD_LEN { - return Err(crate::error::ColumnarError::Serialization(format!( - "{context} length {len} exceeds maximum {MAX_FIELD_LEN}" - ))); + return Err(corrupt( + at, + format!("{context} length {len} exceeds maximum {MAX_FIELD_LEN}"), + )); } read_slice(data, cursor, len, context) } +/// Read a length-prefixed UTF-8 string. Bytes that are not UTF-8 are an +/// error, never replacement characters. +fn read_utf8( + data: &[u8], + cursor: &mut usize, + context: &str, +) -> Result { + let at = *cursor; + let bytes = read_length_prefixed(data, cursor, context)?; + String::from_utf8(bytes.to_vec()) + .map_err(|e| corrupt(at, format!("{context} is not UTF-8: {e}"))) +} + /// Decode a row from the columnar wire format back into Values. +/// +/// `Err(WalRowCorrupt)` when the bytes do not decode. The corruption is +/// reported here, where it is detected. pub fn decode_row_from_wal( data: &[u8], ) -> Result, crate::error::ColumnarError> { + decode_row(data).inspect_err(crate::diag::wal_row_corrupt) +} + +/// The body of [`decode_row_from_wal`]. +fn decode_row(data: &[u8]) -> Result, crate::error::ColumnarError> { use nodedb_types::value::Value; let mut values = Vec::new(); let mut cursor = 0; while cursor < data.len() { - let tag_slice = read_slice(data, &mut cursor, 1, "tag")?; - let tag = tag_slice[0]; + let tag_at = cursor; + let [tag] = read_array::<1>(data, &mut cursor, "tag")?; let value = match tag { 0 => Value::Null, - 1 => { - let bytes = read_slice(data, &mut cursor, 8, "i64")?; - let v = i64::from_le_bytes(bytes.try_into().map_err(|_| { - crate::error::ColumnarError::Serialization("truncated i64".into()) - })?); - Value::Integer(v) - } - 2 => { - let bytes = read_slice(data, &mut cursor, 8, "f64")?; - let v = f64::from_le_bytes(bytes.try_into().map_err(|_| { - crate::error::ColumnarError::Serialization("truncated f64".into()) - })?); - Value::Float(v) - } + 1 => Value::Integer(i64::from_le_bytes(read_array(data, &mut cursor, "i64")?)), + 2 => Value::Float(f64::from_le_bytes(read_array(data, &mut cursor, "f64")?)), 3 => { - let bytes = read_slice(data, &mut cursor, 1, "bool")?; - Value::Bool(bytes[0] != 0) - } - 4 | 5 | 8 => { - let bytes = read_length_prefixed( - data, - &mut cursor, - match tag { - 4 => "string", - 5 => "bytes", - 8 => "uuid", - _ => unreachable!(), - }, - )?; - match tag { - 4 => Value::String(String::from_utf8_lossy(bytes).into_owned()), - 5 => Value::Bytes(bytes.to_vec()), - 8 => Value::Uuid(String::from_utf8_lossy(bytes).into_owned()), - _ => unreachable!(), - } + let [b] = read_array::<1>(data, &mut cursor, "bool")?; + Value::Bool(b != 0) } + 4 => Value::String(read_utf8(data, &mut cursor, "string")?), + 5 => Value::Bytes(read_length_prefixed(data, &mut cursor, "bytes")?.to_vec()), + 8 => Value::Uuid(read_utf8(data, &mut cursor, "uuid")?), 6 => { - let bytes = read_slice(data, &mut cursor, 8, "timestamp")?; - let micros = i64::from_le_bytes(bytes.try_into().map_err(|_| { - crate::error::ColumnarError::Serialization("truncated timestamp".into()) - })?); + let micros = i64::from_le_bytes(read_array(data, &mut cursor, "timestamp")?); Value::DateTime(nodedb_types::datetime::NdbDateTime::from_micros(micros)) } - 7 => { - let bytes = read_slice(data, &mut cursor, 16, "decimal")?; - let mut arr = [0u8; 16]; - arr.copy_from_slice(bytes); - Value::Decimal(rust_decimal::Decimal::deserialize(arr)) - } + 7 => Value::Decimal(rust_decimal::Decimal::deserialize(read_array( + data, + &mut cursor, + "decimal", + )?)), 9 => { - let count_bytes = read_slice(data, &mut cursor, 4, "vector count")?; - let count = u32::from_le_bytes(count_bytes.try_into().map_err(|_| { - crate::error::ColumnarError::Serialization("truncated vector count".into()) - })?) as usize; - let remaining_values = data.len().saturating_sub(cursor) / 4; - let max_count = (MAX_FIELD_LEN / 4).min(remaining_values); + let count_at = cursor; + let count = + u32::from_le_bytes(read_array(data, &mut cursor, "vector count")?) as usize; + let remaining_values = data.len().saturating_sub(cursor) / 8; + let max_count = (MAX_FIELD_LEN / 8).min(remaining_values); if count > max_count { - return Err(crate::error::ColumnarError::Serialization(format!( - "vector count {count} exceeds maximum {max_count}" - ))); + return Err(corrupt( + count_at, + format!("vector count {count} exceeds maximum {max_count}"), + )); } let capacity = checked_decode_capacity( count, size_of::(), data.len().saturating_sub(cursor), - 4, + 8, max_count, usize::MAX, ) .ok_or_else(|| { - crate::error::ColumnarError::Serialization( - "vector count exceeds decode allocation bounds".into(), - ) + corrupt(count_at, "vector count exceeds decode allocation bounds") })?; let mut arr = Vec::with_capacity(capacity); for _ in 0..count { - let fb = read_slice(data, &mut cursor, 4, "vector f32")?; - let f = f32::from_le_bytes(fb.try_into().map_err(|_| { - crate::error::ColumnarError::Serialization("truncated f32".into()) - })?); - arr.push(Value::Float(f as f64)); + let f = f64::from_le_bytes(read_array(data, &mut cursor, "vector f64")?); + arr.push(Value::Float(f)); } Value::Array(arr) } 10 => { + let json_at = cursor; let json_bytes = read_length_prefixed(data, &mut cursor, "json")?; - sonic_rs::from_slice(json_bytes).unwrap_or(Value::Null) - } - _ => { - return Err(crate::error::ColumnarError::Serialization(format!( - "unknown WAL value tag: {tag}" - ))); + sonic_rs::from_slice(json_bytes) + .map_err(|e| corrupt(json_at, format!("json value does not decode: {e}")))? } + _ => return Err(corrupt(tag_at, format!("unknown WAL value tag: {tag}"))), }; values.push(value); @@ -339,10 +345,54 @@ mod tests { bytes.extend_from_slice(&u32::MAX.to_le_bytes()); assert!(matches!( decode_row_from_wal(&bytes), - Err(crate::error::ColumnarError::Serialization(_)) + Err(crate::error::ColumnarError::WalRowCorrupt { .. }) )); } + #[test] + fn a_string_that_is_not_utf8_is_refused() { + let mut bytes = vec![4]; + bytes.extend_from_slice(&2u32.to_le_bytes()); + bytes.extend_from_slice(&[0xC3, 0x28]); + assert!(matches!( + decode_row_from_wal(&bytes), + Err(crate::error::ColumnarError::WalRowCorrupt { offset: 1, .. }) + )); + } + + #[test] + fn a_uuid_that_is_not_utf8_is_refused() { + let mut bytes = vec![8]; + bytes.extend_from_slice(&1u32.to_le_bytes()); + bytes.push(0xFF); + assert!(matches!( + decode_row_from_wal(&bytes), + Err(crate::error::ColumnarError::WalRowCorrupt { .. }) + )); + } + + #[test] + fn a_json_value_that_does_not_decode_is_refused() { + let json = b"{not json"; + let mut bytes = vec![10]; + bytes.extend_from_slice(&(json.len() as u32).to_le_bytes()); + bytes.extend_from_slice(json); + assert!(matches!( + decode_row_from_wal(&bytes), + Err(crate::error::ColumnarError::WalRowCorrupt { .. }) + )); + } + + #[test] + fn arrays_decode_to_the_values_they_were() { + let values = vec![ + Value::Array(vec![Value::Float(0.1), Value::Float(-2.5)]), + Value::Array(vec![Value::String("a".into()), Value::String("b".into())]), + ]; + let encoded = encode_row_for_wal(&values).expect("encode"); + assert_eq!(decode_row_from_wal(&encoded).expect("decode"), values); + } + #[test] fn wal_record_roundtrip() { let records = vec![ diff --git a/nodedb-columnar/src/writer/block.rs b/nodedb-columnar/src/writer/block.rs index 3cf93b5a3..d70d064a9 100644 --- a/nodedb-columnar/src/writer/block.rs +++ b/nodedb-columnar/src/writer/block.rs @@ -4,23 +4,41 @@ use nodedb_codec::{ColumnCodec, ResolvedColumnCodec}; use nodedb_mem::ScopedMemory; -use nodedb_types::columnar::ColumnType; use crate::error::ColumnarError; -use crate::format::{BLOCK_SIZE, BlockStats}; +use crate::format::{BLOCK_SIZE, BlockLayout, BlockStats}; use crate::memtable::ColumnData; use super::encode::{ encode_f64_with_validity, encode_i64_with_validity, encode_validity_bitmap, prepend_validity, }; -use super::stats::{compute_string_block_stats, numeric_min_max_f64, numeric_min_max_i64}; +use super::stats::{ + StringBlock, compute_string_block_stats, numeric_min_max_f64, numeric_min_max_i64, +}; + +/// The block layout `encode_single_block` writes for `col_data`. +pub(super) fn block_layout(col_data: &ColumnData) -> BlockLayout { + match col_data { + ColumnData::Int64 { .. } | ColumnData::Timestamp { .. } => BlockLayout::Int64, + ColumnData::Float64 { .. } => BlockLayout::Float64, + ColumnData::Bool { .. } => BlockLayout::PackedBool, + ColumnData::String { .. } + | ColumnData::Bytes { .. } + | ColumnData::Json { .. } + | ColumnData::Geometry { .. } => BlockLayout::VarLen, + ColumnData::Decimal { .. } | ColumnData::Uuid { .. } | ColumnData::Vector { .. } => { + BlockLayout::FixedWidth + } + ColumnData::DictEncoded { .. } => BlockLayout::DictIds, + } +} /// Encode all blocks for a single column, appending to `buf`. /// Returns per-block statistics. pub(super) fn encode_column_blocks( buf: &mut Vec, + col_name: &str, col_data: &ColumnData, - col_type: &ColumnType, codec: ResolvedColumnCodec, row_count: usize, memory: &ScopedMemory, @@ -35,8 +53,8 @@ pub(super) fn encode_column_blocks( let block_row_count = end - start; let (compressed, stats) = encode_single_block( + col_name, col_data, - col_type, codec, start, end, @@ -56,9 +74,12 @@ pub(super) fn encode_column_blocks( } /// Encode a single block of rows for a column. +/// +/// Only string columns carry min/max bounds. A bytes column carries none, +/// so no predicate prunes its blocks by a byte compare. fn encode_single_block( + col_name: &str, col_data: &ColumnData, - _col_type: &ColumnType, codec: ResolvedColumnCodec, start: usize, end: usize, @@ -138,7 +159,8 @@ fn encode_single_block( let compressed = nodedb_codec::encode_bytes_pipeline(string_bytes, codec.into_column_codec())?; - let stats = compute_string_block_stats( + let stats = compute_string_block_stats(StringBlock { + column: col_name, data, offsets, valid_slice, @@ -146,7 +168,7 @@ fn encode_single_block( end, null_count, block_row_count, - ); + })?; let block_offsets: Vec = offsets[start..=end] .iter() diff --git a/nodedb-columnar/src/writer/segment_writer.rs b/nodedb-columnar/src/writer/segment_writer.rs index a39d4938a..beab51176 100644 --- a/nodedb-columnar/src/writer/segment_writer.rs +++ b/nodedb-columnar/src/writer/segment_writer.rs @@ -10,7 +10,7 @@ use crate::error::ColumnarError; use crate::format::{ColumnMeta, HEADER_SIZE, SegmentFooter, SegmentHeader}; use crate::memtable::ColumnData; -use super::block::encode_column_blocks; +use super::block::{block_layout, encode_column_blocks}; use super::encode::compute_schema_hash; /// Profile tag values for the segment footer. @@ -83,8 +83,8 @@ impl SegmentWriter { // Encode blocks. let block_stats = encode_column_blocks( &mut buf, + &col_def.name, col_data, - &col_def.column_type, codec, row_count, &self.memory, @@ -107,6 +107,7 @@ impl SegmentWriter { offset: col_start - HEADER_SIZE as u64, length: col_end - col_start, codec: effective_codec, + layout: block_layout(col_data), block_count: block_stats.len() as u32, block_stats, dictionary, diff --git a/nodedb-columnar/src/writer/stats.rs b/nodedb-columnar/src/writer/stats.rs index 619cb7c8d..a723bc36f 100644 --- a/nodedb-columnar/src/writer/stats.rs +++ b/nodedb-columnar/src/writer/stats.rs @@ -2,6 +2,7 @@ //! Statistics computation helpers: numeric min/max, string zone-maps, bloom filters. +use crate::error::ColumnarError; use crate::format::{BlockStats, BloomFilter}; use crate::predicate::{BLOOM_BITS_DEFAULT, BLOOM_K_DEFAULT, bloom_insert}; @@ -13,26 +14,47 @@ pub(super) const STRING_BOUND_MAX_BYTES: usize = 32; /// with ≤16 distinct values get no bloom. const BLOOM_DISTINCT_THRESHOLD: usize = 16; +/// The cells of one string column block, for its statistics. +pub(super) struct StringBlock<'a> { + /// Column name, for the error a non-UTF-8 cell returns. + pub column: &'a str, + pub data: &'a [u8], + pub offsets: &'a [u32], + pub valid_slice: &'a [bool], + pub start: usize, + pub end: usize, + pub null_count: u32, + pub block_row_count: usize, +} + /// Compute `BlockStats` for a string column block. /// /// Iterates over all non-null values in `[start, end)`, building lexicographic /// min/max (each truncated to `STRING_BOUND_MAX_BYTES` bytes) and a 256-byte /// bloom filter for equality-predicate fast-reject. +/// +/// `Err(StringCellNotUtf8)` for a cell that is not UTF-8. The memtable push +/// stores only UTF-8, so such a cell is corruption. Bounds built from it +/// would prune rows a predicate matches. pub(super) fn compute_string_block_stats( - data: &[u8], - offsets: &[u32], - valid_slice: &[bool], - start: usize, - end: usize, - null_count: u32, - block_row_count: usize, -) -> BlockStats { - let mut str_min: Option = None; - let mut str_max: Option = None; + block: StringBlock<'_>, +) -> Result { + let StringBlock { + column, + data, + offsets, + valid_slice, + start, + end, + null_count, + block_row_count, + } = block; + let mut cell_min: Option<&str> = None; + let mut cell_max: Option<&str> = None; let mut distinct = std::collections::HashSet::new(); - let mut has_non_null = false; + let mut cells: Vec<&str> = Vec::with_capacity(end - start); - // First pass: compute min/max and count distinct values. + // First pass: check UTF-8, compute min/max and count distinct values. for (&is_valid, row_idx) in valid_slice.iter().zip(start..end) { if !is_valid { continue; @@ -40,58 +62,56 @@ pub(super) fn compute_string_block_stats( let b_start = offsets[row_idx] as usize; let b_end = offsets[row_idx + 1] as usize; let raw = &data[b_start..b_end]; - let s = std::str::from_utf8(raw).unwrap_or(""); - - has_non_null = true; + let s = std::str::from_utf8(raw).map_err(|_| { + let err = ColumnarError::StringCellNotUtf8 { + column: column.to_string(), + row: row_idx, + }; + crate::diag::string_cell_not_utf8(&err); + err + })?; + cells.push(s); // Track distinct values up to threshold+1 (stop counting early). if distinct.len() <= BLOOM_DISTINCT_THRESHOLD { distinct.insert(raw); } - let truncated = truncate_to_char_boundary(s, STRING_BOUND_MAX_BYTES); - - match &str_min { - None => str_min = Some(truncated.to_owned()), - Some(cur) if truncated < cur.as_str() => str_min = Some(truncated.to_owned()), - _ => {} + if cell_min.is_none_or(|cur| s < cur) { + cell_min = Some(s); } - match &str_max { - None => str_max = Some(truncated.to_owned()), - Some(cur) if truncated > cur.as_str() => str_max = Some(truncated.to_owned()), - _ => {} + if cell_max.is_none_or(|cur| s > cur) { + cell_max = Some(s); } } + // A prefix of the least cell is a lower bound. The upper bound must not + // fall below the greatest cell, so a truncated maximum is raised. + let str_min = cell_min.map(|s| truncate_to_char_boundary(s, STRING_BOUND_MAX_BYTES).to_owned()); + let str_max = cell_max.and_then(|s| upper_bound_prefix(s, STRING_BOUND_MAX_BYTES)); // Only build bloom filter for high-cardinality blocks where zone maps // alone cannot efficiently prune. Low-cardinality columns (≤16 distinct) // are better served by dict encoding + integer comparison. - let bloom_opt = if has_non_null && distinct.len() > BLOOM_DISTINCT_THRESHOLD { + let bloom_opt = if !cells.is_empty() && distinct.len() > BLOOM_DISTINCT_THRESHOLD { let byte_count = (BLOOM_BITS_DEFAULT as usize).div_ceil(8); let mut bloom = BloomFilter { k: BLOOM_K_DEFAULT, m: BLOOM_BITS_DEFAULT, bytes: vec![0u8; byte_count], }; - for (&is_valid, row_idx) in valid_slice.iter().zip(start..end) { - if !is_valid { - continue; - } - let b_start = offsets[row_idx] as usize; - let b_end = offsets[row_idx + 1] as usize; - let s = std::str::from_utf8(&data[b_start..b_end]).unwrap_or(""); + for s in &cells { bloom_insert(&mut bloom, s); } Some(bloom) } else { None }; - BlockStats::string_block( + Ok(BlockStats::string_block( null_count, block_row_count as u32, str_min, str_max, bloom_opt, - ) + )) } /// Truncate a string to at most `max_bytes` bytes, preserving valid UTF-8 by @@ -107,6 +127,36 @@ pub(super) fn truncate_to_char_boundary(s: &str, max_bytes: usize) -> &str { &s[..boundary] } +/// A string of about `max_bytes` bytes that is not less than `s`. +/// +/// `s` itself when it fits. Otherwise its prefix with the last character +/// raised by one, so every string with that prefix sorts below it. `None` +/// when no character of the prefix can be raised: the block then carries +/// no upper bound. +pub(super) fn upper_bound_prefix(s: &str, max_bytes: usize) -> Option { + if s.len() <= max_bytes { + return Some(s.to_owned()); + } + let mut chars: Vec = truncate_to_char_boundary(s, max_bytes).chars().collect(); + while let Some(last) = chars.pop() { + if let Some(next) = next_char(last) { + chars.push(next); + return Some(chars.into_iter().collect()); + } + } + None +} + +/// The character after `c` in code point order, skipping the surrogate +/// range. `None` for `char::MAX`. +fn next_char(c: char) -> Option { + let mut code = u32::from(c) + 1; + if (0xD800..=0xDFFF).contains(&code) { + code = 0xE000; + } + char::from_u32(code) +} + /// Compute min/max for i64 values, skipping nulls. pub(super) fn numeric_min_max_i64(values: &[i64], valid: &[bool]) -> (i64, i64) { let mut min = i64::MAX; @@ -144,3 +194,61 @@ pub(super) fn numeric_min_max_f64(values: &[f64], valid: &[bool]) -> (f64, f64) (min, max) } } + +#[cfg(test)] +mod tests { + use super::*; + + fn stats_of(cells: &[&[u8]]) -> Result { + let mut data = Vec::new(); + let mut offsets = vec![0u32]; + for cell in cells { + data.extend_from_slice(cell); + offsets.push(data.len() as u32); + } + let valid = vec![true; cells.len()]; + compute_string_block_stats(StringBlock { + column: "s", + data: &data, + offsets: &offsets, + valid_slice: &valid, + start: 0, + end: cells.len(), + null_count: 0, + block_row_count: cells.len(), + }) + } + + #[test] + fn a_cell_that_is_not_utf8_refuses_the_block() { + let result = stats_of(&[b"ok", &[0xC3, 0x28]]); + assert!(matches!( + result, + Err(ColumnarError::StringCellNotUtf8 { ref column, row: 1 }) if column == "s" + )); + } + + #[test] + fn a_long_maximum_keeps_an_upper_bound_above_every_cell() { + let long = "a".repeat(STRING_BOUND_MAX_BYTES + 8); + let stats = stats_of(&[b"a", long.as_bytes()]).expect("stats"); + let max = stats.str_max.expect("max"); + assert!(max.as_str() >= long.as_str(), "{max} is below {long}"); + assert_eq!(stats.str_min.as_deref(), Some("a")); + } + + #[test] + fn the_upper_bound_raises_the_last_raisable_character() { + let s = format!("{}{}", "b".repeat(STRING_BOUND_MAX_BYTES - 1), char::MAX); + let tail = format!("{s}zzz"); + assert_eq!( + upper_bound_prefix(&tail, STRING_BOUND_MAX_BYTES + 3), + Some(format!("{}c", "b".repeat(STRING_BOUND_MAX_BYTES - 2))) + ); + assert_eq!(upper_bound_prefix("short", 32), Some("short".to_string())); + assert_eq!( + upper_bound_prefix(&char::MAX.to_string().repeat(3), 4), + None + ); + } +} diff --git a/nodedb-crdt/src/error.rs b/nodedb-crdt/src/error.rs index 8257fa08c..66bd26f27 100644 --- a/nodedb-crdt/src/error.rs +++ b/nodedb-crdt/src/error.rs @@ -35,13 +35,15 @@ pub enum CrdtError { #[error("CRDT import has too many operations: {actual} > {limit}")] ImportOperationLimitExceeded { limit: usize, actual: usize }, - /// The imported update depends on operations this document has never seen, - /// so Loro buffered them as causally pending instead of applying them. + /// The imported blob carries changes that depend on operations this + /// document has never seen, so Loro buffered those changes as causally + /// pending instead of applying them. /// - /// The document state did NOT advance. Reporting such an import as success - /// is silent data loss: the caller acknowledges a write that was never - /// applied and may never be, since the missing predecessors are not part of - /// this document's operation history. + /// The ready changes in the same blob still apply, so the document state + /// can advance under this error. The pending changes do not. Reporting + /// such an import as success is silent data loss: the caller acknowledges + /// a write that was never applied and may never be, since the missing + /// predecessors are not part of this document's operation history. #[error("CRDT import depends on operations absent from this document")] ImportPendingDependencies, @@ -137,6 +139,33 @@ pub enum CrdtError { field: String, }, + /// A delta writes a root container that is not a map. + /// + /// Every collection is a root map of row maps. A root text, list, movable + /// list, tree, or counter holds no rows, so no constraint can check it. + #[error( + "CRDT delta writes root {container_type} container `{container}`; \ + a collection is a root map of row maps" + )] + NonMapRootContainer { + container: String, + container_type: String, + }, + + /// A delta sets a collection row to a value that is not a map. + /// + /// A row is a map of fields. Any other row value has no fields, so NOT + /// NULL and CHECK constraints cannot run on it. + #[error( + "CRDT delta sets row `{row_id}` in collection `{collection}` to {value}; \ + a row is a map of fields" + )] + NonMapRowValue { + collection: String, + row_id: String, + value: String, + }, + /// Auth context has expired — agent must re-authenticate before syncing. #[error("auth expired: user {user_id} must re-authenticate (expired at {expired_at})")] AuthExpired { user_id: u64, expired_at: u64 }, diff --git a/nodedb-crdt/src/lib.rs b/nodedb-crdt/src/lib.rs index d38961cb7..2fd97b4ac 100644 --- a/nodedb-crdt/src/lib.rs +++ b/nodedb-crdt/src/lib.rs @@ -46,6 +46,9 @@ pub use policy::{ }; pub use row_lookup::RowLookup; pub use signing::{DeltaSigner, DeviceRegistry}; -pub use state::{CrdtDeltaPreview, CrdtDeltaPreviewLimits, CrdtState, ImportAdmission}; +pub use state::{ + CrdtDeltaPreview, CrdtDeltaPreviewLimits, CrdtState, ImportAdmission, TrackedImport, + WriteSetImport, +}; pub use state::{DEFAULT_MAX_DELTA_BYTES, DEFAULT_MAX_POST_IMAGE_BYTES}; pub use validator::{ValidationOutcome, Validator}; diff --git a/nodedb-crdt/src/state/changed_rows.rs b/nodedb-crdt/src/state/changed_rows.rs new file mode 100644 index 000000000..23b10162a --- /dev/null +++ b/nodedb-crdt/src/state/changed_rows.rs @@ -0,0 +1,373 @@ +// SPDX-License-Identifier: Apache-2.0 + +//! Imports that report which rows of one collection they changed. +//! +//! A caller that keeps a derived index over rows (full-text, spatial) needs +//! the rows an import moved, without a rescan of the collection. +//! +//! The rows come from the operations the import added to the oplog, read the +//! same way as an import's write-set (see [`super::write_set`]). No +//! subscription, second diff, or checkout runs, so a shallow document works +//! like any other and later local writes pay nothing for the tracking. + +use std::collections::BTreeSet; + +use super::core::CrdtState; +use super::import_admission::ImportAdmission; +use super::write_set::collect_imported_rows; +use crate::error::Result; + +/// The outcome of a tracked import and the rows it changed. +/// +/// `changed_rows` is filled for every outcome. A blob can apply its ready +/// changes and still return `ImportPendingDependencies`. Derived indexes +/// must follow the rows that did apply. +#[derive(Debug)] +#[must_use] +pub struct TrackedImport { + /// The import result, as the untracked import returns it. + pub outcome: Result, + /// Row ids of the collection an applied operation of the import wrote: + /// rows added, removed, replaced, or changed in any nested container. + pub changed_rows: BTreeSet, +} + +impl CrdtState { + /// [`Self::import`] that also reports the rows of `collection` it changed. + pub fn import_tracked(&self, collection: &str, data: &[u8]) -> TrackedImport { + self.track_changed_rows(collection, || self.import(data)) + } + + /// [`Self::import_local`] that also reports the rows of `collection` it changed. + pub fn import_local_tracked(&self, collection: &str, data: &[u8]) -> TrackedImport { + self.track_changed_rows(collection, || self.import_local(data)) + } + + /// Run `import` and collect the rows of `collection` it wrote. + fn track_changed_rows( + &self, + collection: &str, + import: impl FnOnce() -> Result, + ) -> TrackedImport { + // Shape faults are the write-set's concern. A derived index follows + // every row the import moved, whatever shape the delta wrote. + let (outcome, imported) = collect_imported_rows(&self.doc, Some(collection), import); + TrackedImport { + outcome, + changed_rows: imported + .rows + .into_iter() + .map(|(_, row_id)| row_id) + .collect(), + } + } +} + +#[cfg(test)] +mod tests { + use loro::{ExportMode, IdSpan, LoroValue}; + + use super::CrdtState; + use crate::error::CrdtError; + + fn put(state: &CrdtState, row: &str, text: &str) { + state + .upsert("docs", row, &[("body", LoroValue::from(text))]) + .expect("upsert"); + } + + fn rows(tracked: &super::TrackedImport) -> Vec<&str> { + tracked.changed_rows.iter().map(String::as_str).collect() + } + + /// A source with rows `a`, `b`, `c` and a target holding the same state. + fn synced_pair() -> (CrdtState, CrdtState) { + let source = CrdtState::new(1).expect("source"); + put(&source, "a", "one"); + put(&source, "b", "two"); + put(&source, "c", "three"); + let target = CrdtState::new(2).expect("target"); + target + .import(&source.export_snapshot().expect("snapshot")) + .expect("seed import"); + (source, target) + } + + #[test] + fn a_first_import_reports_every_row() { + let source = CrdtState::new(1).expect("source"); + put(&source, "a", "one"); + put(&source, "b", "two"); + let target = CrdtState::new(2).expect("target"); + + let tracked = target.import_tracked("docs", &source.export_snapshot().expect("snapshot")); + tracked.outcome.as_ref().expect("import"); + assert_eq!(rows(&tracked), vec!["a", "b"]); + } + + #[test] + fn a_delta_reports_exactly_the_rows_it_touched() { + let (source, target) = synced_pair(); + let before = source.oplog_version_vector(); + put(&source, "b", "changed"); + source.delete("docs", "a").expect("delete"); + let delta = source.export_updates_since(&before).expect("delta"); + + let tracked = target.import_tracked("docs", &delta); + tracked.outcome.as_ref().expect("import"); + assert_eq!(rows(&tracked), vec!["a", "b"]); + assert!(!target.row_exists("docs", "a")); + } + + #[test] + fn a_shallow_snapshot_and_a_following_delta_report_their_rows() { + let mut source = CrdtState::new(1).expect("source"); + put(&source, "a", "one"); + put(&source, "b", "two"); + put(&source, "c", "three"); + source.compact_history().expect("compact"); + let target = CrdtState::new(2).expect("target"); + + let tracked = + target.import_local_tracked("docs", &source.export_snapshot().expect("snapshot")); + tracked.outcome.as_ref().expect("shallow import"); + assert_eq!(rows(&tracked), vec!["a", "b", "c"]); + assert!(target.doc.is_shallow()); + + let before = source.oplog_version_vector(); + put(&source, "d", "four"); + source.delete("docs", "a").expect("delete"); + let delta = source.export_updates_since(&before).expect("delta"); + + let tracked = target.import_tracked("docs", &delta); + tracked.outcome.as_ref().expect("delta into shallow doc"); + assert_eq!(rows(&tracked), vec!["a", "d"]); + } + + #[test] + fn a_shallow_snapshot_into_a_synced_doc_reports_only_rows_it_changed() { + let (mut source, target) = synced_pair(); + source.compact_history().expect("compact"); + put(&source, "b", "changed"); + let snapshot = source.export_snapshot().expect("snapshot"); + assert!(source.doc.is_shallow()); + + let tracked = target.import_tracked("docs", &snapshot); + tracked.outcome.as_ref().expect("import"); + assert_eq!(rows(&tracked), vec!["b"]); + assert_eq!( + target.read_field("docs", "b", "body"), + Some(LoroValue::from("changed")) + ); + } + + #[test] + fn a_nested_container_change_names_its_row() { + let (source, target) = synced_pair(); + let before = source.oplog_version_vector(); + source + .list_insert_fields( + "docs", + "a", + "blocks", + 0, + &[("text".to_owned(), LoroValue::from("block"))], + ) + .expect("list insert"); + let delta = source.export_updates_since(&before).expect("delta"); + + let tracked = target.import_tracked("docs", &delta); + tracked.outcome.as_ref().expect("import"); + assert_eq!(rows(&tracked), vec!["a"]); + } + + #[test] + fn a_row_deleted_and_recreated_in_one_delta_is_reported() { + let (source, target) = synced_pair(); + let before = source.oplog_version_vector(); + source.delete("docs", "a").expect("delete"); + put(&source, "a", "reborn"); + let delta = source.export_updates_since(&before).expect("delta"); + + let tracked = target.import_tracked("docs", &delta); + tracked.outcome.as_ref().expect("import"); + assert_eq!(rows(&tracked), vec!["a"]); + assert_eq!( + target.read_field("docs", "a", "body"), + Some(LoroValue::from("reborn")) + ); + } + + #[test] + fn rows_of_another_collection_are_not_reported() { + let source = CrdtState::new(1).expect("source"); + put(&source, "a", "one"); + source + .upsert("other", "x", &[("body", LoroValue::from("elsewhere"))]) + .expect("other upsert"); + let target = CrdtState::new(2).expect("target"); + + let tracked = target.import_tracked("docs", &source.export_snapshot().expect("snapshot")); + tracked.outcome.as_ref().expect("first import"); + assert_eq!(rows(&tracked), vec!["a"]); + + let before = source.oplog_version_vector(); + source + .upsert("other", "y", &[("body", LoroValue::from("elsewhere"))]) + .expect("other upsert"); + put(&source, "b", "two"); + let delta = source.export_updates_since(&before).expect("delta"); + + let tracked = target.import_tracked("docs", &delta); + tracked.outcome.as_ref().expect("delta"); + assert_eq!(rows(&tracked), vec!["b"]); + } + + #[test] + fn an_uncommitted_local_write_is_not_reported_as_imported() { + let (source, target) = synced_pair(); + put(&target, "local", "mine"); + assert_ne!(target.doc.get_pending_txn_len(), 0); + let before = source.oplog_version_vector(); + put(&source, "c", "changed"); + let delta = source.export_updates_since(&before).expect("delta"); + + let tracked = target.import_tracked("docs", &delta); + tracked.outcome.as_ref().expect("import"); + assert_eq!(rows(&tracked), vec!["c"]); + assert!(target.row_exists("docs", "local")); + } + + #[test] + fn an_uncommitted_local_write_on_an_empty_doc_is_not_reported() { + let source = CrdtState::new(1).expect("source"); + put(&source, "a", "one"); + let target = CrdtState::new(2).expect("target"); + put(&target, "local", "mine"); + + let tracked = target.import_tracked("docs", &source.export_snapshot().expect("snapshot")); + tracked.outcome.as_ref().expect("import"); + assert_eq!(rows(&tracked), vec!["a"]); + } + + #[test] + fn a_replayed_delta_reports_no_rows() { + let (source, target) = synced_pair(); + let before = source.oplog_version_vector(); + put(&source, "b", "changed"); + let delta = source.export_updates_since(&before).expect("delta"); + let first = target.import_tracked("docs", &delta); + first.outcome.as_ref().expect("first import"); + + let replay = target.import_tracked("docs", &delta); + replay.outcome.as_ref().expect("replay"); + assert!(replay.changed_rows.is_empty()); + } + + #[test] + fn a_tracked_import_leaves_later_local_writes_unchanged() { + let (source, tracked) = synced_pair(); + let plain = CrdtState::new(2).expect("plain target"); + plain + .import(&source.export_snapshot().expect("snapshot")) + .expect("seed import"); + let before = source.oplog_version_vector(); + put(&source, "b", "changed"); + let delta = source.export_updates_since(&before).expect("delta"); + + let imported = tracked.import_tracked("docs", &delta); + imported.outcome.as_ref().expect("tracked import"); + assert_eq!(rows(&imported), vec!["b"]); + plain.import(&delta).expect("plain import"); + + let local_write = |state: &CrdtState| { + let before = state.oplog_version_vector(); + put(state, "local", "mine"); + state.export_updates_since(&before).expect("local delta") + }; + assert_eq!(local_write(&tracked), local_write(&plain)); + } + + /// An edit inside an existing list-valued row names that row, as the + /// first import of the row does. The write-set still refuses the + /// non-map row. + #[test] + fn an_edit_inside_a_list_valued_row_reports_the_row() { + let source = CrdtState::new(1).expect("source"); + let list = source + .doc + .get_map("docs") + .insert_container("l", loro::LoroList::new()) + .expect("list row"); + list.push("first").expect("push"); + source.doc.commit(); + let seed = source.export_snapshot().expect("seed snapshot"); + + let target = CrdtState::new(2).expect("target"); + let tracked = target.import_tracked("docs", &seed); + tracked.outcome.as_ref().expect("seed import"); + assert_eq!(rows(&tracked), vec!["l"]); + + let before = source.oplog_version_vector(); + list.push("second").expect("push"); + source.doc.commit(); + let delta = source.export_updates_since(&before).expect("delta"); + + let tracked = target.import_tracked("docs", &delta); + tracked.outcome.as_ref().expect("delta import"); + assert_eq!(rows(&tracked), vec!["l"]); + + let checked = CrdtState::new(3).expect("checked target"); + checked.import(&seed).expect("seed import"); + let imported = checked.import_with_write_set(&delta); + imported.outcome.as_ref().expect("delta import"); + assert!( + matches!( + imported.write_set, + Err(CrdtError::NonMapRowValue { ref row_id, .. }) if row_id == "l" + ), + "{:?}", + imported.write_set + ); + } + + #[test] + fn a_partially_pending_import_reports_the_rows_that_applied() { + let peer1 = CrdtState::new(1).expect("peer 1"); + put(&peer1, "x", "ready"); + let peer3 = CrdtState::new(3).expect("peer 3"); + put(&peer3, "y", "withheld"); + let withheld = peer3.export_snapshot().expect("peer 3 snapshot"); + + // The hub's own write to `z` depends on peer 3's history. + let hub = CrdtState::new(2).expect("hub"); + hub.import(&peer1.export_snapshot().expect("peer 1 snapshot")) + .expect("hub imports peer 1"); + hub.import(&withheld).expect("hub imports peer 3"); + put(&hub, "z", "dependent"); + hub.doc.commit(); + let vv = hub.oplog_version_vector(); + let end = |peer: u64| vv.get(&peer).copied().unwrap_or(0); + let blob = hub + .doc + .export(ExportMode::updates_in_range(vec![ + IdSpan::new(1, 0, end(1)), + IdSpan::new(2, 0, end(2)), + ])) + .expect("range export"); + + let target = CrdtState::new(4).expect("target"); + let tracked = target.import_tracked("docs", &blob); + assert!(matches!( + tracked.outcome, + Err(CrdtError::ImportPendingDependencies) + )); + assert_eq!(rows(&tracked), vec!["x"]); + assert!(target.row_exists("docs", "x")); + + let tracked = target.import_tracked("docs", &withheld); + tracked.outcome.as_ref().expect("withheld import"); + assert_eq!(rows(&tracked), vec!["y", "z"]); + assert!(target.row_exists("docs", "z")); + } +} diff --git a/nodedb-crdt/src/state/mod.rs b/nodedb-crdt/src/state/mod.rs index cfef620a9..2ed961b06 100644 --- a/nodedb-crdt/src/state/mod.rs +++ b/nodedb-crdt/src/state/mod.rs @@ -7,6 +7,7 @@ //! where each row is itself a `LoroMap` of field→value. pub mod bitemporal_archive; +pub mod changed_rows; pub mod core; pub(crate) mod document_cell; pub mod frontier_digest; @@ -14,10 +15,13 @@ pub mod history; pub(crate) mod import_admission; pub mod preview; pub mod rekey; +pub mod remove_fields; pub(crate) mod restore_containers; +pub mod row_image; pub mod snapshot; pub mod write_set; +pub use changed_rows::TrackedImport; pub use core::CrdtState; pub use import_admission::{ CrdtImportLimits, DEFAULT_MAX_IMPORT_BYTES, DEFAULT_MAX_IMPORT_OPS, ImportAdmission, @@ -26,3 +30,5 @@ pub use preview::{ CrdtDeltaPreview, CrdtDeltaPreviewLimits, DEFAULT_MAX_DELTA_BYTES, DEFAULT_MAX_ENCODED_DELTA_OPS, DEFAULT_MAX_POST_IMAGE_BYTES, }; +pub use row_image::RowImage; +pub use write_set::WriteSetImport; diff --git a/nodedb-crdt/src/state/preview.rs b/nodedb-crdt/src/state/preview.rs index 3b880574b..27d04bb75 100644 --- a/nodedb-crdt/src/state/preview.rs +++ b/nodedb-crdt/src/state/preview.rs @@ -18,6 +18,7 @@ use crate::loro_value::loro_to_value; use super::core::CrdtState; use super::document_cell::DocumentCell; use super::import_admission::{CrdtImportLimits, admit_import}; +use super::write_set::collect_write_set; /// Maximum raw CRDT delta bytes accepted by the default authoritative preview. pub const DEFAULT_MAX_DELTA_BYTES: usize = 1024 * 1024; @@ -128,12 +129,16 @@ impl CrdtState { // The source is quiescent, so Loro's fork cannot publish source state. let fork = self.doc.fork(); - let before_frontier = fork.state_frontiers(); let before_oplog = fork.oplog_vv(); if before_oplog != authoritative_oplog { return Err(CrdtError::PreviewInvalidOperationRange); } - let status = fork.import(delta).map_err(|error| match error { + // The write-set comes from the operations the import adds to the fork. + // It runs after byte and imported-operation caps have bounded the + // import work. `max_write_set_entries` is semantic cardinality + // enforcement, not an allocation short-circuit. + let (imported, write_set) = collect_write_set(&fork, || fork.import(delta)); + let status = imported.map_err(|error| match error { loro::LoroError::ImportUpdatesThatDependsOnOutdatedVersion => { CrdtError::PreviewPendingDependencies } @@ -149,6 +154,9 @@ impl CrdtState { if imported_ops != admission.new_operations { return Err(CrdtError::PreviewInvalidOperationRange); } + // A delta that writes outside the root-map-of-row-maps shape is + // refused with the write-set's own typed error. + let write_set = write_set?; let resulting_frontier = fork.state_frontiers(); let fork_state = CrdtState { @@ -156,11 +164,6 @@ impl CrdtState { peer_id: self.peer_id, _single_owner: std::marker::PhantomData, }; - // Loro exposes diffs only as a complete DiffBatch. This runs after - // byte and imported-operation caps have bounded the fork/import work; - // `max_write_set_entries` is semantic cardinality enforcement, not an - // allocation short-circuit. - let write_set = fork_state.write_set_since(&before_frontier)?; if write_set.len() > limits.max_write_set_entries { return Err(CrdtError::PreviewWriteSetLimitExceeded { limit: limits.max_write_set_entries, diff --git a/nodedb-crdt/src/state/remove_fields.rs b/nodedb-crdt/src/state/remove_fields.rs new file mode 100644 index 000000000..d6831ac21 --- /dev/null +++ b/nodedb-crdt/src/state/remove_fields.rs @@ -0,0 +1,134 @@ +// SPDX-License-Identifier: Apache-2.0 + +//! Removing named scalar fields from one row, leaving every other key intact. + +use loro::{LoroMap, ValueOrContainer}; + +use crate::error::{CrdtError, Result}; + +use super::core::CrdtState; + +impl CrdtState { + /// Delete the named scalar `fields` from a row. Returns how many were + /// present and deleted. + /// + /// The inverse of `set_fields`: untouched keys keep their values. An + /// absent row or an absent field is a no-op and authors no operation. + /// + /// A container-valued key is refused with `ScalarFieldShadowsContainer`, + /// as `set_fields` refuses it: deleting it discards its nested CRDT state + /// (e.g. a row's block list). Every field is checked before any is + /// deleted, so a refused call deletes nothing. + pub fn remove_fields(&self, collection: &str, row_id: &str, fields: &[&str]) -> Result { + let coll = self.doc.get_map(collection); + let row: LoroMap = match coll.get(row_id) { + Some(ValueOrContainer::Container(loro::Container::Map(m))) => m, + _ => return Ok(0), + }; + if let Some(field) = fields + .iter() + .find(|field| matches!(row.get(field), Some(ValueOrContainer::Container(_)))) + { + return Err(CrdtError::ScalarFieldShadowsContainer { + collection: collection.to_string(), + row_id: row_id.to_string(), + field: (*field).to_string(), + }); + } + let mut removed = 0; + for field in fields { + if row.get(field).is_some() { + row.delete(field) + .map_err(|e| CrdtError::Loro(e.to_string()))?; + removed += 1; + } + } + Ok(removed) + } +} + +#[cfg(test)] +mod tests { + use loro::LoroValue; + + use super::*; + + #[test] + fn remove_fields_keeps_untouched_fields() { + let state = CrdtState::new(1).expect("state"); + state + .upsert( + "c", + "r", + &[("a", LoroValue::I64(1)), ("b", LoroValue::I64(2))], + ) + .expect("upsert"); + assert_eq!(state.remove_fields("c", "r", &["a"]).expect("remove"), 1); + assert_eq!(state.read_field("c", "r", "a"), None); + assert_eq!(state.read_field("c", "r", "b"), Some(LoroValue::I64(2))); + } + + #[test] + fn remove_fields_on_absent_row_authors_nothing() { + let state = CrdtState::new(1).expect("state"); + let before = state.local_op_counter(); + assert_eq!( + state.remove_fields("c", "missing", &["a"]).expect("remove"), + 0 + ); + assert_eq!(state.local_op_counter(), before); + assert!(!state.row_exists("c", "missing")); + } + + #[test] + fn remove_fields_refuses_a_container_key_and_deletes_nothing() { + let state = CrdtState::new(1).expect("state"); + state + .upsert("c", "r", &[("a", LoroValue::I64(1))]) + .expect("upsert"); + let row = match state.doc.get_map("c").get("r") { + Some(ValueOrContainer::Container(loro::Container::Map(m))) => m, + other => panic!("expected a row map, got {other:?}"), + }; + row.insert_container("blocks", loro::LoroList::new()) + .expect("nested list"); + match state.remove_fields("c", "r", &["a", "blocks"]) { + Err(CrdtError::ScalarFieldShadowsContainer { field, .. }) => { + assert_eq!(field, "blocks"); + } + other => panic!("expected ScalarFieldShadowsContainer, got {other:?}"), + } + assert_eq!(state.read_field("c", "r", "a"), Some(LoroValue::I64(1))); + } + + #[test] + fn remove_fields_counts_only_present_fields() { + let state = CrdtState::new(1).expect("state"); + state + .upsert( + "c", + "r", + &[("a", LoroValue::I64(1)), ("b", LoroValue::Null)], + ) + .expect("upsert"); + assert_eq!( + state + .remove_fields("c", "r", &["a", "b", "missing", "a"]) + .expect("remove"), + 2 + ); + assert_eq!(state.read_field("c", "r", "a"), None); + assert_eq!(state.read_field("c", "r", "b"), None); + } + + #[test] + fn remove_fields_with_only_absent_fields_authors_nothing() { + let state = CrdtState::new(1).expect("state"); + state + .upsert("c", "r", &[("a", LoroValue::I64(1))]) + .expect("upsert"); + let before = state.local_op_counter(); + assert_eq!(state.remove_fields("c", "r", &["zz"]).expect("remove"), 0); + assert_eq!(state.local_op_counter(), before); + } +} diff --git a/nodedb-crdt/src/state/row_image.rs b/nodedb-crdt/src/state/row_image.rs new file mode 100644 index 000000000..96b357cd6 --- /dev/null +++ b/nodedb-crdt/src/state/row_image.rs @@ -0,0 +1,189 @@ +// SPDX-License-Identifier: Apache-2.0 + +//! The state of one row that a scalar write can change, captured before the +//! write and put back when the write is abandoned. + +use loro::{LoroMap, LoroValue, ValueOrContainer}; + +use crate::error::{CrdtError, Result}; + +use super::core::CrdtState; + +/// One row as `upsert` and `set_fields` see it. +/// +/// Both writes change only scalar fields. They refuse a field held by a +/// nested container. So the scalar fields are the whole pre-image that such a +/// write has to put back. +#[derive(Debug, Clone, PartialEq)] +pub enum RowImage { + /// The collection held no entry under the row id. + Absent, + /// The row id held a plain value. + Value(LoroValue), + /// The row id held a map with these scalar fields. Container-valued + /// fields are left out: a scalar write does not change them. + Fields(Vec<(String, LoroValue)>), +} + +fn loro_error(error: loro::LoroError) -> CrdtError { + CrdtError::Loro(error.to_string()) +} + +/// The scalar fields of `row`, in key order. +fn scalar_fields(row: &LoroMap) -> Vec<(String, LoroValue)> { + row.keys() + .filter_map(|key| match row.get(&key) { + Some(ValueOrContainer::Value(value)) => Some((key.to_string(), value)), + _ => None, + }) + .collect() +} + +impl CrdtState { + /// Capture the state of `row_id` that a scalar write can change. + /// + /// A row held by a non-map container is refused with `NonMapRowValue`. + /// A scalar write replaces that container with a map, and no row image + /// can put the container back. + pub fn row_image(&self, collection: &str, row_id: &str) -> Result { + let coll = self.doc.get_map(collection); + match coll.get(row_id) { + None => Ok(RowImage::Absent), + Some(ValueOrContainer::Value(value)) => Ok(RowImage::Value(value)), + Some(ValueOrContainer::Container(loro::Container::Map(row))) => { + Ok(RowImage::Fields(scalar_fields(&row))) + } + Some(ValueOrContainer::Container(other)) => Err(CrdtError::NonMapRowValue { + collection: collection.to_string(), + row_id: row_id.to_string(), + value: format!("a {:?} container", other.get_type()), + }), + } + } + + /// Put `row_id` back to `image` with new operations. + /// + /// Container-valued fields keep their state. A `Fields` image needs the + /// row to still be a map: a scalar write keeps the row's map, so any + /// other row shape is refused with `NonMapRowValue`. + pub fn restore_row_image( + &self, + collection: &str, + row_id: &str, + image: &RowImage, + ) -> Result<()> { + let coll = self.doc.get_map(collection); + match image { + RowImage::Absent => { + if coll.get(row_id).is_some() { + coll.delete(row_id).map_err(loro_error)?; + } + Ok(()) + } + RowImage::Value(value) => coll.insert(row_id, value.clone()).map_err(loro_error), + RowImage::Fields(fields) => { + let row = match coll.get(row_id) { + Some(ValueOrContainer::Container(loro::Container::Map(row))) => row, + other => { + return Err(CrdtError::NonMapRowValue { + collection: collection.to_string(), + row_id: row_id.to_string(), + value: match other { + None => "nothing".to_string(), + Some(ValueOrContainer::Value(value)) => { + format!("the scalar {value:?}") + } + Some(ValueOrContainer::Container(container)) => { + format!("a {:?} container", container.get_type()) + } + }, + }); + } + }; + for (key, _) in scalar_fields(&row) { + if !fields.iter().any(|(field, _)| *field == key) { + row.delete(&key).map_err(loro_error)?; + } + } + for (field, value) in fields { + match row.get(field) { + Some(ValueOrContainer::Value(current)) if current == *value => {} + _ => { + row.insert(field, value.clone()).map_err(loro_error)?; + } + } + } + Ok(()) + } + } + } +} + +#[cfg(test)] +mod tests { + use super::*; + + fn string(value: &str) -> LoroValue { + LoroValue::String(value.to_string().into()) + } + + #[test] + fn an_absent_row_is_removed_again() { + let state = CrdtState::new(1).expect("state"); + let image = state.row_image("c", "r").expect("image"); + assert_eq!(image, RowImage::Absent); + state + .upsert("c", "r", &[("a", string("x"))]) + .expect("upsert"); + state.restore_row_image("c", "r", &image).expect("restore"); + assert!(state.read_row("c", "r").is_none()); + } + + #[test] + fn a_replaced_row_gets_its_scalar_fields_back() { + let state = CrdtState::new(1).expect("state"); + state + .upsert("c", "r", &[("a", string("x")), ("b", string("y"))]) + .expect("seed"); + let before = state.read_row("c", "r"); + let image = state.row_image("c", "r").expect("image"); + state + .upsert("c", "r", &[("a", string("z")), ("n", string("new"))]) + .expect("upsert"); + state.restore_row_image("c", "r", &image).expect("restore"); + assert_eq!(state.read_row("c", "r"), before); + } + + #[test] + fn a_partial_write_is_put_back() { + let state = CrdtState::new(1).expect("state"); + state.upsert("c", "r", &[("a", string("x"))]).expect("seed"); + let before = state.read_row("c", "r"); + let image = state.row_image("c", "r").expect("image"); + state + .set_fields("c", "r", &[("a", string("y")), ("b", string("z"))]) + .expect("set"); + state.restore_row_image("c", "r", &image).expect("restore"); + assert_eq!(state.read_row("c", "r"), before); + } + + #[test] + fn a_container_field_keeps_its_state() { + let state = CrdtState::new(1).expect("state"); + state.upsert("c", "r", &[("a", string("x"))]).expect("seed"); + state + .list_insert_fields("c", "r", "blocks", 0, &[("t".to_string(), string("b0"))]) + .expect("block"); + let image = state.row_image("c", "r").expect("image"); + assert_eq!( + image, + RowImage::Fields(vec![("a".to_string(), string("x"))]) + ); + state + .upsert("c", "r", &[("a", string("y"))]) + .expect("upsert"); + state.restore_row_image("c", "r", &image).expect("restore"); + assert_eq!(state.read_field("c", "r", "a"), Some(string("x"))); + assert_eq!(state.list_length("c", "r", "blocks").expect("length"), 1); + } +} diff --git a/nodedb-crdt/src/state/snapshot.rs b/nodedb-crdt/src/state/snapshot.rs index d6dd7ffe1..ff07ac87a 100644 --- a/nodedb-crdt/src/state/snapshot.rs +++ b/nodedb-crdt/src/state/snapshot.rs @@ -58,12 +58,15 @@ impl CrdtState { /// The limits are checked before Loro can import the blob, so rejected /// bytes never mutate this document or allocate an import graph. /// - /// An update whose causal predecessors are absent from this document is - /// buffered by Loro as *pending* and leaves the applied state untouched. - /// That is reported as [`CrdtError::ImportPendingDependencies`], never as - /// success: a caller that took `Ok` here would acknowledge a write that - /// was never applied. The buffered operations remain queued inside Loro, - /// so a later import carrying the missing predecessors still converges. + /// A change whose causal predecessors are absent from this document is + /// buffered by Loro as *pending* and does not reach the applied state. + /// A blob that carries any such change returns + /// [`CrdtError::ImportPendingDependencies`], never success: a caller that + /// took `Ok` here would acknowledge a write that was never applied. The + /// ready changes in the same blob still apply, so the state can advance + /// under that error. [`Self::import_tracked`] reports the rows they + /// changed. The buffered changes stay queued inside Loro, so a later + /// import carrying the missing predecessors still converges. /// /// The returned [`ImportAdmission`] reports how much of the blob was new /// and how much Loro trimmed as already-known. An `Ok` whose @@ -104,6 +107,9 @@ impl CrdtState { /// Call this periodically (e.g., every 30 minutes or when memory /// pressure exceeds threshold) to prevent unbounded history growth. pub fn compact_history(&mut self) -> Result<()> { + // `oplog_frontiers` excludes an open auto-commit transaction. Commit + // it first so the shallow root covers every write made so far. + self.doc.commit(); self.compact_to_frontiers(&self.doc.oplog_frontiers()) } diff --git a/nodedb-crdt/src/state/write_set.rs b/nodedb-crdt/src/state/write_set.rs index 590925c17..cfe1fd0f5 100644 --- a/nodedb-crdt/src/state/write_set.rs +++ b/nodedb-crdt/src/state/write_set.rs @@ -1,94 +1,419 @@ // SPDX-License-Identifier: Apache-2.0 -//! Post-import write-set extraction. +//! Write-set extraction from an import. //! -//! Given a Loro version frontier captured immediately before an import, these -//! helpers compute the rows a delta *actually* wrote (independent of any -//! row-id the sender claimed) and assemble a [`ProposedChange`] for a single -//! committed row so it can be re-checked against installed constraints. +//! An import reports the rows it *actually* wrote, independent of any row id +//! the sender claimed, and a committed row is assembled into a +//! [`ProposedChange`] so it can be re-checked against installed constraints. //! -//! The write-set is row-granular only — collection + row-id pairs. The +//! The rows come from the operations the import added to the oplog. Loro adds +//! an operation to the oplog only once its causal predecessors are present, +//! and the oplog version vector also counts an open local transaction. So the +//! version vector difference across the import holds exactly: +//! - the ready operations of this blob. +//! - earlier pending operations this blob unblocked. +//! +//! It never holds an operation the document already had, nor a local write. +//! `ImportStatus::success` is not used: it also covers ranges the document +//! already knew. +//! +//! Each new span is read back with `LoroDoc::export_json_in_id_span`: +//! - an operation on a root map names the row key it set or deleted. +//! - an operation on any other container names the row its container path +//! passes through (`LoroDoc::get_path_to_container`). +//! +//! A container whose row the same import deleted has no path. The root-map +//! delete already names that row. Nothing subscribes to the document, so Loro +//! never turns on diff recording, and later local writes cost the same as on a +//! document that never ran a tracked import. +//! +//! An import into an empty applied state with no open local transaction makes +//! every present row new, so the rows are read from the state directly. That +//! also covers a shallow snapshot, whose history before its shallow root cannot +//! be read back. +//! +//! A row is reported when an applied operation wrote it, also when that write +//! lost a concurrent conflict and left the visible value unchanged. +//! +//! The write-set is row-granular only: collection + row-id pairs. The //! validator re-reads the full row, so field-level detail is unnecessary here. //! Ordering is deterministic (`BTreeSet` for the write-set, sorted field vec //! for the change) because every replica must agree on the same result. - -use std::collections::BTreeSet; - -use loro::event::Diff; -use loro::{ContainerID, ContainerType, Frontiers, LoroValue}; +//! +//! Every collection is a root map of row maps. A write outside that shape +//! names no row the validator can check, so the write-set refuses the import: +//! - an operation on a root text, list, movable list, tree, or counter, or on +//! a container under one, is [`CrdtError::NonMapRootContainer`]. +//! - a root-map row set to a scalar or a non-map container, or an operation +//! on such a row container, is [`CrdtError::NonMapRowValue`]. + +use std::collections::{BTreeSet, HashMap}; + +use loro::{ + Container, ContainerID, ContainerType, Frontiers, IdSpan, Index, JsonMapOp, JsonOp, + JsonOpContent, LoroDoc, LoroValue, ValueOrContainer, +}; use nodedb_types::Surrogate; use crate::error::{CrdtError, Result}; use crate::validator::ProposedChange; use super::core::CrdtState; +use super::import_admission::ImportAdmission; + +/// The outcome of an import and the rows it wrote. +/// +/// `write_set` is filled for every outcome. A blob can apply its ready +/// changes and still return `ImportPendingDependencies`. +#[derive(Debug)] +#[must_use] +pub struct WriteSetImport { + /// The import result, as [`CrdtState::import`] returns it. + pub outcome: Result, + /// Sorted `(collection, row_id)` pairs an applied operation of the import + /// wrote: rows added, removed, replaced, or changed in any nested + /// container. An error when an applied operation breaks the collection + /// shape: the caller must refuse the import. + pub write_set: Result>, +} -impl CrdtState { - /// Capture the current version frontier. Take this *before* an import so a - /// later [`write_set_since`](Self::write_set_since) can diff against it. - pub fn frontier(&self) -> Frontiers { - self.doc.state_frontiers() +/// A write outside the root-map-of-row-maps shape. +/// +/// Ordered so every replica reports the same fault for one import. +#[derive(Debug, Clone, PartialEq, Eq, PartialOrd, Ord)] +pub(in crate::state) enum ShapeFault { + /// A root container that is not a map, or a container under one. + NonMapRoot { + container: String, + container_type: String, + }, + /// A root-map row whose value is not a map. + NonMapRow { + collection: String, + row_id: String, + value: String, + }, +} + +impl ShapeFault { + fn into_error(self) -> CrdtError { + match self { + Self::NonMapRoot { + container, + container_type, + } => CrdtError::NonMapRootContainer { + container, + container_type, + }, + Self::NonMapRow { + collection, + row_id, + value, + } => CrdtError::NonMapRowValue { + collection, + row_id, + value, + }, + } } +} - /// Compute the `(collection, row_id)` pairs written between `before` and the - /// current state — the rows a just-imported delta actually touched. - /// - /// Two container shapes surface a written row: - /// - the collection's root map gains/updates a key (the row-id), and - /// - the row's own (normal) map container. - /// - /// For the row container, `get_path_to_container` returns the full path - /// from root to target, *including* a trailing element for the target - /// itself: `[(collection_root, Key(collection)), (row_container, - /// Key(row_id))]`. The collection name therefore comes from the - /// second-to-last element's root ContainerID, and the row-id from the - /// last element's key index. Reading both from the same (first) element - /// would resolve the collection root as if it were a row named after the - /// collection. - /// - /// Both shapes map to the same `(collection, row_id)`. Loro's - /// `ContainerID` / `Index` / `Diff` are foreign enums; the fallthrough arm - /// intentionally ignores irrelevant shapes (sequence/tree indices, non-map - /// roots, list/text/tree/counter diffs). - pub fn write_set_since(&self, before: &Frontiers) -> Result> { - let after = self.doc.state_frontiers(); - let batch = self - .doc - .diff(before, &after) - .map_err(|e| CrdtError::Loro(e.to_string()))?; +/// The rows an import wrote and the shape faults among its writes. +#[derive(Debug, Default)] +pub(in crate::state) struct ImportedRows { + pub(in crate::state) rows: BTreeSet<(String, String)>, + pub(in crate::state) faults: BTreeSet, +} - let mut keys = BTreeSet::<(String, String)>::new(); - for (cid, diff) in batch.iter() { - match cid { - ContainerID::Root { - name, - container_type: ContainerType::Map, - } => { - if let Diff::Map(md) = diff { - for k in md.updated.keys() { - keys.insert((name.to_string(), k.to_string())); - } +/// Run `import` on `doc` and collect the `(collection, row_id)` pairs its +/// applied operations wrote. +/// +/// `scope` limits the pairs to one collection. `None` covers every root map. +/// Shape faults are collected for every collection. +pub(in crate::state) fn collect_imported_rows( + doc: &LoroDoc, + scope: Option<&str>, + import: impl FnOnce() -> T, +) -> (T, ImportedRows) { + if doc.state_frontiers().is_empty() && doc.get_pending_txn_len() == 0 { + let outcome = import(); + return (outcome, present_rows(doc, scope)); + } + + let before = doc.oplog_vv(); + let outcome = import(); + let after = doc.oplog_vv(); + let mut reader = RowReader { + doc, + scope, + imported: ImportedRows::default(), + container_places: HashMap::new(), + }; + for (peer, end) in after.iter() { + let start = before.get(peer).copied().unwrap_or(0); + if *end <= start { + continue; + } + for change in doc.export_json_in_id_span(IdSpan::new(*peer, start, *end)) { + for op in &change.ops { + reader.read(op); + } + } + } + (outcome, reader.imported) +} + +/// Run `import` on `doc` and collect the sorted `(collection, row_id)` pairs +/// it wrote across every collection, or the first shape fault. +pub(in crate::state) fn collect_write_set( + doc: &LoroDoc, + import: impl FnOnce() -> T, +) -> (T, Result>) { + let (outcome, imported) = collect_imported_rows(doc, None, import); + let write_set = match imported.faults.into_iter().next() { + Some(fault) => Err(fault.into_error()), + None => Ok(imported.rows.into_iter().collect()), + }; + (outcome, write_set) +} + +/// Every row present in `scope`, or in every root map when `scope` is `None`. +fn present_rows(doc: &LoroDoc, scope: Option<&str>) -> ImportedRows { + let mut imported = ImportedRows::default(); + let mut add_collection = |collection: &str| { + let map = doc.get_map(collection); + for row_id in map.keys() { + let row_id = row_id.to_string(); + if let Some(value) = map.get(&row_id).and_then(|row| non_map_row(&row)) { + imported.faults.insert(ShapeFault::NonMapRow { + collection: collection.to_owned(), + row_id: row_id.clone(), + value, + }); + } + imported.rows.insert((collection.to_owned(), row_id)); + } + }; + let mut root_faults = BTreeSet::new(); + if let LoroValue::Map(roots) = doc.get_value() { + for root in roots.values() { + let LoroValue::Container(id) = root else { + continue; + }; + let ContainerID::Root { + name, + container_type, + } = id + else { + continue; + }; + if id.is_mergeable() { + continue; + } + if *container_type != ContainerType::Map { + root_faults.insert(ShapeFault::NonMapRoot { + container: name.to_string(), + container_type: format!("{container_type:?}"), + }); + } else if scope.is_none() { + add_collection(name.as_str()); + } + } + } + if let Some(collection) = scope { + add_collection(collection); + } + imported.faults.extend(root_faults); + imported +} + +/// A description of a row value that is not a map, or `None` for a map. +fn non_map_row(row: &ValueOrContainer) -> Option { + match row { + ValueOrContainer::Container(Container::Map(_)) => None, + ValueOrContainer::Container(other) => Some(format!("a {:?} container", other.get_type())), + ValueOrContainer::Value(value) => Some(format!("the scalar {value:?}")), + } +} + +/// A description of a map-insert value that is not a map container, or +/// `None` for a map container. +fn non_map_inserted(value: &LoroValue) -> Option { + match value { + LoroValue::Container(id) if id.container_type() == ContainerType::Map => None, + LoroValue::Container(id) => Some(format!("a {:?} container", id.container_type())), + other => Some(format!("the scalar {other:?}")), + } +} + +/// Where a non-root container sits. +#[derive(Debug, Clone)] +enum ContainerPlace { + /// Inside row `row_id` of the root map `collection`. + Row(String, String), + /// No live path: the import deleted the row that held it. + Detached, + /// Outside a root-map row. + Misplaced(ShapeFault), +} + +/// Maps imported operations to the rows they wrote. +struct RowReader<'a> { + doc: &'a LoroDoc, + scope: Option<&'a str>, + imported: ImportedRows, + /// Place of each non-root container already resolved. + container_places: HashMap, +} + +impl RowReader<'_> { + fn read(&mut self, op: &JsonOp) { + let scope = self.scope; + let in_scope = |collection: &str| scope.is_none_or(|wanted| wanted == collection); + + if let ContainerID::Root { + name, + container_type, + } = &op.container + && !op.container.is_mergeable() + { + if *container_type != ContainerType::Map { + self.imported.faults.insert(ShapeFault::NonMapRoot { + container: name.to_string(), + container_type: format!("{container_type:?}"), + }); + return; + } + let row = match &op.content { + JsonOpContent::Map(JsonMapOp::Insert { key, value }) => { + if let Some(value) = non_map_inserted(value) { + self.imported.faults.insert(ShapeFault::NonMapRow { + collection: name.to_string(), + row_id: key.clone(), + value, + }); } + Some(key) } - ContainerID::Normal { .. } => { - // Path is root → target, last element being the target - // container itself: the owning collection root is the - // second-to-last element, the row-id the last element's key. - if let Some(path) = self.doc.get_path_to_container(cid) - && path.len() >= 2 - && let (Some((parent, _)), Some((_, idx))) = - (path.get(path.len() - 2), path.last()) - && let (Some((name, _)), Some(key)) = (parent.as_root(), idx.as_key()) - { - keys.insert((name.to_string(), key.to_string())); - } + JsonOpContent::Map(JsonMapOp::Delete { key }) => Some(key), + _ => None, + }; + if let Some(row_id) = row + && in_scope(name.as_str()) + { + self.imported + .rows + .insert((name.to_string(), row_id.clone())); + } + return; + } + + let doc = self.doc; + let place = self + .container_places + .entry(op.container.clone()) + .or_insert_with(|| place_of_container(doc, &op.container)); + match place { + ContainerPlace::Row(collection, row_id) => { + if in_scope(collection.as_str()) { + self.imported + .rows + .insert((collection.clone(), row_id.clone())); + } + } + ContainerPlace::Detached => {} + ContainerPlace::Misplaced(fault) => { + // An edit inside a non-map row still wrote that row. The + // fault refuses the write-set, and a tracked import reports + // the row. + if let ShapeFault::NonMapRow { + collection, row_id, .. + } = &*fault + && in_scope(collection.as_str()) + { + self.imported + .rows + .insert((collection.clone(), row_id.clone())); } - _ => {} + self.imported.faults.insert(fault.clone()); } } + } +} - // `BTreeSet` ⇒ deterministic sorted output. - Ok(keys.into_iter().collect()) +/// Where a non-root container sits. +/// +/// The path runs from the root: `[(root, Key(collection)), (row, +/// Key(row_id)), ...]`. The row container must be a map. A container under +/// any other root kind is misplaced. +fn place_of_container(doc: &LoroDoc, container: &ContainerID) -> ContainerPlace { + let Some(path) = doc.get_path_to_container(container) else { + return ContainerPlace::Detached; + }; + match path.as_slice() { + [ + ( + ContainerID::Root { + name, + container_type: ContainerType::Map, + }, + _, + ), + (row, Index::Key(row_id)), + .., + ] => { + if row.container_type() != ContainerType::Map { + return ContainerPlace::Misplaced(ShapeFault::NonMapRow { + collection: name.to_string(), + row_id: row_id.to_string(), + value: format!("a {:?} container", row.container_type()), + }); + } + ContainerPlace::Row(name.to_string(), row_id.to_string()) + } + [ + ( + root @ ContainerID::Root { + name, + container_type, + }, + _, + ), + .., + ] if *container_type != ContainerType::Map && !root.is_mergeable() => { + ContainerPlace::Misplaced(ShapeFault::NonMapRoot { + container: name.to_string(), + container_type: format!("{container_type:?}"), + }) + } + // A path that does not reach a root-map row names no row. + _ => ContainerPlace::Detached, + } +} + +impl CrdtState { + /// Current applied-state frontier. Equal frontiers mean equal applied + /// state for two documents built from the same history. + pub fn frontier(&self) -> Frontiers { + self.doc.state_frontiers() + } + + /// [`Self::import`] that also reports the `(collection, row_id)` pairs + /// the import wrote. + pub fn import_with_write_set(&self, data: &[u8]) -> WriteSetImport { + let (outcome, write_set) = collect_write_set(&self.doc, || self.import(data)); + WriteSetImport { outcome, write_set } + } + + /// [`Self::import_local`] that also reports the `(collection, row_id)` + /// pairs the import wrote. + /// + /// For bytes this process or its backups exported: a checkpoint, or a + /// collection snapshot. Sync deltas from a peer go through + /// [`Self::import_with_write_set`], which keeps the peer ceilings. + pub fn import_local_with_write_set(&self, data: &[u8]) -> WriteSetImport { + let (outcome, write_set) = collect_write_set(&self.doc, || self.import_local(data)); + WriteSetImport { outcome, write_set } } /// Assemble a [`ProposedChange`] from a committed row's current fields. @@ -131,17 +456,27 @@ mod tests { LoroValue::I64(v) } - // (1) A single-row insert surfaces exactly that row. + fn pair(collection: &str, row: &str) -> (String, String) { + (collection.to_string(), row.to_string()) + } + + fn write_set_of(dst: &CrdtState, data: &[u8]) -> Vec<(String, String)> { + let imported = dst.import_with_write_set(data); + imported.outcome.as_ref().expect("import"); + imported.write_set.expect("write set") + } + + // (1) A single-row delta surfaces exactly that row. #[test] fn single_insert_write_set() { - let state = CrdtState::new(1).unwrap(); - let before = state.frontier(); - state - .upsert("users", "u1", &[("name", s("Alice"))]) + let src = CrdtState::new(2).unwrap(); + src.upsert("users", "u1", &[("name", s("Alice"))]).unwrap(); + let delta = src + .export_updates_since(&loro::VersionVector::default()) .unwrap(); - let ws = state.write_set_since(&before).unwrap(); - assert_eq!(ws, vec![("users".to_string(), "u1".to_string())]); + let dst = CrdtState::new(1).unwrap(); + assert_eq!(write_set_of(&dst, &delta), vec![pair("users", "u1")]); } // (2) A two-row blob imported from a peer surfaces both rows, sorted. @@ -153,16 +488,9 @@ mod tests { let snapshot = src.export_snapshot().unwrap(); let dst = CrdtState::new(1).unwrap(); - let before = dst.frontier(); - dst.import(&snapshot).unwrap(); - - let ws = dst.write_set_since(&before).unwrap(); assert_eq!( - ws, - vec![ - ("users".to_string(), "a".to_string()), - ("users".to_string(), "b".to_string()), - ] + write_set_of(&dst, &snapshot), + vec![pair("users", "a"), pair("users", "b")] ); } @@ -174,12 +502,9 @@ mod tests { let snapshot = src.export_snapshot().unwrap(); let dst = CrdtState::new(1).unwrap(); - let before = dst.frontier(); - dst.import(&snapshot).unwrap(); - - let ws = dst.write_set_since(&before).unwrap(); - assert_eq!(ws, vec![("orders".to_string(), "real-1".to_string())]); - assert!(!ws.contains(&("orders".to_string(), "fake-claimed".to_string()))); + let ws = write_set_of(&dst, &snapshot); + assert_eq!(ws, vec![pair("orders", "real-1")]); + assert!(!ws.contains(&pair("orders", "fake-claimed"))); } // (4) A cross-collection blob surfaces the real collection it wrote into. @@ -190,18 +515,14 @@ mod tests { let snapshot = src.export_snapshot().unwrap(); let dst = CrdtState::new(1).unwrap(); - let before = dst.frontier(); - dst.import(&snapshot).unwrap(); - - let ws = dst.write_set_since(&before).unwrap(); - assert_eq!(ws, vec![("secret".to_string(), "s1".to_string())]); + assert_eq!(write_set_of(&dst, &snapshot), vec![pair("secret", "s1")]); } // (5) A peer delta that merges ONE field of an existing row surfaces that // one row, and the post-import row still carries ALL fields (so NOT // NULL can see the untouched required fields). Modeled via a // field-level merge (`insert` on the existing row container) exported - // as an incremental delta — NOT an `upsert`, which is whole-row + // as an incremental delta, not an `upsert`, which is whole-row // replace and would wipe the untouched field. #[test] fn update_only_write_set_and_full_change() { @@ -212,7 +533,6 @@ mod tests { let dst = CrdtState::new(1).unwrap(); dst.import(&snapshot).unwrap(); - let before = dst.frontier(); // Field-level merge of only `name` on the EXISTING row container. let src_vv = src.doc.oplog_vv(); @@ -223,10 +543,7 @@ mod tests { _ => panic!("row container missing"), } let delta = src.export_updates_since(&src_vv).unwrap(); - dst.import(&delta).unwrap(); - - let ws = dst.write_set_since(&before).unwrap(); - assert_eq!(ws, vec![("users".to_string(), "u1".to_string())]); + assert_eq!(write_set_of(&dst, &delta), vec![pair("users", "u1")]); let change = dst .build_change_from_row("users", "u1", Surrogate::ZERO) @@ -249,14 +566,16 @@ mod tests { let snapshot = src.export_snapshot().unwrap(); let dst = CrdtState::new(1).unwrap(); - let before = dst.frontier(); - dst.import(&snapshot).unwrap(); - dst.write_set_since(&before).unwrap() + write_set_of(&dst, &snapshot) } let first = run(); let second = run(); assert_eq!(first, second); + assert_eq!( + first, + vec![pair("secret", "s1"), pair("users", "a"), pair("users", "b")] + ); } // (7) A deleted row yields no change to validate. @@ -274,4 +593,317 @@ mod tests { .is_none() ); } + + // (8) A delta that only changes a list nested under a row names that row. + #[test] + fn nested_only_change_is_in_the_write_set() { + let src = CrdtState::new(2).unwrap(); + src.upsert("pages", "p1", &[("title", s("t"))]).unwrap(); + src.list_insert_fields("pages", "p1", "blocks", 0, &[("text".to_owned(), s("one"))]) + .unwrap(); + let dst = CrdtState::new(1).unwrap(); + dst.import(&src.export_snapshot().unwrap()).unwrap(); + + let before = src.oplog_version_vector(); + match src.doc.get_map("pages").get("p1") { + Some(loro::ValueOrContainer::Container(loro::Container::Map(row))) => { + match row.get("blocks") { + Some(loro::ValueOrContainer::Container(loro::Container::MovableList(list))) => { + match list.get(0) { + Some(loro::ValueOrContainer::Container(loro::Container::Map( + block, + ))) => { + block.insert("text", s("two")).unwrap(); + } + _ => panic!("block map missing"), + } + } + _ => panic!("block list missing"), + } + } + _ => panic!("row container missing"), + } + let delta = src.export_updates_since(&before).unwrap(); + + assert_eq!(write_set_of(&dst, &delta), vec![pair("pages", "p1")]); + } + + // (9) A shallow source works for the first import and for later deltas. + #[test] + fn shallow_source_reports_its_rows() { + let mut src = CrdtState::new(2).unwrap(); + src.upsert("users", "a", &[("x", n(1))]).unwrap(); + src.upsert("users", "b", &[("x", n(2))]).unwrap(); + src.compact_history().unwrap(); + let snapshot = src.export_snapshot().unwrap(); + + let dst = CrdtState::new(1).unwrap(); + let first = dst.import_with_write_set(&snapshot); + first.outcome.as_ref().expect("shallow import"); + assert!(dst.doc.is_shallow()); + assert_eq!( + first.write_set.expect("write set"), + vec![pair("users", "a"), pair("users", "b")] + ); + + let before = src.oplog_version_vector(); + src.upsert("users", "c", &[("x", n(3))]).unwrap(); + src.delete("users", "a").unwrap(); + let delta = src.export_updates_since(&before).unwrap(); + assert_eq!( + write_set_of(&dst, &delta), + vec![pair("users", "a"), pair("users", "c")] + ); + } + + // (10) Rows of each collection are attributed to that collection. + #[test] + fn rows_are_attributed_to_their_own_collection() { + let src = CrdtState::new(2).unwrap(); + src.upsert("users", "u1", &[("x", n(1))]).unwrap(); + src.upsert("orders", "o1", &[("amt", n(1))]).unwrap(); + let dst = CrdtState::new(1).unwrap(); + dst.import(&src.export_snapshot().unwrap()).unwrap(); + + let before = src.oplog_version_vector(); + src.upsert("users", "u2", &[("x", n(2))]).unwrap(); + src.upsert("orders", "u1", &[("amt", n(5))]).unwrap(); + src.list_insert_fields("orders", "o1", "lines", 0, &[("sku".to_owned(), s("k"))]) + .unwrap(); + let delta = src.export_updates_since(&before).unwrap(); + + assert_eq!( + write_set_of(&dst, &delta), + vec![ + pair("orders", "o1"), + pair("orders", "u1"), + pair("users", "u2") + ] + ); + } + + // (11) A local write is not part of an import's write set. + #[test] + fn local_write_is_not_in_the_write_set() { + let src = CrdtState::new(2).unwrap(); + src.upsert("users", "u1", &[("x", n(1))]).unwrap(); + let dst = CrdtState::new(1).unwrap(); + dst.upsert("users", "local", &[("x", n(9))]).unwrap(); + + assert_eq!( + write_set_of(&dst, &src.export_snapshot().unwrap()), + vec![pair("users", "u1")] + ); + } + + fn row_map(state: &CrdtState, collection: &str, row: &str) -> loro::LoroMap { + match state.doc.get_map(collection).get(row) { + Some(loro::ValueOrContainer::Container(loro::Container::Map(row))) => row, + other => panic!("row {collection}/{row} is not a map: {other:?}"), + } + } + + /// A source and a target that share one row, and the source's version. + fn synced_users() -> (CrdtState, CrdtState, loro::VersionVector) { + let src = CrdtState::new(2).unwrap(); + src.upsert("users", "u1", &[("name", s("a"))]).unwrap(); + let dst = CrdtState::new(1).unwrap(); + dst.import(&src.export_snapshot().unwrap()).unwrap(); + let before = src.oplog_version_vector(); + (src, dst, before) + } + + fn write_set_error(dst: &CrdtState, data: &[u8]) -> CrdtError { + let imported = dst.import_with_write_set(data); + imported.outcome.as_ref().expect("import"); + imported.write_set.expect_err("shape fault") + } + + // (12) An operation on a root text container is refused. + #[test] + fn root_text_delta_is_refused() { + let (src, dst, before) = synced_users(); + src.doc.get_text("notes").insert(0, "hi").unwrap(); + let delta = src.export_updates_since(&before).unwrap(); + + match write_set_error(&dst, &delta) { + CrdtError::NonMapRootContainer { + container, + container_type, + } => { + assert_eq!(container, "notes"); + assert_eq!(container_type, "Text"); + } + other => panic!("expected NonMapRootContainer, got {other:?}"), + } + } + + // (13) A root text in a first snapshot import is refused too. + #[test] + fn root_text_snapshot_is_refused() { + let src = CrdtState::new(2).unwrap(); + src.upsert("users", "u1", &[("name", s("a"))]).unwrap(); + src.doc.get_text("notes").insert(0, "hi").unwrap(); + let dst = CrdtState::new(1).unwrap(); + + match write_set_error(&dst, &src.export_snapshot().unwrap()) { + CrdtError::NonMapRootContainer { container, .. } => assert_eq!(container, "notes"), + other => panic!("expected NonMapRootContainer, got {other:?}"), + } + } + + // (14) A row set to a scalar is refused, naming the collection and row. + #[test] + fn scalar_row_is_refused() { + let (src, dst, before) = synced_users(); + src.doc.get_map("users").insert("u2", 5).unwrap(); + let delta = src.export_updates_since(&before).unwrap(); + + match write_set_error(&dst, &delta) { + CrdtError::NonMapRowValue { + collection, row_id, .. + } => { + assert_eq!(collection, "users"); + assert_eq!(row_id, "u2"); + } + other => panic!("expected NonMapRowValue, got {other:?}"), + } + } + + // (15) A scalar row in a first snapshot import is refused too. + #[test] + fn scalar_row_snapshot_is_refused() { + let src = CrdtState::new(2).unwrap(); + src.doc.get_map("users").insert("u1", 5).unwrap(); + let dst = CrdtState::new(1).unwrap(); + + let err = write_set_error(&dst, &src.export_snapshot().unwrap()); + assert!( + matches!(err, CrdtError::NonMapRowValue { ref row_id, .. } if row_id == "u1"), + "{err:?}" + ); + } + + // (16) A row set to a non-map container is refused, and so is a later + // edit inside that container. + #[test] + fn non_map_row_container_is_refused() { + let (src, dst, before) = synced_users(); + let text = src + .doc + .get_map("users") + .insert_container("u3", loro::LoroText::new()) + .unwrap(); + text.insert(0, "x").unwrap(); + let delta = src.export_updates_since(&before).unwrap(); + + let err = write_set_error(&dst, &delta); + assert!( + matches!(err, CrdtError::NonMapRowValue { ref row_id, .. } if row_id == "u3"), + "{err:?}" + ); + } + + /// A collection snapshot past the peer operation ceiling imports through + /// the local write-set import, and reports its row. The peer import + /// refuses it. + #[test] + fn a_local_write_set_import_admits_past_the_peer_ceilings() { + let doc = loro::LoroDoc::new(); + doc.set_peer_id(2).unwrap(); + let row = doc + .get_map("docs") + .insert_container("r", loro::LoroMap::new()) + .unwrap(); + let body = row.insert_container("body", loro::LoroText::new()).unwrap(); + body.insert(0, &"x".repeat(crate::state::DEFAULT_MAX_IMPORT_OPS + 1)) + .unwrap(); + doc.commit(); + let snapshot = doc.export(loro::ExportMode::Snapshot).unwrap(); + + let peer = CrdtState::new(1).unwrap(); + assert!( + matches!( + peer.import_with_write_set(&snapshot).outcome, + Err(CrdtError::ImportOperationLimitExceeded { .. }) + ), + "a peer import keeps the ceilings" + ); + + let local = CrdtState::new(1).unwrap(); + let imported = local.import_local_with_write_set(&snapshot); + imported.outcome.as_ref().expect("local import"); + assert_eq!( + imported.write_set.expect("write set"), + vec![pair("docs", "r")] + ); + } + + // (17) A nested change under a map row stays accepted. + #[test] + fn nested_field_change_is_accepted() { + let (src, dst, before) = synced_users(); + row_map(&src, "users", "u1") + .insert_container("tags", loro::LoroList::new()) + .unwrap() + .push("t") + .unwrap(); + let delta = src.export_updates_since(&before).unwrap(); + + assert_eq!(write_set_of(&dst, &delta), vec![pair("users", "u1")]); + } + + // (18) A write that loses a concurrent conflict still names its row. + #[test] + fn concurrent_conflict_loser_is_reported() { + // Peer 2 wins a same-lamport conflict against peer 1. + let loser = CrdtState::new(1).unwrap(); + loser.upsert("users", "u1", &[("name", s("base"))]).unwrap(); + let winner = CrdtState::new(2).unwrap(); + winner.import(&loser.export_snapshot().unwrap()).unwrap(); + + let loser_before = loser.oplog_version_vector(); + row_map(&loser, "users", "u1") + .insert("name", s("lose")) + .unwrap(); + row_map(&winner, "users", "u1") + .insert("name", s("win")) + .unwrap(); + winner.doc.commit(); + let losing_delta = loser.export_updates_since(&loser_before).unwrap(); + + assert_eq!( + write_set_of(&winner, &losing_delta), + vec![pair("users", "u1")] + ); + let change = winner + .build_change_from_row("users", "u1", Surrogate::ZERO) + .unwrap(); + assert!(change.fields.contains(&("name".to_string(), s("win")))); + } + + // (19) A nested edit under a row the same import deletes names that row + // once, through the root-map delete. + #[test] + fn nested_op_under_a_row_deleted_in_the_same_import() { + let src = CrdtState::new(2).unwrap(); + src.upsert("pages", "p1", &[("title", s("t"))]).unwrap(); + src.list_insert_fields("pages", "p1", "blocks", 0, &[("text".to_owned(), s("one"))]) + .unwrap(); + let dst = CrdtState::new(1).unwrap(); + dst.import(&src.export_snapshot().unwrap()).unwrap(); + + let before = src.oplog_version_vector(); + row_map(&src, "pages", "p1") + .insert("title", s("u")) + .unwrap(); + src.delete("pages", "p1").unwrap(); + let delta = src.export_updates_since(&before).unwrap(); + + assert_eq!(write_set_of(&dst, &delta), vec![pair("pages", "p1")]); + assert!( + dst.build_change_from_row("pages", "p1", Surrogate::ZERO) + .is_none() + ); + } } diff --git a/nodedb-fts/src/analyzer/language/stemmer.rs b/nodedb-fts/src/analyzer/language/stemmer.rs index ddf179934..62dbc4654 100644 --- a/nodedb-fts/src/analyzer/language/stemmer.rs +++ b/nodedb-fts/src/analyzer/language/stemmer.rs @@ -4,7 +4,7 @@ use rust_stemmers::{Algorithm, Stemmer}; -use crate::analyzer::pipeline::{TextAnalyzer, tokenize_with_stemmer}; +use crate::analyzer::pipeline::{TextAnalyzer, tokenize_raw, tokenize_with_stemmer}; use super::stop_words; @@ -97,10 +97,7 @@ impl NoStemAnalyzer { impl TextAnalyzer for NoStemAnalyzer { fn analyze(&self, text: &str) -> Vec { let stop_list = stop_words::stop_words(&self.lang_code); - // Use English stemmer as no-op: it won't affect non-English words meaningfully. - // The stop word list does the language-specific work. - let stemmer = Stemmer::create(Algorithm::English); - tokenize_with_stemmer(text, &stemmer, &self.lang_code, stop_list) + tokenize_raw(text, &self.lang_code, stop_list) } fn name(&self) -> &str { @@ -157,4 +154,10 @@ mod tests { // "यह" and "है" are Hindi stop words. assert!(!tokens.iter().any(|t| t == "यह" || t == "है")); } + + #[test] + fn no_stem_keeps_latin_words_as_written() { + let analyzer = NoStemAnalyzer::new("indonesian").unwrap(); + assert_eq!(analyzer.analyze("Running dogs"), vec!["running", "dogs"]); + } } diff --git a/nodedb-fts/src/analyzer/pipeline.rs b/nodedb-fts/src/analyzer/pipeline.rs index 82bd716d9..48b59132a 100644 --- a/nodedb-fts/src/analyzer/pipeline.rs +++ b/nodedb-fts/src/analyzer/pipeline.rs @@ -48,51 +48,10 @@ pub fn tokenize_no_stem(text: &str) -> Vec { tokenize_raw(text, "en", en_stops) } -/// Raw tokenization shared by `tokenize_no_stem` and language-specific variants. -/// Same pipeline as `tokenize_with_stemmer` but skips the stemming step. +/// Tokenize without stemming: the pipeline of [`tokenize_with_stemmer`] +/// with the stemming stage left out. pub(crate) fn tokenize_raw(text: &str, lang: &str, stop_list: &[&str]) -> Vec { - let mut normalized = String::with_capacity(text.len()); - for c in text.chars() { - if script::is_cjk(c) || script::is_hangul_jamo(c) || script::is_thai(c) { - for lc in c.to_lowercase() { - normalized.push(lc); - } - } else { - for decomposed in c.nfd() { - if unicode_normalization::char::is_combining_mark(decomposed) { - continue; - } - for lc in decomposed.to_lowercase() { - normalized.push(lc); - } - } - } - } - - let mut tokens = Vec::new(); - for word in normalized.split(|c: char| !c.is_alphanumeric() && c != '-' && c != '_') { - let trimmed = word.trim_matches(|c: char| c == '-' || c == '_'); - if trimmed.is_empty() || trimmed.len() <= 1 { - continue; - } - if trimmed.chars().any(script::needs_segmentation) { - let cjk_tokens = if matches!(lang, "ja" | "zh" | "ko" | "th") { - super::language::cjk::segmenter::segment(trimmed, lang) - } else { - bigram::tokenize_cjk(trimmed) - }; - for token in cjk_tokens { - if !token.is_empty() && !is_stop_word_in_list(&token, stop_list) { - tokens.push(token); - } - } - continue; - } - if !is_stop_word_in_list(trimmed, stop_list) { - tokens.push(trimmed.to_string()); - } - } - tokens + tokenize(text, None, lang, stop_list) } /// Shared tokenization pipeline used by both the standard `analyze()` function @@ -104,6 +63,12 @@ pub(crate) fn tokenize_with_stemmer( lang: &str, stop_list: &[&str], ) -> Vec { + tokenize(text, Some(stemmer), lang, stop_list) +} + +/// The one analysis pipeline. A `None` stemmer keeps each non-CJK token as +/// the normalized word. +fn tokenize(text: &str, stemmer: Option<&Stemmer>, lang: &str, stop_list: &[&str]) -> Vec { // Stage 1-2: Normalize and lowercase. // Process char-by-char: for CJK/Hangul characters, preserve as-is (lowercased). // For others, apply NFD + strip combining marks to handle diacritics (café → cafe). @@ -163,11 +128,7 @@ pub(crate) fn tokenize_with_stemmer( } } } else if run.len() > 1 && !is_stop_word_in_list(&run, stop_list) { - // Non-CJK run → stem. - let stemmed = stemmer.stem(&run); - if !stemmed.is_empty() { - tokens.push(stemmed.into_owned()); - } + push_stemmed(&mut tokens, stemmer, &run); } i = run_end; } @@ -185,15 +146,24 @@ pub(crate) fn tokenize_with_stemmer( } // Stage 7: Snowball stemming. - let stemmed = stemmer.stem(trimmed); - if !stemmed.is_empty() { - tokens.push(stemmed.into_owned()); - } + push_stemmed(&mut tokens, stemmer, trimmed); } tokens } +/// Append `word` stemmed by `stemmer`, or unchanged when there is none. An +/// empty stem is dropped. +fn push_stemmed(tokens: &mut Vec, stemmer: Option<&Stemmer>, word: &str) { + let token = match stemmer { + Some(stemmer) => stemmer.stem(word).into_owned(), + None => word.to_owned(), + }; + if !token.is_empty() { + tokens.push(token); + } +} + /// Check if a word is in a sorted stop word list via binary search. fn is_stop_word_in_list(word: &str, list: &[&str]) -> bool { list.binary_search(&word).is_ok() diff --git a/nodedb-fts/src/backend/memory.rs b/nodedb-fts/src/backend/memory.rs index a9dac7723..e9bb8924a 100644 --- a/nodedb-fts/src/backend/memory.rs +++ b/nodedb-fts/src/backend/memory.rs @@ -6,8 +6,9 @@ //! matching the `&self` trait signature. Rebuilt from documents on cold //! start — acceptable for edge-scale datasets. //! -//! Keys are fully structural tuples `(database_id, tid, collection, …)` — -//! database and tenant isolation never depends on lexical-prefix ordering. +//! Keys are fully structural tuples `(database_id, tid, collection, field, …)` +//! — database, tenant, and index isolation never depends on lexical-prefix +//! ordering. use std::cell::RefCell; use std::collections::HashMap; @@ -17,6 +18,7 @@ use nodedb_types::Surrogate; use crate::backend::FtsBackend; use crate::posting::Posting; +use crate::scope::IndexScope; /// In-memory backend error (infallible in practice, but trait requires it). #[derive(Debug)] @@ -28,27 +30,30 @@ impl fmt::Display for MemoryError { } } -type QuadKey = (u64, u64, String, String); -type DocLenKey = (u64, u64, String, Surrogate); -type TripleKey = (u64, u64, String); +/// `(database_id, tid, collection, field)`: one index. +type IndexKey = (u64, u64, String, String); +/// An index key plus a term, meta subkey, or segment id. +type SubKey = (IndexKey, String); +/// An index key plus a document. +type DocLenKey = (IndexKey, Surrogate); /// In-memory FTS backend backed by HashMaps keyed by -/// `(database_id, tid, collection, …)` tuples. +/// `(database_id, tid, collection, field, …)` tuples. /// /// Uses `RefCell` for interior mutability so the `FtsBackend` trait /// can use `&self` uniformly (redb has its own transactional isolation). #[derive(Debug, Default)] pub struct MemoryBackend { - /// `(database_id, tid, collection, term) → posting list`. - postings: RefCell>>, - /// `(database_id, tid, collection, doc_id) → token count`. + /// `(index, term) → posting list`. + postings: RefCell>>, + /// `(index, doc_id) → token count`. doc_lengths: RefCell>, - /// `(database_id, tid, collection) → (doc_count, total_token_sum)`. - stats: RefCell>, - /// `(database_id, tid, collection, subkey) → blob` for docmap, fieldnorms, analyzer, language. - meta: RefCell>>, - /// `(database_id, tid, collection, segment_id) → compressed segment bytes`. - segments: RefCell>>, + /// `index → (doc_count, total_token_sum)`. + stats: RefCell>, + /// `(index, subkey) → blob` for fieldnorms, analyzer, language. + meta: RefCell>>, + /// `(index, segment_id) → compressed segment bytes`. + segments: RefCell>>, } impl MemoryBackend { @@ -57,16 +62,48 @@ impl MemoryBackend { } } -fn quad(database_id: u64, tid: u64, collection: &str, sub: &str) -> QuadKey { - (database_id, tid, collection.to_string(), sub.to_string()) +fn index_key(database_id: u64, tid: u64, index: IndexScope<'_>) -> IndexKey { + ( + database_id, + tid, + index.collection().to_string(), + index.field_key().to_string(), + ) } -fn doc_len_key(database_id: u64, tid: u64, collection: &str, doc_id: Surrogate) -> DocLenKey { - (database_id, tid, collection.to_string(), doc_id) +fn sub_key(database_id: u64, tid: u64, index: IndexScope<'_>, sub: &str) -> SubKey { + (index_key(database_id, tid, index), sub.to_string()) } -fn triple(database_id: u64, tid: u64, collection: &str) -> TripleKey { - (database_id, tid, collection.to_string()) +fn doc_len_key(database_id: u64, tid: u64, index: IndexScope<'_>, doc_id: Surrogate) -> DocLenKey { + (index_key(database_id, tid, index), doc_id) +} + +/// Whether `key` belongs to `(database_id, tid)`, and to `collection` when given. +fn owned_by(key: &IndexKey, database_id: u64, tid: u64, collection: Option<&str>) -> bool { + key.0 == database_id && key.1 == tid && collection.is_none_or(|c| key.2 == c) +} + +impl MemoryBackend { + /// Drop every entry owned by `(database_id, tid[, collection])`. + fn purge_owned(&self, database_id: u64, tid: u64, collection: Option<&str>) -> usize { + let mut postings = self.postings.borrow_mut(); + let mut doc_lengths = self.doc_lengths.borrow_mut(); + let before = postings.len() + doc_lengths.len(); + postings.retain(|k, _| !owned_by(&k.0, database_id, tid, collection)); + doc_lengths.retain(|k, _| !owned_by(&k.0, database_id, tid, collection)); + self.stats + .borrow_mut() + .retain(|k, _| !owned_by(k, database_id, tid, collection)); + self.meta + .borrow_mut() + .retain(|k, _| !owned_by(&k.0, database_id, tid, collection)); + self.segments + .borrow_mut() + .retain(|k, _| !owned_by(&k.0, database_id, tid, collection)); + let after = postings.len() + doc_lengths.len(); + before - after + } } impl FtsBackend for MemoryBackend { @@ -76,13 +113,13 @@ impl FtsBackend for MemoryBackend { &self, database_id: u64, tid: u64, - collection: &str, + index: IndexScope<'_>, term: &str, ) -> Result, Self::Error> { Ok(self .postings .borrow() - .get(&quad(database_id, tid, collection, term)) + .get(&sub_key(database_id, tid, index, term)) .cloned() .unwrap_or_default()) } @@ -91,11 +128,11 @@ impl FtsBackend for MemoryBackend { &self, database_id: u64, tid: u64, - collection: &str, + index: IndexScope<'_>, term: &str, postings: &[Posting], ) -> Result<(), Self::Error> { - let key = quad(database_id, tid, collection, term); + let key = sub_key(database_id, tid, index, term); let mut map = self.postings.borrow_mut(); if postings.is_empty() { map.remove(&key); @@ -109,12 +146,12 @@ impl FtsBackend for MemoryBackend { &self, database_id: u64, tid: u64, - collection: &str, + index: IndexScope<'_>, term: &str, ) -> Result<(), Self::Error> { self.postings .borrow_mut() - .remove(&quad(database_id, tid, collection, term)); + .remove(&sub_key(database_id, tid, index, term)); Ok(()) } @@ -122,13 +159,13 @@ impl FtsBackend for MemoryBackend { &self, database_id: u64, tid: u64, - collection: &str, + index: IndexScope<'_>, doc_id: Surrogate, ) -> Result, Self::Error> { Ok(self .doc_lengths .borrow() - .get(&doc_len_key(database_id, tid, collection, doc_id)) + .get(&doc_len_key(database_id, tid, index, doc_id)) .copied()) } @@ -136,13 +173,13 @@ impl FtsBackend for MemoryBackend { &self, database_id: u64, tid: u64, - collection: &str, + index: IndexScope<'_>, doc_id: Surrogate, length: u32, ) -> Result<(), Self::Error> { self.doc_lengths .borrow_mut() - .insert(doc_len_key(database_id, tid, collection, doc_id), length); + .insert(doc_len_key(database_id, tid, index, doc_id), length); Ok(()) } @@ -150,12 +187,12 @@ impl FtsBackend for MemoryBackend { &self, database_id: u64, tid: u64, - collection: &str, + index: IndexScope<'_>, doc_id: Surrogate, ) -> Result<(), Self::Error> { self.doc_lengths .borrow_mut() - .remove(&doc_len_key(database_id, tid, collection, doc_id)); + .remove(&doc_len_key(database_id, tid, index, doc_id)); Ok(()) } @@ -163,14 +200,15 @@ impl FtsBackend for MemoryBackend { &self, database_id: u64, tid: u64, - collection: &str, + index: IndexScope<'_>, ) -> Result, Self::Error> { + let key = index_key(database_id, tid, index); Ok(self .postings .borrow() .keys() - .filter(|(d, t, c, _)| *d == database_id && *t == tid && c == collection) - .map(|(_, _, _, term)| term.clone()) + .filter(|(k, _)| *k == key) + .map(|(_, term)| term.clone()) .collect()) } @@ -178,12 +216,12 @@ impl FtsBackend for MemoryBackend { &self, database_id: u64, tid: u64, - collection: &str, + index: IndexScope<'_>, ) -> Result<(u32, u64), Self::Error> { Ok(self .stats .borrow() - .get(&triple(database_id, tid, collection)) + .get(&index_key(database_id, tid, index)) .copied() .unwrap_or((0, 0))) } @@ -192,12 +230,12 @@ impl FtsBackend for MemoryBackend { &self, database_id: u64, tid: u64, - collection: &str, + index: IndexScope<'_>, doc_len: u32, ) -> Result<(), Self::Error> { let mut stats = self.stats.borrow_mut(); let entry = stats - .entry(triple(database_id, tid, collection)) + .entry(index_key(database_id, tid, index)) .or_insert((0, 0)); entry.0 += 1; entry.1 += doc_len as u64; @@ -208,12 +246,12 @@ impl FtsBackend for MemoryBackend { &self, database_id: u64, tid: u64, - collection: &str, + index: IndexScope<'_>, doc_len: u32, ) -> Result<(), Self::Error> { let mut stats = self.stats.borrow_mut(); let entry = stats - .entry(triple(database_id, tid, collection)) + .entry(index_key(database_id, tid, index)) .or_insert((0, 0)); entry.0 = entry.0.saturating_sub(1); entry.1 = entry.1.saturating_sub(doc_len as u64); @@ -224,13 +262,13 @@ impl FtsBackend for MemoryBackend { &self, database_id: u64, tid: u64, - collection: &str, + index: IndexScope<'_>, subkey: &str, ) -> Result>, Self::Error> { Ok(self .meta .borrow() - .get(&quad(database_id, tid, collection, subkey)) + .get(&sub_key(database_id, tid, index, subkey)) .cloned()) } @@ -238,13 +276,13 @@ impl FtsBackend for MemoryBackend { &self, database_id: u64, tid: u64, - collection: &str, + index: IndexScope<'_>, subkey: &str, value: &[u8], ) -> Result<(), Self::Error> { self.meta .borrow_mut() - .insert(quad(database_id, tid, collection, subkey), value.to_vec()); + .insert(sub_key(database_id, tid, index, subkey), value.to_vec()); Ok(()) } @@ -252,14 +290,13 @@ impl FtsBackend for MemoryBackend { &self, database_id: u64, tid: u64, - collection: &str, + index: IndexScope<'_>, segment_id: &str, data: &[u8], ) -> Result<(), Self::Error> { - self.segments.borrow_mut().insert( - quad(database_id, tid, collection, segment_id), - data.to_vec(), - ); + self.segments + .borrow_mut() + .insert(sub_key(database_id, tid, index, segment_id), data.to_vec()); Ok(()) } @@ -267,13 +304,13 @@ impl FtsBackend for MemoryBackend { &self, database_id: u64, tid: u64, - collection: &str, + index: IndexScope<'_>, segment_id: &str, ) -> Result>, Self::Error> { Ok(self .segments .borrow() - .get(&quad(database_id, tid, collection, segment_id)) + .get(&sub_key(database_id, tid, index, segment_id)) .cloned()) } @@ -281,14 +318,15 @@ impl FtsBackend for MemoryBackend { &self, database_id: u64, tid: u64, - collection: &str, + index: IndexScope<'_>, ) -> Result, Self::Error> { + let key = index_key(database_id, tid, index); Ok(self .segments .borrow() .keys() - .filter(|(d, t, c, _)| *d == database_id && *t == tid && c == collection) - .map(|(_, _, _, seg)| seg.clone()) + .filter(|(k, _)| *k == key) + .map(|(_, seg)| seg.clone()) .collect()) } @@ -296,12 +334,12 @@ impl FtsBackend for MemoryBackend { &self, database_id: u64, tid: u64, - collection: &str, + index: IndexScope<'_>, segment_id: &str, ) -> Result<(), Self::Error> { self.segments .borrow_mut() - .remove(&quad(database_id, tid, collection, segment_id)); + .remove(&sub_key(database_id, tid, index, segment_id)); Ok(()) } @@ -311,41 +349,11 @@ impl FtsBackend for MemoryBackend { tid: u64, collection: &str, ) -> Result { - let matches_dtc = |d: u64, t: u64, c: &str| d == database_id && t == tid && c == collection; - - let mut postings = self.postings.borrow_mut(); - let mut doc_lengths = self.doc_lengths.borrow_mut(); - let before = postings.len() + doc_lengths.len(); - postings.retain(|k, _| !matches_dtc(k.0, k.1, &k.2)); - doc_lengths.retain(|k, _| !matches_dtc(k.0, k.1, &k.2)); - self.stats - .borrow_mut() - .remove(&triple(database_id, tid, collection)); - self.meta - .borrow_mut() - .retain(|k, _| !matches_dtc(k.0, k.1, &k.2)); - self.segments - .borrow_mut() - .retain(|k, _| !matches_dtc(k.0, k.1, &k.2)); - let after = postings.len() + doc_lengths.len(); - Ok(before - after) + Ok(self.purge_owned(database_id, tid, Some(collection))) } fn purge_tenant(&self, database_id: u64, tid: u64) -> Result { - let matches_dt = |d: u64, t: u64| d == database_id && t == tid; - - let mut postings = self.postings.borrow_mut(); - let mut doc_lengths = self.doc_lengths.borrow_mut(); - let before = postings.len() + doc_lengths.len(); - postings.retain(|k, _| !matches_dt(k.0, k.1)); - doc_lengths.retain(|k, _| !matches_dt(k.0, k.1)); - self.stats.borrow_mut().retain(|k, _| !matches_dt(k.0, k.1)); - self.meta.borrow_mut().retain(|k, _| !matches_dt(k.0, k.1)); - self.segments - .borrow_mut() - .retain(|k, _| !matches_dt(k.0, k.1)); - let after = postings.len() + doc_lengths.len(); - Ok(before - after) + Ok(self.purge_owned(database_id, tid, None)) } } @@ -355,6 +363,16 @@ mod tests { const DB: u64 = 0; const T: u64 = 1; + const COL: IndexScope<'static> = IndexScope::document("col"); + const OTHER: IndexScope<'static> = IndexScope::document("other"); + + fn posting(position: u32) -> Vec { + vec![Posting { + doc_id: Surrogate(1), + term_freq: 1, + positions: vec![position], + }] + } #[test] fn roundtrip_postings() { @@ -365,10 +383,10 @@ mod tests { positions: vec![0, 5], }]; backend - .write_postings(DB, T, "col", "hello", &postings) + .write_postings(DB, T, COL, "hello", &postings) .unwrap(); - let read = backend.read_postings(DB, T, "col", "hello").unwrap(); + let read = backend.read_postings(DB, T, COL, "hello").unwrap(); assert_eq!(read.len(), 1); assert_eq!(read[0].doc_id, Surrogate(1)); } @@ -377,18 +395,16 @@ mod tests { fn roundtrip_doc_lengths() { let backend = MemoryBackend::new(); backend - .write_doc_length(DB, T, "col", Surrogate(1), 42) + .write_doc_length(DB, T, COL, Surrogate(1), 42) .unwrap(); assert_eq!( - backend.read_doc_length(DB, T, "col", Surrogate(1)).unwrap(), + backend.read_doc_length(DB, T, COL, Surrogate(1)).unwrap(), Some(42) ); - backend - .remove_doc_length(DB, T, "col", Surrogate(1)) - .unwrap(); + backend.remove_doc_length(DB, T, COL, Surrogate(1)).unwrap(); assert_eq!( - backend.read_doc_length(DB, T, "col", Surrogate(1)).unwrap(), + backend.read_doc_length(DB, T, COL, Surrogate(1)).unwrap(), None ); } @@ -396,85 +412,102 @@ mod tests { #[test] fn incremental_stats() { let backend = MemoryBackend::new(); - backend.increment_stats(DB, T, "col", 10).unwrap(); - backend.increment_stats(DB, T, "col", 20).unwrap(); - assert_eq!(backend.collection_stats(DB, T, "col").unwrap(), (2, 30)); + backend.increment_stats(DB, T, COL, 10).unwrap(); + backend.increment_stats(DB, T, COL, 20).unwrap(); + assert_eq!(backend.collection_stats(DB, T, COL).unwrap(), (2, 30)); - backend.decrement_stats(DB, T, "col", 10).unwrap(); - assert_eq!(backend.collection_stats(DB, T, "col").unwrap(), (1, 20)); + backend.decrement_stats(DB, T, COL, 10).unwrap(); + assert_eq!(backend.collection_stats(DB, T, COL).unwrap(), (1, 20)); } #[test] fn stats_saturating_sub() { let backend = MemoryBackend::new(); - backend.decrement_stats(DB, T, "col", 100).unwrap(); - assert_eq!(backend.collection_stats(DB, T, "col").unwrap(), (0, 0)); + backend.decrement_stats(DB, T, COL, 100).unwrap(); + assert_eq!(backend.collection_stats(DB, T, COL).unwrap(), (0, 0)); + } + + #[test] + fn field_scopes_are_isolated_from_the_document_scope() { + let backend = MemoryBackend::new(); + let title = IndexScope::field("col", "title").unwrap(); + backend + .write_postings(DB, T, title, "rust", &posting(0)) + .unwrap(); + backend.increment_stats(DB, T, title, 3).unwrap(); + backend + .write_doc_length(DB, T, title, Surrogate(1), 3) + .unwrap(); + + assert!( + backend + .read_postings(DB, T, COL, "rust") + .unwrap() + .is_empty() + ); + assert_eq!(backend.collection_stats(DB, T, COL).unwrap(), (0, 0)); + assert_eq!( + backend.read_doc_length(DB, T, COL, Surrogate(1)).unwrap(), + None + ); + assert_eq!( + backend.collection_terms(DB, T, title).unwrap(), + vec!["rust"] + ); + assert!(backend.collection_terms(DB, T, COL).unwrap().is_empty()); } #[test] fn purge_clears_stats_and_isolates_collections() { let backend = MemoryBackend::new(); - backend.increment_stats(DB, T, "col", 10).unwrap(); + let title = IndexScope::field("col", "title").unwrap(); + backend.increment_stats(DB, T, COL, 10).unwrap(); + backend.increment_stats(DB, T, title, 2).unwrap(); backend - .write_doc_length(DB, T, "col", Surrogate(1), 10) + .write_doc_length(DB, T, COL, Surrogate(1), 10) .unwrap(); backend - .write_postings( - DB, - T, - "col", - "hello", - &[Posting { - doc_id: Surrogate(1), - term_freq: 1, - positions: vec![0], - }], - ) + .write_postings(DB, T, COL, "hello", &posting(0)) + .unwrap(); + backend + .write_postings(DB, T, title, "hello", &posting(0)) .unwrap(); - backend.increment_stats(DB, T, "other", 7).unwrap(); + backend.increment_stats(DB, T, OTHER, 7).unwrap(); backend - .write_doc_length(DB, T, "other", Surrogate(1), 7) + .write_doc_length(DB, T, OTHER, Surrogate(1), 7) .unwrap(); backend - .write_postings( - DB, - T, - "other", - "world", - &[Posting { - doc_id: Surrogate(1), - term_freq: 1, - positions: vec![0], - }], - ) + .write_postings(DB, T, OTHER, "world", &posting(0)) .unwrap(); backend.purge_collection(DB, T, "col").unwrap(); - assert_eq!(backend.collection_stats(DB, T, "col").unwrap(), (0, 0)); + assert_eq!(backend.collection_stats(DB, T, COL).unwrap(), (0, 0)); + assert_eq!(backend.collection_stats(DB, T, title).unwrap(), (0, 0)); + assert!( + backend + .read_postings(DB, T, COL, "hello") + .unwrap() + .is_empty() + ); assert!( backend - .read_postings(DB, T, "col", "hello") + .read_postings(DB, T, title, "hello") .unwrap() .is_empty() ); assert_eq!( - backend.read_doc_length(DB, T, "col", Surrogate(1)).unwrap(), + backend.read_doc_length(DB, T, COL, Surrogate(1)).unwrap(), None ); - assert_eq!(backend.collection_stats(DB, T, "other").unwrap(), (1, 7)); + assert_eq!(backend.collection_stats(DB, T, OTHER).unwrap(), (1, 7)); assert_eq!( - backend - .read_postings(DB, T, "other", "world") - .unwrap() - .len(), + backend.read_postings(DB, T, OTHER, "world").unwrap().len(), 1 ); assert_eq!( - backend - .read_doc_length(DB, T, "other", Surrogate(1)) - .unwrap(), + backend.read_doc_length(DB, T, OTHER, Surrogate(1)).unwrap(), Some(7) ); } @@ -483,33 +516,13 @@ mod tests { fn collection_terms() { let backend = MemoryBackend::new(); backend - .write_postings( - DB, - T, - "col", - "hello", - &[Posting { - doc_id: Surrogate(1), - term_freq: 1, - positions: vec![0], - }], - ) + .write_postings(DB, T, COL, "hello", &posting(0)) .unwrap(); backend - .write_postings( - DB, - T, - "col", - "world", - &[Posting { - doc_id: Surrogate(1), - term_freq: 1, - positions: vec![1], - }], - ) + .write_postings(DB, T, COL, "world", &posting(1)) .unwrap(); - let mut terms = backend.collection_terms(DB, T, "col").unwrap(); + let mut terms = backend.collection_terms(DB, T, COL).unwrap(); terms.sort(); assert_eq!(terms, vec!["hello", "world"]); } @@ -518,124 +531,82 @@ mod tests { fn segment_roundtrip() { let backend = MemoryBackend::new(); let data = b"compressed segment bytes"; - backend.write_segment(DB, T, "col", "id1", data).unwrap(); + backend.write_segment(DB, T, COL, "id1", data).unwrap(); assert_eq!( - backend.read_segment(DB, T, "col", "id1").unwrap(), + backend.read_segment(DB, T, COL, "id1").unwrap(), Some(data.to_vec()) ); - assert_eq!(backend.read_segment(DB, T, "col", "missing").unwrap(), None); + assert_eq!(backend.read_segment(DB, T, COL, "missing").unwrap(), None); } #[test] - fn segment_list_filters_by_collection() { + fn segment_list_filters_by_index() { let backend = MemoryBackend::new(); - backend.write_segment(DB, T, "col", "a", b"a").unwrap(); - backend.write_segment(DB, T, "col", "b", b"b").unwrap(); - backend.write_segment(DB, T, "other", "c", b"c").unwrap(); + let title = IndexScope::field("col", "title").unwrap(); + backend.write_segment(DB, T, COL, "a", b"a").unwrap(); + backend.write_segment(DB, T, COL, "b", b"b").unwrap(); + backend.write_segment(DB, T, OTHER, "c", b"c").unwrap(); + backend.write_segment(DB, T, title, "d", b"d").unwrap(); - let mut segs = backend.list_segments(DB, T, "col").unwrap(); + let mut segs = backend.list_segments(DB, T, COL).unwrap(); segs.sort(); assert_eq!(segs, vec!["a", "b"]); - let other = backend.list_segments(DB, T, "other").unwrap(); - assert_eq!(other, vec!["c"]); + assert_eq!(backend.list_segments(DB, T, OTHER).unwrap(), vec!["c"]); + assert_eq!(backend.list_segments(DB, T, title).unwrap(), vec!["d"]); } #[test] fn segment_remove() { let backend = MemoryBackend::new(); - backend.write_segment(DB, T, "col", "id1", b"data").unwrap(); - backend.remove_segment(DB, T, "col", "id1").unwrap(); - assert_eq!(backend.read_segment(DB, T, "col", "id1").unwrap(), None); + backend.write_segment(DB, T, COL, "id1", b"data").unwrap(); + backend.remove_segment(DB, T, COL, "id1").unwrap(); + assert_eq!(backend.read_segment(DB, T, COL, "id1").unwrap(), None); } #[test] fn purge_clears_segments() { let backend = MemoryBackend::new(); - backend.write_segment(DB, T, "col", "a", b"a").unwrap(); - backend.write_segment(DB, T, "other", "b", b"b").unwrap(); + backend.write_segment(DB, T, COL, "a", b"a").unwrap(); + backend.write_segment(DB, T, OTHER, "b", b"b").unwrap(); backend.purge_collection(DB, T, "col").unwrap(); - assert!(backend.list_segments(DB, T, "col").unwrap().is_empty()); - assert_eq!(backend.list_segments(DB, T, "other").unwrap().len(), 1); + assert!(backend.list_segments(DB, T, COL).unwrap().is_empty()); + assert_eq!(backend.list_segments(DB, T, OTHER).unwrap().len(), 1); } #[test] fn purge_tenant_isolates_tenants() { let backend = MemoryBackend::new(); - backend.increment_stats(DB, 1, "col", 5).unwrap(); - backend.increment_stats(DB, 2, "col", 7).unwrap(); + backend.increment_stats(DB, 1, COL, 5).unwrap(); + backend.increment_stats(DB, 2, COL, 7).unwrap(); backend - .write_postings( - DB, - 1, - "col", - "t", - &[Posting { - doc_id: Surrogate(1), - term_freq: 1, - positions: vec![0], - }], - ) + .write_postings(DB, 1, COL, "t", &posting(0)) .unwrap(); backend - .write_postings( - DB, - 2, - "col", - "t", - &[Posting { - doc_id: Surrogate(1), - term_freq: 1, - positions: vec![0], - }], - ) + .write_postings(DB, 2, COL, "t", &posting(0)) .unwrap(); backend.purge_tenant(DB, 1).unwrap(); - assert_eq!(backend.collection_stats(DB, 1, "col").unwrap(), (0, 0)); - assert!(backend.read_postings(DB, 1, "col", "t").unwrap().is_empty()); - assert_eq!(backend.collection_stats(DB, 2, "col").unwrap(), (1, 7)); - assert_eq!(backend.read_postings(DB, 2, "col", "t").unwrap().len(), 1); + assert_eq!(backend.collection_stats(DB, 1, COL).unwrap(), (0, 0)); + assert!(backend.read_postings(DB, 1, COL, "t").unwrap().is_empty()); + assert_eq!(backend.collection_stats(DB, 2, COL).unwrap(), (1, 7)); + assert_eq!(backend.read_postings(DB, 2, COL, "t").unwrap().len(), 1); } #[test] fn databases_isolated() { let backend = MemoryBackend::new(); - backend.increment_stats(0, T, "col", 5).unwrap(); - backend.increment_stats(9, T, "col", 7).unwrap(); - backend - .write_postings( - 0, - T, - "col", - "t", - &[Posting { - doc_id: Surrogate(1), - term_freq: 1, - positions: vec![0], - }], - ) - .unwrap(); - backend - .write_postings( - 9, - T, - "col", - "t", - &[Posting { - doc_id: Surrogate(1), - term_freq: 1, - positions: vec![0], - }], - ) - .unwrap(); + backend.increment_stats(0, T, COL, 5).unwrap(); + backend.increment_stats(9, T, COL, 7).unwrap(); + backend.write_postings(0, T, COL, "t", &posting(0)).unwrap(); + backend.write_postings(9, T, COL, "t", &posting(0)).unwrap(); backend.purge_tenant(0, T).unwrap(); - assert_eq!(backend.collection_stats(0, T, "col").unwrap(), (0, 0)); - assert!(backend.read_postings(0, T, "col", "t").unwrap().is_empty()); + assert_eq!(backend.collection_stats(0, T, COL).unwrap(), (0, 0)); + assert!(backend.read_postings(0, T, COL, "t").unwrap().is_empty()); // Same tenant in a different database must be unaffected. - assert_eq!(backend.collection_stats(9, T, "col").unwrap(), (1, 7)); - assert_eq!(backend.read_postings(9, T, "col", "t").unwrap().len(), 1); + assert_eq!(backend.collection_stats(9, T, COL).unwrap(), (1, 7)); + assert_eq!(backend.read_postings(9, T, COL, "t").unwrap().len(), 1); } } diff --git a/nodedb-fts/src/backend/traits.rs b/nodedb-fts/src/backend/traits.rs index 5268513ce..49df66e72 100644 --- a/nodedb-fts/src/backend/traits.rs +++ b/nodedb-fts/src/backend/traits.rs @@ -3,6 +3,7 @@ use nodedb_types::Surrogate; use crate::posting::Posting; +use crate::scope::IndexScope; /// Storage backend abstraction for the full-text search engine. /// @@ -15,6 +16,13 @@ use crate::posting::Posting; /// and tenants structurally — no boundary may depend on lexical-prefix /// ordering of a composed string key. /// +/// Per-index methods take an [`IndexScope`]: a collection's whole-document +/// index or one of its field indexes. Backends key every per-index entry by +/// both the collection and the scope's field key, so no two scopes share +/// postings, lengths, stats, metadata, or segments. Collection-level +/// configuration (analyzer, language, fuzzy) lives in the metadata of +/// `IndexScope::document(collection)`. +/// /// Write methods take `&self` (not `&mut self`) because: /// - Redb provides transactional isolation internally — concurrent writes /// are safe through redb's MVCC. @@ -24,21 +32,21 @@ pub trait FtsBackend { /// Error type for backend operations. type Error: std::fmt::Display; - /// Read the posting list for a term in a collection. + /// Read the posting list for a term in an index. fn read_postings( &self, database_id: u64, tid: u64, - collection: &str, + index: IndexScope<'_>, term: &str, ) -> Result, Self::Error>; - /// Write/replace the posting list for a term in a collection. + /// Write/replace the posting list for a term in an index. fn write_postings( &self, database_id: u64, tid: u64, - collection: &str, + index: IndexScope<'_>, term: &str, postings: &[Posting], ) -> Result<(), Self::Error>; @@ -48,47 +56,62 @@ pub trait FtsBackend { &self, database_id: u64, tid: u64, - collection: &str, + index: IndexScope<'_>, term: &str, ) -> Result<(), Self::Error>; - /// Read the document length (token count) for a document. + /// Read the document length (token count) of a document in an index. fn read_doc_length( &self, database_id: u64, tid: u64, - collection: &str, + index: IndexScope<'_>, doc_id: Surrogate, ) -> Result, Self::Error>; - /// Write/replace the document length for a document. + /// Read the document lengths of `doc_ids` in an index, parallel to + /// `doc_ids`. A backend with transactions reads them in one. + fn read_doc_lengths( + &self, + database_id: u64, + tid: u64, + index: IndexScope<'_>, + doc_ids: &[Surrogate], + ) -> Result>, Self::Error> { + doc_ids + .iter() + .map(|doc_id| self.read_doc_length(database_id, tid, index, *doc_id)) + .collect() + } + + /// Write/replace the document length of a document in an index. fn write_doc_length( &self, database_id: u64, tid: u64, - collection: &str, + index: IndexScope<'_>, doc_id: Surrogate, length: u32, ) -> Result<(), Self::Error>; - /// Remove a document's length entry. + /// Remove a document's length entry from an index. fn remove_doc_length( &self, database_id: u64, tid: u64, - collection: &str, + index: IndexScope<'_>, doc_id: Surrogate, ) -> Result<(), Self::Error>; - /// Get all term names in a collection (for fuzzy matching). + /// Get all term names in an index (for fuzzy matching). fn collection_terms( &self, database_id: u64, tid: u64, - collection: &str, + index: IndexScope<'_>, ) -> Result, Self::Error>; - /// Get total document count and sum of all document lengths for a collection. + /// Get total document count and sum of all document lengths of an index. /// Returns `(doc_count, total_token_sum)`. /// /// Implementations should maintain these incrementally for O(1) lookup. @@ -96,26 +119,26 @@ pub trait FtsBackend { &self, database_id: u64, tid: u64, - collection: &str, + index: IndexScope<'_>, ) -> Result<(u32, u64), Self::Error>; - /// Increment collection stats after indexing a document. + /// Increment index stats after indexing a document. /// `doc_len` is the number of tokens in the newly indexed document. fn increment_stats( &self, database_id: u64, tid: u64, - collection: &str, + index: IndexScope<'_>, doc_len: u32, ) -> Result<(), Self::Error>; - /// Decrement collection stats after removing a document. + /// Decrement index stats after removing a document. /// `doc_len` is the token count of the removed document. fn decrement_stats( &self, database_id: u64, tid: u64, - collection: &str, + index: IndexScope<'_>, doc_len: u32, ) -> Result<(), Self::Error>; @@ -125,7 +148,7 @@ pub trait FtsBackend { &self, database_id: u64, tid: u64, - collection: &str, + index: IndexScope<'_>, subkey: &str, ) -> Result>, Self::Error>; @@ -134,18 +157,18 @@ pub trait FtsBackend { &self, database_id: u64, tid: u64, - collection: &str, + index: IndexScope<'_>, subkey: &str, value: &[u8], ) -> Result<(), Self::Error>; - /// Write a segment blob. `segment_id` is a stable per-collection + /// Write a segment blob. `segment_id` is a stable per-index /// identifier (e.g., `"L{level}:{id:016x}"`). fn write_segment( &self, database_id: u64, tid: u64, - collection: &str, + index: IndexScope<'_>, segment_id: &str, data: &[u8], ) -> Result<(), Self::Error>; @@ -155,16 +178,16 @@ pub trait FtsBackend { &self, database_id: u64, tid: u64, - collection: &str, + index: IndexScope<'_>, segment_id: &str, ) -> Result>, Self::Error>; - /// List all segment ids for a collection. + /// List all segment ids of an index. fn list_segments( &self, database_id: u64, tid: u64, - collection: &str, + index: IndexScope<'_>, ) -> Result, Self::Error>; /// Remove a segment blob. @@ -172,11 +195,12 @@ pub trait FtsBackend { &self, database_id: u64, tid: u64, - collection: &str, + index: IndexScope<'_>, segment_id: &str, ) -> Result<(), Self::Error>; - /// Remove all entries for a collection. Returns count of removed entries. + /// Remove all entries of every index of a collection. Returns count of + /// removed entries. fn purge_collection( &self, database_id: u64, diff --git a/nodedb-fts/src/document_text.rs b/nodedb-fts/src/document_text.rs new file mode 100644 index 000000000..3a770aed8 --- /dev/null +++ b/nodedb-fts/src/document_text.rs @@ -0,0 +1,145 @@ +// SPDX-License-Identifier: Apache-2.0 + +//! A document's indexable text: its top-level string fields, in field-name order. + +use std::borrow::Cow; +use std::collections::BTreeMap; + +use crate::scope::IndexScope; + +/// The `(field, text)` pairs of one document, sorted by field name. +#[derive(Debug, Clone, Default, PartialEq, Eq)] +pub struct DocumentText { + /// Sorted by field name, names unique. + fields: Vec<(String, String)>, +} + +impl DocumentText { + /// Collect `(field, text)` pairs. A repeated name keeps its last text. + pub fn from_fields>(fields: I) -> Self { + let sorted: BTreeMap = fields.into_iter().collect(); + Self { + fields: sorted.into_iter().collect(), + } + } + + /// The `(field, text)` pairs, sorted by field name. + pub fn fields(&self) -> &[(String, String)] { + &self.fields + } + + /// Consume into the sorted `(field, text)` pairs. + pub fn into_fields(self) -> Vec<(String, String)> { + self.fields + } + + /// Whether no field holds any text. + pub fn is_empty(&self) -> bool { + self.fields.iter().all(|(_, t)| t.is_empty()) + } + + /// Every field's text, in field-name order, joined by one space. + pub fn whole(&self) -> String { + let mut out = String::new(); + for (i, (_, text)) in self.fields.iter().enumerate() { + if i != 0 { + out.push(' '); + } + out.push_str(text); + } + out + } + + /// Each field that has an index of its own, with its scope. + pub fn field_scopes<'s>( + &'s self, + collection: &'s str, + ) -> impl Iterator, &'s str)> { + self.fields + .iter() + .filter_map(move |(f, t)| IndexScope::field(collection, f).map(|s| (s, t.as_str()))) + } + + /// The text one index holds for this document: empty when the field is absent. + pub fn text_of(&self, index: IndexScope<'_>) -> Cow<'_, str> { + match index.field_name() { + None => Cow::Owned(self.whole()), + Some(f) => Cow::Borrowed( + self.fields + .iter() + .find(|(n, _)| n == f) + .map_or("", |(_, t)| t.as_str()), + ), + } + } +} + +#[cfg(test)] +mod tests { + use super::*; + + fn pairs() -> Vec<(String, String)> { + vec![ + ("title".into(), "Rust book".into()), + ("body".into(), "fearless concurrency".into()), + ("author".into(), "".into()), + ("zeta".into(), "last".into()), + ] + } + + /// The join every deployment indexed before per-field scopes: string + /// values sorted by field name, joined by one space. + fn reference_join(mut texts: Vec<(String, String)>) -> String { + texts.sort_by(|a, b| a.0.cmp(&b.0)); + texts + .into_iter() + .map(|(_, t)| t) + .collect::>() + .join(" ") + } + + #[test] + fn whole_matches_the_field_name_ordered_join() { + let text = DocumentText::from_fields(pairs()); + assert_eq!(text.whole(), reference_join(pairs())); + assert_eq!(text.whole(), " fearless concurrency Rust book last"); + } + + #[test] + fn repeated_name_keeps_last_text() { + let text = DocumentText::from_fields([ + ("a".to_string(), "one".to_string()), + ("a".to_string(), "two".to_string()), + ]); + assert_eq!(text.fields(), &[("a".to_string(), "two".to_string())]); + } + + #[test] + fn text_of_resolves_document_and_field_scopes() { + let text = DocumentText::from_fields(pairs()); + let title = IndexScope::field("c", "title").expect("non-empty field"); + let missing = IndexScope::field("c", "missing").expect("non-empty field"); + assert_eq!(text.text_of(title), "Rust book"); + assert_eq!(text.text_of(missing), ""); + assert_eq!(text.text_of(IndexScope::document("c")), text.whole()); + } + + #[test] + fn field_scopes_skip_empty_names() { + let text = DocumentText::from_fields([ + (String::new(), "orphan".to_string()), + ("title".to_string(), "kept".to_string()), + ]); + let scopes: Vec<_> = text.field_scopes("c").collect(); + assert_eq!(scopes.len(), 1); + assert_eq!(scopes[0].0.field_name(), Some("title")); + assert_eq!(scopes[0].1, "kept"); + } + + #[test] + fn empty_when_no_field_has_text() { + assert!(DocumentText::default().is_empty()); + assert!(DocumentText::from_fields([("a".to_string(), String::new())]).is_empty()); + assert!(!DocumentText::from_fields(pairs()).is_empty()); + } +} diff --git a/nodedb-fts/src/index/analyzer_config.rs b/nodedb-fts/src/index/analyzer_config.rs index 0f36b5c91..f2f47ede2 100644 --- a/nodedb-fts/src/index/analyzer_config.rs +++ b/nodedb-fts/src/index/analyzer_config.rs @@ -2,7 +2,8 @@ //! Per-collection analyzer configuration stored in backend metadata. //! -//! Uses structural `(tid, collection, subkey)` meta blobs: +//! Uses structural meta blobs of the collection's whole-document index +//! (`IndexScope::document(collection)`), shared by every field index: //! - `subkey = "analyzer"` → analyzer name (e.g. "german", "standard") //! - `subkey = "language"` → lang code (e.g. "de", "ja") //! - `subkey = "fuzzy"` → `"1"` when the collection defaults to fuzzy matching @@ -14,6 +15,7 @@ use crate::analyzer::pipeline::{TextAnalyzer, analyze}; use crate::analyzer::standard::StandardAnalyzer; use crate::backend::FtsBackend; use crate::index::FtsIndex; +use crate::scope::IndexScope; impl FtsIndex { /// Set the analyzer for a collection. Persists to backend metadata. @@ -27,7 +29,7 @@ impl FtsIndex { self.backend.write_meta( database_id, tid, - collection, + IndexScope::document(collection), "analyzer", analyzer_name.as_bytes(), ) @@ -44,7 +46,7 @@ impl FtsIndex { self.backend.write_meta( database_id, tid, - collection, + IndexScope::document(collection), "language", lang_code.as_bytes(), ) @@ -62,7 +64,7 @@ impl FtsIndex { self.backend.write_meta( database_id, tid, - collection, + IndexScope::document(collection), "fuzzy", if fuzzy { b"1" } else { b"0" }, ) @@ -77,7 +79,7 @@ impl FtsIndex { ) -> Result { Ok(self .backend - .read_meta(database_id, tid, collection, "fuzzy")? + .read_meta(database_id, tid, IndexScope::document(collection), "fuzzy")? .is_some_and(|bytes| bytes.as_slice() == b"1")) } @@ -88,10 +90,12 @@ impl FtsIndex { tid: u64, collection: &str, ) -> Result, B::Error> { - match self - .backend - .read_meta(database_id, tid, collection, "analyzer")? - { + match self.backend.read_meta( + database_id, + tid, + IndexScope::document(collection), + "analyzer", + )? { Some(bytes) => Ok(std::str::from_utf8(&bytes).ok().map(String::from)), None => Ok(None), } @@ -104,10 +108,12 @@ impl FtsIndex { tid: u64, collection: &str, ) -> Result, B::Error> { - match self - .backend - .read_meta(database_id, tid, collection, "language")? - { + match self.backend.read_meta( + database_id, + tid, + IndexScope::document(collection), + "language", + )? { Some(bytes) => Ok(std::str::from_utf8(&bytes).ok().map(String::from)), None => Ok(None), } diff --git a/nodedb-fts/src/index/error.rs b/nodedb-fts/src/index/error.rs index fb5583e83..a6474f3db 100644 --- a/nodedb-fts/src/index/error.rs +++ b/nodedb-fts/src/index/error.rs @@ -79,6 +79,34 @@ pub enum FtsIndexError { #[error("FTS segment error: {0}")] Segment(crate::lsm::segment::error::SegmentError), + /// A stored segment of the index fails validation when it is opened. + /// + /// The segment's postings cannot be read, so a read that skipped it + /// would answer from part of the index. + #[error("FTS segment {segment_id} is corrupt: {source}")] + CorruptSegment { + segment_id: String, + source: crate::lsm::segment::error::SegmentError, + }, + + /// The backend lists a segment of the index that it does not hold. + #[error("FTS segment {segment_id} is listed but missing")] + MissingSegment { segment_id: String }, + + /// A stored index-state blob (backend metadata `subkey`) fails to decode. + #[error("FTS index state '{subkey}' is corrupt: {detail}")] + CorruptState { + subkey: &'static str, + detail: String, + }, + + /// An index-state blob (backend metadata `subkey`) fails to encode. + #[error("FTS index state '{subkey}' failed to encode: {detail}")] + StateEncode { + subkey: &'static str, + detail: String, + }, + /// Memory budget exhausted for the FTS engine. /// /// The operation requires more memory than the engine's remaining budget diff --git a/nodedb-fts/src/index/fieldnorm.rs b/nodedb-fts/src/index/fieldnorm.rs index 2f5560806..630e0ee2b 100644 --- a/nodedb-fts/src/index/fieldnorm.rs +++ b/nodedb-fts/src/index/fieldnorm.rs @@ -11,21 +11,22 @@ use nodedb_types::Surrogate; use crate::backend::FtsBackend; use crate::codec::smallfloat; use crate::index::FtsIndex; +use crate::scope::IndexScope; impl FtsIndex { /// Get the fieldnorm (SmallFloat-encoded doc length) for a doc. /// /// Returns the decoded approximate u32 length, or `None` if not stored. - pub fn read_fieldnorm( + pub fn read_fieldnorm<'a>( &self, database_id: u64, tid: u64, - collection: &str, + index: impl Into>, doc_id: Surrogate, ) -> Result, B::Error> { let data = self .backend - .read_meta(database_id, tid, collection, "fieldnorms")?; + .read_meta(database_id, tid, index.into(), "fieldnorms")?; match data { Some(bytes) if (doc_id.0 as usize) < bytes.len() => { Ok(Some(smallfloat::decode(bytes[doc_id.0 as usize]))) @@ -35,17 +36,18 @@ impl FtsIndex { } /// Write a fieldnorm byte for a surrogate. Grows the array if needed. - pub fn write_fieldnorm( + pub fn write_fieldnorm<'a>( &self, database_id: u64, tid: u64, - collection: &str, + index: impl Into>, doc_id: Surrogate, doc_length: u32, ) -> Result<(), B::Error> { + let index = index.into(); let mut data = self .backend - .read_meta(database_id, tid, collection, "fieldnorms")? + .read_meta(database_id, tid, index, "fieldnorms")? .unwrap_or_default(); let idx = doc_id.0 as usize; @@ -55,7 +57,7 @@ impl FtsIndex { data[idx] = smallfloat::encode(doc_length); self.backend - .write_meta(database_id, tid, collection, "fieldnorms", &data) + .write_meta(database_id, tid, index, "fieldnorms", &data) } } diff --git a/nodedb-fts/src/index/stats.rs b/nodedb-fts/src/index/stats.rs index 7b5209f12..2c12de130 100644 --- a/nodedb-fts/src/index/stats.rs +++ b/nodedb-fts/src/index/stats.rs @@ -4,21 +4,22 @@ use crate::backend::FtsBackend; use crate::index::FtsIndex; +use crate::scope::IndexScope; impl FtsIndex { - /// Get total document count and average document length for a collection. + /// Get total document count and average document length of one index. /// - /// Returns `(total_docs, avg_doc_len)`. If the collection is empty, + /// Returns `(total_docs, avg_doc_len)`. If the index is empty, /// returns `(0, 1.0)` to avoid division by zero. - pub fn index_stats( + pub fn index_stats<'a>( &self, database_id: u64, tid: u64, - collection: &str, + index: impl Into>, ) -> Result<(u32, f32), B::Error> { let (count, total_len) = self .backend - .collection_stats(database_id, tid, collection)?; + .collection_stats(database_id, tid, index.into())?; let avg = if count > 0 { total_len as f32 / count as f32 } else { diff --git a/nodedb-fts/src/index/synonym_groups.rs b/nodedb-fts/src/index/synonym_groups.rs index 3664e46f0..ecd1af2e8 100644 --- a/nodedb-fts/src/index/synonym_groups.rs +++ b/nodedb-fts/src/index/synonym_groups.rs @@ -18,9 +18,10 @@ use crate::analyzer::synonym::SynonymMap; use crate::backend::FtsBackend; use crate::index::writer::FtsIndex; +use crate::scope::IndexScope; -/// Sentinel collection name for synonym group meta storage. -const SYNONYM_GROUPS_COLLECTION: &str = "_synonym_groups"; +/// Sentinel index whose metadata holds the synonym groups. +const SYNONYM_GROUPS_COLLECTION: IndexScope<'static> = IndexScope::document("_synonym_groups"); /// Special meta subkey that holds the JSON array of all group names. const INDEX_SUBKEY: &str = "_index"; @@ -178,6 +179,31 @@ impl FtsIndex { Ok(expanded) } + /// Expand each analyzed query token into its word group: the token + /// first, then its distinct synonyms. One group per token, in token + /// order. A document matches a word when it holds any term of the group. + pub fn expand_query_groups( + &self, + database_id: u64, + tid: u64, + tokens: Vec, + ) -> Result>, B::Error> { + let groups = self.list_synonym_groups(database_id, tid)?; + if groups.is_empty() { + return Ok(tokens.into_iter().map(|token| vec![token]).collect()); + } + let map = self.build_synonym_map_for_tenant(database_id, tid, &groups); + Ok(tokens + .into_iter() + .map(|token| { + let mut group = map.expand(std::slice::from_ref(&token)); + let mut seen = std::collections::HashSet::new(); + group.retain(|term| seen.insert(term.clone())); + group + }) + .collect()) + } + // ── internal helpers ────────────────────────────────────────────────────── fn read_name_index(&self, database_id: u64, tid: u64) -> Result, B::Error> { diff --git a/nodedb-fts/src/index/writer.rs b/nodedb-fts/src/index/writer.rs deleted file mode 100644 index 6c9e8d938..000000000 --- a/nodedb-fts/src/index/writer.rs +++ /dev/null @@ -1,489 +0,0 @@ -// SPDX-License-Identifier: Apache-2.0 - -//! Core FtsIndex: indexing and document management over any backend. - -use std::collections::HashMap; -use std::sync::Arc; -use std::sync::atomic::{AtomicU64, Ordering}; - -use nodedb_types::Surrogate; -use tracing::debug; - -use crate::backend::FtsBackend; - -use crate::block::CompactPosting; -use crate::codec::smallfloat; -use crate::index::error::{FtsIndexError, MAX_INDEXABLE_SURROGATE}; -use crate::lsm::compaction; -use crate::lsm::memtable::{Memtable, MemtableConfig}; -use crate::lsm::segment::writer as seg_writer; -use crate::posting::Bm25Params; -use nodedb_mem::MemoryGovernor; - -/// Full-text search index generic over storage backend. -/// -/// Provides identical indexing, search, and highlighting logic -/// for Origin (redb), Lite (in-memory), and WASM deployments. -/// -/// Writes accumulate in an in-memory `Memtable`. When the memtable -/// exceeds its threshold, it is flushed to an immutable segment -/// stored via the backend. Queries merge the active memtable with -/// all persisted segments. -/// -/// [`MemoryGovernor`] enforces per-engine memory budgets on large -/// allocations (compaction, segment merge, query term collection). -pub struct FtsIndex { - pub(crate) backend: B, - pub(crate) bm25_params: Bm25Params, - pub(crate) memtable: Memtable, - /// Monotonic segment ID counter. - next_segment_id: AtomicU64, - /// Memory governor for budget enforcement. - pub(crate) governor: Arc, -} - -impl FtsIndex { - /// Create a new FTS index with the given backend, default BM25 params, - /// and a memory governor. - pub fn new(backend: B, governor: Arc) -> Self { - Self { - backend, - bm25_params: Bm25Params::default(), - memtable: Memtable::new(MemtableConfig::default()), - next_segment_id: AtomicU64::new(1), - governor, - } - } - - /// Create a new FTS index with custom BM25 parameters and a memory governor. - pub fn with_params(backend: B, params: Bm25Params, governor: Arc) -> Self { - Self { - backend, - bm25_params: params, - memtable: Memtable::new(MemtableConfig::default()), - next_segment_id: AtomicU64::new(1), - governor, - } - } - - /// Access the underlying backend. - pub fn backend(&self) -> &B { - &self.backend - } - - /// Mutable access to the underlying backend. - pub fn backend_mut(&mut self) -> &mut B { - &mut self.backend - } - - /// Access the active memtable (for LSM query merging). - pub fn memtable(&self) -> &Memtable { - &self.memtable - } - - /// Index a document's text content. - /// - /// Returns `Err(FtsIndexError::SurrogateOutOfRange)` if `doc_id` is - /// `Surrogate::ZERO` (the unassigned sentinel) or exceeds - /// `MAX_INDEXABLE_SURROGATE`. The FTS memtable uses the surrogate's raw - /// `u32` value as a direct array index into per-doc fieldnorm storage; - /// values near `u32::MAX` would cause multi-GiB allocations. Rejecting - /// out-of-range surrogates at this boundary is the correct fix — not a - /// `debug_assert!`, which would be a silent-wrap equivalent. - pub fn index_document( - &self, - database_id: u64, - tid: u64, - collection: &str, - doc_id: Surrogate, - text: &str, - ) -> Result<(), FtsIndexError> { - let raw = doc_id.as_u32(); - if raw == 0 || raw > MAX_INDEXABLE_SURROGATE { - return Err(FtsIndexError::SurrogateOutOfRange { surrogate: doc_id }); - } - - let tokens = self - .analyze_for_collection(database_id, tid, collection, text) - .map_err(FtsIndexError::backend)?; - if tokens.is_empty() { - return Ok(()); - } - - let mut term_data: HashMap<&str, (u32, Vec)> = HashMap::new(); - for (pos, token) in tokens.iter().enumerate() { - let entry = term_data.entry(token.as_str()).or_insert((0, Vec::new())); - entry.0 += 1; - entry.1.push(pos as u32); - } - - let doc_len = tokens.len() as u32; - - for (term, (freq, positions)) in &term_data { - let compact = CompactPosting { - doc_id, - term_freq: *freq, - fieldnorm: smallfloat::encode(doc_len), - positions: positions.clone(), - }; - let scoped_term = memtable_key(database_id, tid, collection, term); - self.memtable.insert(&scoped_term, compact); - } - self.memtable.record_doc(doc_id, doc_len); - - // Write document length, fieldnorm, and update incremental stats. - self.backend - .write_doc_length(database_id, tid, collection, doc_id, doc_len) - .map_err(FtsIndexError::backend)?; - self.write_fieldnorm(database_id, tid, collection, doc_id, doc_len) - .map_err(FtsIndexError::backend)?; - self.backend - .increment_stats(database_id, tid, collection, doc_len) - .map_err(FtsIndexError::backend)?; - - if self.memtable.should_flush() { - self.flush_memtable(database_id, tid, collection)?; - } - - debug!(database_id, tid, %collection, doc_id = doc_id.0, tokens = tokens.len(), terms = term_data.len(), "indexed document"); - Ok(()) - } - - /// Flush the active memtable to an immutable segment in the backend. - /// - /// Calling this before serializing the index guarantees that all posting - /// data written since the last spill threshold is captured in the backend's - /// segment storage rather than the in-memory memtable. Callers that - /// checkpoint the index (e.g., NodeDB-Lite flush) must call this once per - /// active index before persisting. - pub fn flush_memtable( - &self, - database_id: u64, - tid: u64, - collection: &str, - ) -> Result<(), FtsIndexError> { - let drained = self.memtable.drain(); - if drained.is_empty() { - return Ok(()); - } - - let segment_bytes = seg_writer::flush_to_segment(drained)?; - let seg_id = self.next_segment_id.fetch_add(1, Ordering::Relaxed); - let id = compaction::segment_id(seg_id, 0); - self.backend - .write_segment(database_id, tid, collection, &id, &segment_bytes) - .map_err(FtsIndexError::backend)?; - - debug!(database_id, tid, %collection, seg_id, bytes = segment_bytes.len(), "flushed memtable to segment"); - Ok(()) - } - - /// Remove a document from the index. - pub fn remove_document( - &self, - database_id: u64, - tid: u64, - collection: &str, - doc_id: Surrogate, - ) -> Result<(), B::Error> { - let doc_len = self - .backend - .read_doc_length(database_id, tid, collection, doc_id)?; - - self.memtable.remove_doc(doc_id); - self.backend - .remove_doc_length(database_id, tid, collection, doc_id)?; - - if let Some(len) = doc_len { - self.backend - .decrement_stats(database_id, tid, collection, len)?; - } - - Ok(()) - } - - /// Purge all entries for a collection. Returns count of removed entries. - pub fn purge_collection( - &self, - database_id: u64, - tid: u64, - collection: &str, - ) -> Result { - self.memtable - .drain_collection(&memtable_collection_prefix(database_id, tid, collection)); - self.backend.purge_collection(database_id, tid, collection) - } - - /// Purge all entries for a `(database_id, tenant)` across every collection. - pub fn purge_tenant(&self, database_id: u64, tid: u64) -> Result { - self.memtable - .drain_collection(&memtable_tenant_prefix(database_id, tid)); - self.backend.purge_tenant(database_id, tid) - } -} - -/// Memtable key format: `"{database_id}:{tid}:{collection}:{term}"`. The -/// memtable is a single in-memory map shared across databases and tenants, -/// so keys must carry the full database + tenant + collection scope. -pub(crate) fn memtable_key(database_id: u64, tid: u64, collection: &str, term: &str) -> String { - format!("{database_id}:{tid}:{collection}:{term}") -} - -/// Prefix used by `drain_collection` to remove all memtable entries for -/// a given `(database_id, tid, collection)`. -pub(crate) fn memtable_collection_prefix(database_id: u64, tid: u64, collection: &str) -> String { - format!("{database_id}:{tid}:{collection}:") -} - -/// Prefix used to remove every memtable entry for a given `(database_id, tenant)`. -pub(crate) fn memtable_tenant_prefix(database_id: u64, tid: u64) -> String { - format!("{database_id}:{tid}:") -} - -#[cfg(test)] -mod tests { - use nodedb_types::Surrogate; - - use crate::backend::memory::MemoryBackend; - use crate::test_support::test_governor; - - use super::*; - - const DB: u64 = 0; - const T: u64 = 1; - - fn make_index() -> FtsIndex { - FtsIndex::new(MemoryBackend::new(), test_governor()) - } - - #[test] - fn flush_propagates_term_too_long_as_typed_error() { - let backend = MemoryBackend::new(); - let idx = FtsIndex { - backend, - bm25_params: Bm25Params::default(), - memtable: Memtable::new(MemtableConfig { - max_postings: 1, - max_terms: 1, - }), - next_segment_id: AtomicU64::new(1), - governor: test_governor(), - }; - - // Insert a single posting under a term whose byte length exceeds the - // u16 segment-format cap. Bypasses the analyzer (which would tokenize - // away most pathological inputs); we want to exercise the flush-path - // boundary check directly. - let oversize_term = "x".repeat(crate::lsm::segment::format::MAX_TERM_LEN + 1); - idx.memtable.insert( - &super::memtable_key(DB, T, "docs", &oversize_term), - CompactPosting { - doc_id: Surrogate(1), - term_freq: 1, - fieldnorm: 1, - positions: vec![0], - }, - ); - idx.memtable.record_doc(Surrogate(1), 1); - - let err = idx - .flush_memtable(DB, T, "docs") - .expect_err("flush must reject oversize term"); - let key_overhead = super::memtable_key(DB, T, "docs", "").len(); - match err { - FtsIndexError::TermTooLong { len, max } => { - assert_eq!(len, oversize_term.len() + key_overhead); - assert_eq!(max, crate::lsm::segment::format::MAX_TERM_LEN); - } - other => panic!("expected TermTooLong, got {other:?}"), - } - } - - #[test] - fn index_writes_to_memtable() { - let idx = make_index(); - idx.index_document(DB, T, "docs", Surrogate(1), "hello world greeting") - .unwrap(); - - assert!(!idx.memtable.is_empty()); - assert!(idx.memtable.posting_count() > 0); - } - - #[test] - fn memtable_flush_on_threshold() { - let backend = MemoryBackend::new(); - let idx = FtsIndex { - backend, - bm25_params: Bm25Params::default(), - memtable: Memtable::new(MemtableConfig { - max_postings: 5, - max_terms: 100, - }), - next_segment_id: AtomicU64::new(1), - governor: test_governor(), - }; - - idx.index_document( - DB, - T, - "docs", - Surrogate(1), - "alpha bravo charlie delta echo foxtrot", - ) - .unwrap(); - - assert!(idx.memtable.is_empty()); - let segments = idx.backend.list_segments(DB, T, "docs").unwrap(); - assert!(!segments.is_empty(), "segment should have been written"); - } - - #[test] - fn index_surrogate_stored() { - let idx = make_index(); - // Surrogates must be in 1..=MAX_INDEXABLE_SURROGATE. Surrogate::ZERO is - // the unassigned sentinel and is now rejected at index time. - idx.index_document(DB, T, "docs", Surrogate(10), "hello world greeting") - .unwrap(); - idx.index_document(DB, T, "docs", Surrogate(11), "hello rust language") - .unwrap(); - - let (count, _) = idx.backend.collection_stats(DB, T, "docs").unwrap(); - assert_eq!(count, 2); - } - - #[test] - fn remove_decrements_stats() { - let idx = make_index(); - idx.index_document(DB, T, "docs", Surrogate(10), "hello world") - .unwrap(); - idx.index_document(DB, T, "docs", Surrogate(11), "hello rust") - .unwrap(); - - idx.remove_document(DB, T, "docs", Surrogate(10)).unwrap(); - - let (count, _) = idx.backend.collection_stats(DB, T, "docs").unwrap(); - assert_eq!(count, 1); - } - - #[test] - fn index_updates_stats() { - let idx = make_index(); - idx.index_document(DB, T, "docs", Surrogate(10), "hello world greeting") - .unwrap(); - idx.index_document(DB, T, "docs", Surrogate(11), "hello rust language") - .unwrap(); - - let (count, total) = idx.backend.collection_stats(DB, T, "docs").unwrap(); - assert_eq!(count, 2); - assert!(total > 0); - } - - #[test] - fn purge_collection_preserves_others() { - let idx = make_index(); - idx.index_document(DB, T, "col_a", Surrogate(1), "alpha bravo") - .unwrap(); - idx.index_document(DB, T, "col_b", Surrogate(1), "delta echo") - .unwrap(); - - idx.purge_collection(DB, T, "col_a").unwrap(); - assert_eq!( - idx.backend.collection_stats(DB, T, "col_a").unwrap(), - (0, 0) - ); - assert!(idx.backend.collection_stats(DB, T, "col_b").unwrap().0 > 0); - - assert!( - !idx.memtable - .get_postings(&memtable_key(DB, T, "col_b", "delta")) - .is_empty() - ); - assert!( - idx.memtable - .get_postings(&memtable_key(DB, T, "col_a", "alpha")) - .is_empty() - ); - } - - #[test] - fn empty_text_is_noop() { - let idx = make_index(); - idx.index_document(DB, T, "docs", Surrogate(1), "the a is") - .unwrap(); - assert_eq!(idx.backend.collection_stats(DB, T, "docs").unwrap(), (0, 0)); - assert!(idx.memtable.is_empty()); - } - - // ── Surrogate boundary tests ────────────────────────────────────────────── - - /// Spec: Surrogate::ZERO (the unassigned sentinel) must be rejected at index - /// time with FtsIndexError::SurrogateOutOfRange, not written into the index. - #[test] - fn index_document_rejects_zero_surrogate() { - let idx = make_index(); - let err = idx - .index_document(DB, T, "docs", Surrogate(0), "hello world") - .unwrap_err(); - assert!( - matches!(err, FtsIndexError::SurrogateOutOfRange { surrogate } if surrogate == Surrogate(0)), - "expected SurrogateOutOfRange(sur:0), got {err}" - ); - } - - /// Spec: Surrogate(u32::MAX) must be rejected — it is reserved as a sentinel - /// and would also cause a 4 GiB fieldnorm array resize. - #[test] - fn index_document_rejects_u32_max_surrogate() { - let idx = make_index(); - let err = idx - .index_document(DB, T, "docs", Surrogate(u32::MAX), "hello world") - .unwrap_err(); - assert!( - matches!(err, FtsIndexError::SurrogateOutOfRange { .. }), - "expected SurrogateOutOfRange, got {err}" - ); - } - - /// Spec: MAX_INDEXABLE_SURROGATE (u32::MAX - 1) is the last valid surrogate. - /// Indexing with it must succeed. - #[test] - fn index_document_accepts_max_indexable_surrogate() { - // NOTE: The MemoryBackend fieldnorm array would resize to u32::MAX - 1 - // bytes (~4 GiB) in a real call. We test the boundary using a - // surrogate just below the limit to verify the guard passes without - // actually allocating 4 GiB. The exact boundary (MAX_INDEXABLE_SURROGATE) - // is confirmed by the value check: max_sur > MAX and max_sur is rejected, - // while max_sur == MAX_INDEXABLE_SURROGATE is accepted. - // - // We use Surrogate(1) here and confirm the guard's upper boundary - // by separately verifying Surrogate(u32::MAX) is rejected (above). - let idx = make_index(); - // Verify a normal valid surrogate works (guards pass). - idx.index_document(DB, T, "docs", Surrogate(1), "boundary check") - .unwrap(); - // Confirm the constant is correct. - assert_eq!( - crate::index::error::MAX_INDEXABLE_SURROGATE, - u32::MAX - 1, - "MAX_INDEXABLE_SURROGATE must be u32::MAX - 1" - ); - } - - /// Spec: the SurrogateOutOfRange error message must be informative. - #[test] - fn surrogate_out_of_range_error_is_informative() { - let err: FtsIndexError = - FtsIndexError::SurrogateOutOfRange { - surrogate: Surrogate(0), - }; - let msg = err.to_string(); - assert!( - msg.contains("out of the indexable range"), - "error message must mention range: {msg}" - ); - assert!( - msg.contains("unassigned sentinel"), - "error message must explain zero sentinel: {msg}" - ); - } -} diff --git a/nodedb-fts/src/index/writer/document.rs b/nodedb-fts/src/index/writer/document.rs new file mode 100644 index 000000000..ceef0634c --- /dev/null +++ b/nodedb-fts/src/index/writer/document.rs @@ -0,0 +1,528 @@ +// SPDX-License-Identifier: Apache-2.0 + +//! Document writes reuse ordered analyzer output and preserve surrogate bounds. + +use super::FtsIndex; +use crate::{ + backend::FtsBackend, + block::CompactPosting, + codec::smallfloat, + index::error::{FtsIndexError, MAX_INDEXABLE_SURROGATE}, + lsm::segment_deletes::SegmentDeletes, + scope::IndexScope, +}; +use nodedb_types::Surrogate; +use std::collections::HashMap; +use tracing::debug; + +impl FtsIndex { + /// Index a document's text content into one index. The analyzer is the + /// one bound to the index's collection. + /// + /// Returns `Err(FtsIndexError::SurrogateOutOfRange)` for `Surrogate::ZERO` + /// or values exceeding `MAX_INDEXABLE_SURROGATE`. `Surrogate::ZERO` is the + /// unassigned sentinel. Fieldnorm arrays use raw `u32` surrogates as + /// indexes. Values near `u32::MAX` cause multi-GiB allocations. Runtime + /// bounds checks run before analysis and remain active in release builds. + pub fn index_document<'a>( + &self, + database_id: u64, + tid: u64, + index: impl Into>, + doc_id: Surrogate, + text: &str, + ) -> Result<(), FtsIndexError> { + let index = index.into(); + Self::check_surrogate(doc_id)?; + + let tokens = self + .analyze_for_collection(database_id, tid, index.collection(), text) + .map_err(FtsIndexError::backend)?; + self.index_analyzed_document(database_id, tid, index, doc_id, &tokens) + } + + /// Index ordered tokens from this collection's current analyzer. Callers + /// preserve token order and hold analyzer config stable through indexing. + /// + /// Returns `Err(FtsIndexError::SurrogateOutOfRange)` for `Surrogate::ZERO` + /// or values exceeding `MAX_INDEXABLE_SURROGATE`. + pub fn index_analyzed_document<'a>( + &self, + database_id: u64, + tid: u64, + index: impl Into>, + doc_id: Surrogate, + tokens: &[String], + ) -> Result<(), FtsIndexError> { + let index = index.into(); + Self::check_surrogate(doc_id)?; + if tokens.is_empty() { + return Ok(()); + } + + let mut term_data: HashMap<&str, (u32, Vec)> = HashMap::new(); + for (pos, token) in tokens.iter().enumerate() { + let entry = term_data.entry(token.as_str()).or_insert((0, Vec::new())); + entry.0 += 1; + entry.1.push(pos as u32); + } + + let doc_len = tokens.len() as u32; + let fieldnorm = smallfloat::encode(doc_len); + + let term_count = term_data.len(); + self.memtable.insert_doc( + database_id, + tid, + index, + doc_id, + doc_len, + term_data.into_iter().map(|(term, (freq, positions))| { + ( + term, + CompactPosting { + doc_id, + term_freq: freq, + fieldnorm, + positions, + }, + ) + }), + ); + + // Write document length, fieldnorm, and update incremental stats. + self.backend + .write_doc_length(database_id, tid, index, doc_id, doc_len) + .map_err(FtsIndexError::backend)?; + self.write_fieldnorm(database_id, tid, index, doc_id, doc_len) + .map_err(FtsIndexError::backend)?; + self.backend + .increment_stats(database_id, tid, index, doc_len) + .map_err(FtsIndexError::backend)?; + + if self.memtable.should_flush() { + self.flush_all_memtables()?; + } + + debug!( + database_id, + tid, + collection = index.collection(), + field = index.field_key(), + doc_id = doc_id.0, + tokens = tokens.len(), + terms = term_count, + "indexed document" + ); + Ok(()) + } + + /// Remove a document from one index. + /// + /// The memtable drops the document's postings. Postings already flushed + /// to a segment stay in that immutable segment, so the document enters + /// the delete set of every segment the index holds: reads skip those + /// postings and a merge drops them. The corpus stats lose the document's + /// length. A document the index does not hold changes nothing. + pub fn remove_document<'a>( + &self, + database_id: u64, + tid: u64, + index: impl Into>, + doc_id: Surrogate, + ) -> Result<(), FtsIndexError> { + let index = index.into(); + let doc_len = self + .backend + .read_doc_length(database_id, tid, index, doc_id) + .map_err(FtsIndexError::backend)?; + + self.memtable.remove_doc(database_id, tid, index, doc_id); + + let Some(len) = doc_len else { + return Ok(()); + }; + let segments = self + .backend + .list_segments(database_id, tid, index) + .map_err(FtsIndexError::backend)?; + if !segments.is_empty() { + let mut deletes = + SegmentDeletes::load(&self.backend, database_id, tid, index, &segments)?; + deletes.mark_removed(&segments, doc_id); + deletes.store(&self.backend, database_id, tid, index)?; + } + self.backend + .remove_doc_length(database_id, tid, index, doc_id) + .map_err(FtsIndexError::backend)?; + self.backend + .decrement_stats(database_id, tid, index, len) + .map_err(FtsIndexError::backend)?; + + Ok(()) + } + + fn check_surrogate(doc_id: Surrogate) -> Result<(), FtsIndexError> { + let raw = doc_id.as_u32(); + if raw == 0 || raw > MAX_INDEXABLE_SURROGATE { + return Err(FtsIndexError::SurrogateOutOfRange { surrogate: doc_id }); + } + Ok(()) + } +} + +#[cfg(test)] +mod tests { + use nodedb_types::Surrogate; + + use crate::backend::memory::MemoryBackend; + use crate::test_support::test_governor; + + use super::*; + + const DB: u64 = 0; + const T: u64 = 1; + const DOCS: IndexScope<'static> = IndexScope::document("docs"); + + fn make_index() -> FtsIndex { + FtsIndex::new(MemoryBackend::new(), test_governor()) + } + + #[test] + fn index_writes_to_memtable() { + let idx = make_index(); + idx.index_document(DB, T, "docs", Surrogate(1), "hello world greeting") + .unwrap(); + + assert!(!idx.memtable.is_empty()); + assert!(idx.memtable.posting_count() > 0); + } + + #[test] + fn index_surrogate_stored() { + let idx = make_index(); + // Surrogates must be in 1..=MAX_INDEXABLE_SURROGATE. Surrogate::ZERO is the unassigned sentinel and is rejected at index time. + idx.index_document(DB, T, "docs", Surrogate(10), "hello world greeting") + .unwrap(); + idx.index_document(DB, T, "docs", Surrogate(11), "hello rust language") + .unwrap(); + + let (count, _) = idx.backend.collection_stats(DB, T, DOCS).unwrap(); + assert_eq!(count, 2); + } + + #[test] + fn remove_decrements_stats() { + let idx = make_index(); + idx.index_document(DB, T, "docs", Surrogate(10), "hello world") + .unwrap(); + idx.index_document(DB, T, "docs", Surrogate(11), "hello rust") + .unwrap(); + + idx.remove_document(DB, T, "docs", Surrogate(10)).unwrap(); + + let (count, _) = idx.backend.collection_stats(DB, T, DOCS).unwrap(); + assert_eq!(count, 1); + } + + fn hits(idx: &FtsIndex, query: &str) -> Vec { + let mut ids: Vec = idx + .search( + DB, + T, + DOCS, + crate::FtsSearchParams { + query, + top_k: usize::MAX, + fuzzy_enabled: false, + mode: crate::posting::QueryMode::Or, + prefilter: None, + }, + ) + .unwrap() + .into_iter() + .map(|r| r.doc_id.0) + .collect(); + ids.sort_unstable(); + ids + } + + /// A document removed after its postings reached a segment stops + /// matching, leaves the document frequency and the corpus stats, and + /// leaves the segment at the next merge. Indexed again, it matches its + /// new text only. + #[test] + fn remove_after_flush_hides_segment_postings_and_merge_drops_them() { + use crate::lsm::compaction::{CompactLevelParams, SegmentMeta, compact_level, parse_level}; + use crate::lsm::segment::reader::SegmentReader; + use crate::{DocScore, TextQuery}; + + let idx = make_index(); + idx.index_document(DB, T, "docs", Surrogate(1), "rust tokio") + .unwrap(); + idx.index_document(DB, T, "docs", Surrogate(2), "rust axum") + .unwrap(); + idx.flush_memtable(DB, T, "docs").unwrap(); + idx.index_document(DB, T, "docs", Surrogate(3), "rust serde") + .unwrap(); + idx.flush_memtable(DB, T, "docs").unwrap(); + assert_eq!(hits(&idx, "tokio"), vec![1]); + + idx.remove_document(DB, T, "docs", Surrogate(1)).unwrap(); + assert!(hits(&idx, "tokio").is_empty(), "a removed document matches"); + assert_eq!(hits(&idx, "rust"), vec![2, 3]); + let blocks = idx.term_blocks(DB, T, DOCS, &["rust".into()]).unwrap(); + assert_eq!(blocks[0].df, 2, "document frequency counts live documents"); + assert_eq!(idx.backend.collection_stats(DB, T, DOCS).unwrap(), (2, 4)); + let scorer = idx + .doc_scorer( + DB, + T, + DOCS, + TextQuery { + query: "tokio", + fuzzy_enabled: false, + mode: crate::posting::QueryMode::Or, + }, + None, + None, + ) + .unwrap(); + assert_eq!( + scorer.score(&[Surrogate(1)]).unwrap(), + vec![DocScore::Absent] + ); + + idx.index_document(DB, T, "docs", Surrogate(1), "rust hyper") + .unwrap(); + idx.flush_memtable(DB, T, "docs").unwrap(); + assert_eq!(hits(&idx, "hyper"), vec![1]); + assert!(hits(&idx, "tokio").is_empty()); + assert_eq!(hits(&idx, "rust"), vec![1, 2, 3]); + + let segments: Vec = idx + .backend + .list_segments(DB, T, DOCS) + .unwrap() + .into_iter() + .map(|segment_id| SegmentMeta { + level: parse_level(&segment_id), + segment_id, + size: 0, + }) + .collect(); + assert_eq!(segments.len(), 3); + let governor = test_governor(); + let (merged, merged_ids) = compact_level(CompactLevelParams { + backend: &idx.backend, + database_id: DB, + tid: T, + index: DOCS, + segments: &segments, + level: 0, + governor: &governor, + }) + .unwrap() + .expect("three level-0 segments merge"); + let reader = SegmentReader::open(merged.clone()).unwrap(); + assert!( + reader.find_term("tokio").is_none(), + "the merge drops dead postings" + ); + assert_eq!(reader.df("rust"), 3); + + idx.backend + .write_segment(DB, T, DOCS, "L1:00000000000000ff", &merged) + .unwrap(); + for id in &merged_ids { + idx.backend.remove_segment(DB, T, DOCS, id).unwrap(); + } + assert!(hits(&idx, "tokio").is_empty()); + assert_eq!(hits(&idx, "rust"), vec![1, 2, 3]); + + // The merged segment carries no delete set: a later removal hides it. + idx.remove_document(DB, T, "docs", Surrogate(2)).unwrap(); + assert_eq!(hits(&idx, "rust"), vec![1, 3]); + } + + /// A flush after a restart does not reuse a stored segment id. + #[test] + fn a_flush_never_reuses_a_stored_segment_id() { + let idx = make_index(); + idx.index_document(DB, T, "docs", Surrogate(1), "alpha") + .unwrap(); + idx.flush_memtable(DB, T, "docs").unwrap(); + + let restarted = FtsIndex::new(idx.backend, test_governor()); + restarted + .index_document(DB, T, "docs", Surrogate(2), "bravo") + .unwrap(); + restarted.flush_memtable(DB, T, "docs").unwrap(); + assert_eq!( + restarted.backend.list_segments(DB, T, DOCS).unwrap().len(), + 2 + ); + assert_eq!(hits(&restarted, "alpha"), vec![1]); + assert_eq!(hits(&restarted, "bravo"), vec![2]); + } + + #[test] + fn index_updates_stats() { + let idx = make_index(); + idx.index_document(DB, T, "docs", Surrogate(10), "hello world greeting") + .unwrap(); + idx.index_document(DB, T, "docs", Surrogate(11), "hello rust language") + .unwrap(); + + let (count, total) = idx.backend.collection_stats(DB, T, DOCS).unwrap(); + assert_eq!(count, 2); + assert!(total > 0); + } + + #[test] + fn field_index_keeps_its_own_postings_and_stats() { + let idx = make_index(); + let title = IndexScope::field("docs", "title").unwrap(); + idx.index_document(DB, T, title, Surrogate(1), "rust") + .unwrap(); + idx.index_document(DB, T, "docs", Surrogate(1), "rust handbook") + .unwrap(); + idx.index_document(DB, T, "docs", Surrogate(2), "rust nomicon") + .unwrap(); + + assert_eq!(idx.backend.collection_stats(DB, T, title).unwrap(), (1, 1)); + assert_eq!(idx.backend.collection_stats(DB, T, DOCS).unwrap(), (2, 4)); + assert_eq!(idx.memtable.get_postings(DB, T, title, "rust").len(), 1); + assert_eq!(idx.memtable.get_postings(DB, T, DOCS, "rust").len(), 2); + + idx.remove_document(DB, T, title, Surrogate(1)).unwrap(); + assert_eq!(idx.backend.collection_stats(DB, T, title).unwrap(), (0, 0)); + assert!(idx.memtable.get_postings(DB, T, title, "rust").is_empty()); + assert_eq!(idx.memtable.get_postings(DB, T, DOCS, "rust").len(), 2); + } + + #[test] + fn empty_text_is_noop() { + let idx = make_index(); + idx.index_document(DB, T, "docs", Surrogate(1), "the a is") + .unwrap(); + assert_eq!(idx.backend.collection_stats(DB, T, DOCS).unwrap(), (0, 0)); + assert!(idx.memtable.is_empty()); + } + + // ── Surrogate boundary tests ────────────────────────────────────────────── + + /// Spec: Surrogate::ZERO (the unassigned sentinel) must be rejected at index + /// time with FtsIndexError::SurrogateOutOfRange, not written into the index. + #[test] + fn index_document_rejects_zero_surrogate() { + let idx = make_index(); + let err = idx + .index_document(DB, T, "docs", Surrogate(0), "hello world") + .unwrap_err(); + assert!( + matches!(err, FtsIndexError::SurrogateOutOfRange { surrogate } if surrogate == Surrogate(0)), + "expected SurrogateOutOfRange(sur:0), got {err}" + ); + } + + /// Spec: Surrogate(u32::MAX) must be rejected — it is reserved as a sentinel + /// and would also cause a 4 GiB fieldnorm array resize. + #[test] + fn index_document_rejects_u32_max_surrogate() { + let idx = make_index(); + let err = idx + .index_document(DB, T, "docs", Surrogate(u32::MAX), "hello world") + .unwrap_err(); + assert!( + matches!(err, FtsIndexError::SurrogateOutOfRange { .. }), + "expected SurrogateOutOfRange, got {err}" + ); + } + + /// Check the last valid surrogate constant and a representative valid input. + #[test] + fn index_document_accepts_max_indexable_surrogate() { + // Indexing MAX_INDEXABLE_SURROGATE requires multi-GiB fieldnorm arrays. + // This fixture uses Surrogate(1) and checks the constant and sentinel boundary separately. + let idx = make_index(); + // Check a valid surrogate without allocating the largest fieldnorm array. + idx.index_document(DB, T, "docs", Surrogate(1), "boundary check") + .unwrap(); + // Confirm the constant is correct. + assert_eq!( + crate::index::error::MAX_INDEXABLE_SURROGATE, + u32::MAX - 1, + "MAX_INDEXABLE_SURROGATE must be u32::MAX - 1" + ); + } + + /// Spec: the SurrogateOutOfRange error message must be informative. + #[test] + fn surrogate_out_of_range_error_is_informative() { + let err: FtsIndexError = + FtsIndexError::SurrogateOutOfRange { + surrogate: Surrogate(0), + }; + let msg = err.to_string(); + assert!( + msg.contains("out of the indexable range"), + "error message must mention range: {msg}" + ); + assert!( + msg.contains("unassigned sentinel"), + "error message must explain zero sentinel: {msg}" + ); + } + #[test] + fn analyzed_tokens_match_text_indexing_and_surrogate_validation() { + let text_index = make_index(); + let token_index = make_index(); + let text = "Alpha alpha beta gamma"; + let tokens = token_index + .analyze_for_collection(DB, T, "docs", text) + .unwrap(); + text_index + .index_document(DB, T, "docs", Surrogate(1), text) + .unwrap(); + token_index + .index_analyzed_document(DB, T, "docs", Surrogate(1), &tokens) + .unwrap(); + assert_eq!( + text_index.memtable.stats(DB, T, DOCS), + token_index.memtable.stats(DB, T, DOCS) + ); + let mut terms = text_index.memtable.terms(DB, T, DOCS); + terms.sort(); + for term in terms { + let expected = text_index.memtable.get_postings(DB, T, DOCS, &term); + let actual = token_index.memtable.get_postings(DB, T, DOCS, &term); + assert_eq!(actual.len(), expected.len()); + for (actual, expected) in actual.iter().zip(expected) { + assert_eq!(actual.doc_id, expected.doc_id); + assert_eq!(actual.term_freq, expected.term_freq); + assert_eq!(actual.fieldnorm, expected.fieldnorm); + assert_eq!(actual.positions, expected.positions); + } + } + assert_eq!( + text_index.backend.collection_stats(DB, T, DOCS).unwrap(), + token_index.backend.collection_stats(DB, T, DOCS).unwrap() + ); + for id in [Surrogate::ZERO, Surrogate(u32::MAX)] { + assert!(matches!( + token_index.index_analyzed_document(DB, T, "docs", id, &tokens), + Err(FtsIndexError::SurrogateOutOfRange { .. }) + )); + assert!(matches!( + token_index.index_analyzed_document(DB, T, "docs", id, &[]), + Err(FtsIndexError::SurrogateOutOfRange { .. }) + )); + } + let empty = make_index(); + empty + .index_analyzed_document(DB, T, "docs", Surrogate(1), &[]) + .unwrap(); + assert!(empty.memtable.is_empty()); + } +} diff --git a/nodedb-fts/src/index/writer/maintenance.rs b/nodedb-fts/src/index/writer/maintenance.rs new file mode 100644 index 000000000..44f6351a1 --- /dev/null +++ b/nodedb-fts/src/index/writer/maintenance.rs @@ -0,0 +1,349 @@ +// SPDX-License-Identifier: Apache-2.0 + +//! Segment publication and collection purging. + +use super::FtsIndex; +use crate::{ + backend::FtsBackend, + index::error::FtsIndexError, + lsm::{compaction, segment::writer as seg_writer}, + scope::IndexScope, +}; +use std::sync::atomic::Ordering; +use tracing::debug; + +impl FtsIndex { + /// Flush one index's memtable postings to an immutable segment of that + /// index in the backend. Every other index keeps its memtable state. + /// + /// The postings leave the memtable only after the segment is written, so + /// a failed flush loses nothing. + /// + /// Calling this before serializing the index guarantees that all posting + /// data written since the last spill threshold is captured in the backend's + /// segment storage rather than the in-memory memtable. Callers that + /// checkpoint the index (e.g., NodeDB-Lite flush) must call this once per + /// active index before persisting. + pub fn flush_memtable<'a>( + &self, + database_id: u64, + tid: u64, + index: impl Into>, + ) -> Result<(), FtsIndexError> { + let index = index.into(); + let segment = self + .memtable + .with_scope_postings(database_id, tid, index, |postings| { + (!postings.is_empty()).then(|| seg_writer::flush_postings_to_segment(postings)) + }) + .flatten(); + let Some(segment_bytes) = segment else { + self.memtable.drain_scope(database_id, tid, index); + return Ok(()); + }; + let segment_bytes = segment_bytes?; + + // A new segment id is above every id the index holds. The counter + // restarts with the process, so it is raised past the stored ids: + // reusing an id would overwrite a segment and inherit its delete set. + let floor = self + .backend + .list_segments(database_id, tid, index) + .map_err(FtsIndexError::backend)? + .iter() + .map(|id| compaction::parse_segment_number(id)) + .max() + .map_or(0, |n| n.saturating_add(1)); + self.next_segment_id.fetch_max(floor, Ordering::Relaxed); + let seg_id = self.next_segment_id.fetch_add(1, Ordering::Relaxed); + let id = compaction::segment_id(seg_id, 0); + self.backend + .write_segment(database_id, tid, index, &id, &segment_bytes) + .map_err(FtsIndexError::backend)?; + self.memtable.drain_scope(database_id, tid, index); + + debug!( + database_id, + tid, + collection = index.collection(), + field = index.field_key(), + seg_id, + bytes = segment_bytes.len(), + "flushed memtable to segment" + ); + Ok(()) + } + + /// Flush every index the memtable holds, each to its own segment. + pub fn flush_all_memtables(&self) -> Result<(), FtsIndexError> { + for scope in self.memtable.scopes() { + self.flush_memtable(scope.database_id, scope.tid, scope.index())?; + } + Ok(()) + } + + /// Purge every index of a collection. Returns count of removed entries. + pub fn purge_collection( + &self, + database_id: u64, + tid: u64, + collection: &str, + ) -> Result { + self.memtable.drain_collection(database_id, tid, collection); + self.backend.purge_collection(database_id, tid, collection) + } + + /// Purge all entries for a `(database_id, tenant)` across every collection. + pub fn purge_tenant(&self, database_id: u64, tid: u64) -> Result { + self.memtable.drain_tenant(database_id, tid); + self.backend.purge_tenant(database_id, tid) + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::backend::memory::MemoryBackend; + use crate::test_support::test_governor; + use crate::{ + FtsSearchParams, + block::CompactPosting, + lsm::memtable::{Memtable, MemtableConfig}, + posting::{Bm25Params, QueryMode}, + }; + use nodedb_types::Surrogate; + use std::sync::atomic::AtomicU64; + const DB: u64 = 0; + const T: u64 = 1; + const DOCS: IndexScope<'static> = IndexScope::document("docs"); + fn make_index() -> FtsIndex { + FtsIndex::new(MemoryBackend::new(), test_governor()) + } + + fn hits(idx: &FtsIndex, index: IndexScope<'_>, query: &str) -> Vec { + let mut ids: Vec = idx + .search( + DB, + T, + index, + FtsSearchParams { + query, + top_k: usize::MAX, + fuzzy_enabled: false, + mode: QueryMode::Or, + prefilter: None, + }, + ) + .unwrap() + .into_iter() + .map(|r| r.doc_id.0) + .collect(); + ids.sort_unstable(); + ids + } + + #[test] + fn flush_propagates_term_too_long_as_typed_error() { + let backend = MemoryBackend::new(); + let idx = FtsIndex { + backend, + bm25_params: Bm25Params::default(), + memtable: Memtable::new(MemtableConfig { + max_postings: 1, + max_terms: 1, + }), + next_segment_id: AtomicU64::new(1), + governor: test_governor(), + }; + + // Insert a single posting under a term whose byte length exceeds the u16 segment-format cap. Bypasses the analyzer (which would tokenize away most pathological inputs); we want to exercise the flush-path boundary check directly. + let oversize_term = "x".repeat(crate::lsm::segment::format::MAX_TERM_LEN + 1); + idx.memtable.insert( + DB, + T, + DOCS, + &oversize_term, + CompactPosting { + doc_id: Surrogate(1), + term_freq: 1, + fieldnorm: 1, + positions: vec![0], + }, + ); + idx.memtable.record_doc(DB, T, DOCS, Surrogate(1), 1); + + let err = idx + .flush_memtable(DB, T, "docs") + .expect_err("flush must reject oversize term"); + match err { + FtsIndexError::TermTooLong { len, max } => { + assert_eq!(len, oversize_term.len()); + assert_eq!(max, crate::lsm::segment::format::MAX_TERM_LEN); + } + other => panic!("expected TermTooLong, got {other:?}"), + } + assert_eq!( + idx.memtable.get_postings(DB, T, DOCS, &oversize_term).len(), + 1, + "a failed flush keeps the postings in the memtable" + ); + } + #[test] + fn memtable_flush_on_threshold() { + let backend = MemoryBackend::new(); + let idx = FtsIndex { + backend, + bm25_params: Bm25Params::default(), + memtable: Memtable::new(MemtableConfig { + max_postings: 5, + max_terms: 100, + }), + next_segment_id: AtomicU64::new(1), + governor: test_governor(), + }; + + idx.index_document( + DB, + T, + "docs", + Surrogate(1), + "alpha bravo charlie delta echo foxtrot", + ) + .unwrap(); + + assert!(idx.memtable.is_empty()); + let segments = idx.backend.list_segments(DB, T, DOCS).unwrap(); + assert!(!segments.is_empty(), "segment should have been written"); + assert_eq!(hits(&idx, DOCS, "charlie"), vec![1]); + } + #[test] + fn purge_collection_preserves_others() { + let idx = make_index(); + let a = IndexScope::document("col_a"); + let a_title = IndexScope::field("col_a", "title").unwrap(); + let b = IndexScope::document("col_b"); + idx.index_document(DB, T, a, Surrogate(1), "alpha bravo") + .unwrap(); + idx.index_document(DB, T, a_title, Surrogate(1), "alpha") + .unwrap(); + idx.index_document(DB, T, b, Surrogate(1), "delta echo") + .unwrap(); + + idx.purge_collection(DB, T, "col_a").unwrap(); + assert_eq!(idx.backend.collection_stats(DB, T, a).unwrap(), (0, 0)); + assert_eq!( + idx.backend.collection_stats(DB, T, a_title).unwrap(), + (0, 0) + ); + assert!(idx.backend.collection_stats(DB, T, b).unwrap().0 > 0); + + assert!(!idx.memtable.get_postings(DB, T, b, "delta").is_empty()); + assert!(idx.memtable.get_postings(DB, T, a, "alpha").is_empty()); + assert!( + idx.memtable + .get_postings(DB, T, a_title, "alpha") + .is_empty() + ); + } + + /// Flushing one collection's index writes only that index's postings + /// to its segment. The other collection keeps its memtable postings and + /// stats, and both stay searchable. + #[test] + fn flush_of_one_collection_leaves_the_other_intact() { + let idx = make_index(); + let a = IndexScope::document("col_a"); + let b = IndexScope::document("col_b"); + idx.index_document(DB, T, a, Surrogate(1), "alpha bravo") + .unwrap(); + idx.index_document(DB, T, b, Surrogate(2), "alpha charlie") + .unwrap(); + let b_stats = idx.memtable.stats(DB, T, b); + + idx.flush_memtable(DB, T, a).unwrap(); + + assert!(idx.memtable.terms(DB, T, a).is_empty()); + assert_eq!(idx.backend.list_segments(DB, T, a).unwrap().len(), 1); + assert!(idx.backend.list_segments(DB, T, b).unwrap().is_empty()); + assert_eq!(idx.memtable.get_postings(DB, T, b, "alpha").len(), 1); + assert_eq!(idx.memtable.stats(DB, T, b), b_stats); + assert_eq!(b_stats, (1, 2)); + + assert_eq!(hits(&idx, a, "alpha"), vec![1]); + assert_eq!(hits(&idx, b, "alpha"), vec![2]); + assert!(hits(&idx, a, "charlie").is_empty()); + assert_eq!(idx.index_stats(DB, T, a).unwrap().0, 1); + assert_eq!(idx.index_stats(DB, T, b).unwrap().0, 1); + } + + /// Stats of a field index count only the documents that hold that field. + #[test] + fn stats_are_per_index() { + let idx = make_index(); + let title = IndexScope::field("docs", "title").unwrap(); + idx.index_document(DB, T, DOCS, Surrogate(1), "rust book") + .unwrap(); + idx.index_document(DB, T, DOCS, Surrogate(2), "java book") + .unwrap(); + idx.index_document(DB, T, title, Surrogate(1), "rust") + .unwrap(); + + assert_eq!(idx.index_stats(DB, T, DOCS).unwrap(), (2, 2.0)); + assert_eq!(idx.index_stats(DB, T, title).unwrap(), (1, 1.0)); + assert_eq!(idx.memtable.stats(DB, T, DOCS), (2, 4)); + assert_eq!(idx.memtable.stats(DB, T, title), (1, 1)); + } + + /// A low spill threshold flushes every index to its own segment, and + /// every index stays searchable. + #[test] + fn low_threshold_spills_each_index_to_its_own_segment() { + let idx = FtsIndex::with_memtable_config( + MemoryBackend::new(), + MemtableConfig { + max_postings: 4, + max_terms: 100, + }, + test_governor(), + ); + let title = IndexScope::field("docs", "title").unwrap(); + let other = IndexScope::document("other"); + idx.index_document(DB, T, title, Surrogate(1), "rust") + .unwrap(); + idx.index_document(DB, T, other, Surrogate(2), "rust") + .unwrap(); + idx.index_document(DB, T, DOCS, Surrogate(1), "rust handbook") + .unwrap(); + + assert!(idx.memtable.is_empty()); + for index in [title, other, DOCS] { + assert_eq!(idx.backend.list_segments(DB, T, index).unwrap().len(), 1); + } + assert_eq!(hits(&idx, title, "rust"), vec![1]); + assert_eq!(hits(&idx, other, "rust"), vec![2]); + assert_eq!(hits(&idx, DOCS, "handbook"), vec![1]); + assert!(hits(&idx, title, "handbook").is_empty()); + } + + #[test] + fn with_config_sets_bm25_params_and_memtable_thresholds() { + let params = Bm25Params { k1: 2.0, b: 0.5 }; + let idx = FtsIndex::with_config( + MemoryBackend::new(), + params, + MemtableConfig { + max_postings: 1, + max_terms: 100, + }, + test_governor(), + ); + assert_eq!(idx.bm25_params.k1, 2.0); + assert_eq!(idx.bm25_params.b, 0.5); + + idx.index_document(DB, T, DOCS, Surrogate(1), "alpha") + .unwrap(); + assert!(idx.memtable.is_empty()); + assert_eq!(idx.backend.list_segments(DB, T, DOCS).unwrap().len(), 1); + assert_eq!(hits(&idx, DOCS, "alpha"), vec![1]); + } +} diff --git a/nodedb-fts/src/index/writer/mod.rs b/nodedb-fts/src/index/writer/mod.rs new file mode 100644 index 000000000..1076f8380 --- /dev/null +++ b/nodedb-fts/src/index/writer/mod.rs @@ -0,0 +1,9 @@ +// SPDX-License-Identifier: Apache-2.0 + +//! Index writer wiring. + +mod document; +mod maintenance; +mod state; + +pub use state::FtsIndex; diff --git a/nodedb-fts/src/index/writer/state.rs b/nodedb-fts/src/index/writer/state.rs new file mode 100644 index 000000000..091a4d95f --- /dev/null +++ b/nodedb-fts/src/index/writer/state.rs @@ -0,0 +1,97 @@ +// SPDX-License-Identifier: Apache-2.0 + +//! Index state, construction, and backend access. + +use crate::{ + backend::FtsBackend, + lsm::memtable::{Memtable, MemtableConfig}, + posting::Bm25Params, +}; +use nodedb_mem::MemoryGovernor; +use std::sync::{Arc, atomic::AtomicU64}; + +/// Full-text search index generic over storage backend. +/// +/// Indexing, search, and highlighting use the [`FtsBackend`] storage contract. +/// +/// Writes accumulate in an in-memory `Memtable`, one state per index. +/// Threshold-triggered flushing writes each index's postings to an +/// immutable segment of that index through the backend. +/// Queries merge the active memtable with stored segments. +/// The backend determines segment durability. +/// +/// [`MemoryGovernor`] enforces per-engine memory budgets on large +/// allocations (compaction, segment merge, query term collection). +pub struct FtsIndex { + pub(crate) backend: B, + pub(crate) bm25_params: Bm25Params, + pub(crate) memtable: Memtable, + /// Monotonic segment ID counter. + pub(super) next_segment_id: AtomicU64, + /// Memory governor for budget enforcement. + pub(crate) governor: Arc, +} + +impl FtsIndex { + /// Create a new FTS index with the given backend, default BM25 params, and a memory governor. + pub fn new(backend: B, governor: Arc) -> Self { + Self { + backend, + bm25_params: Bm25Params::default(), + memtable: Memtable::new(MemtableConfig::default()), + next_segment_id: AtomicU64::new(1), + governor, + } + } + + /// Create a new FTS index with custom BM25 parameters and a memory governor. + pub fn with_params(backend: B, params: Bm25Params, governor: Arc) -> Self { + Self { + bm25_params: params, + ..Self::new(backend, governor) + } + } + + /// Create a new FTS index whose memtable spills at `memtable` thresholds. + /// Unreachable thresholds keep every posting in the memtable. + pub fn with_memtable_config( + backend: B, + memtable: MemtableConfig, + governor: Arc, + ) -> Self { + Self { + memtable: Memtable::new(memtable), + ..Self::new(backend, governor) + } + } + + /// Create a new FTS index with custom BM25 parameters whose memtable + /// spills at `memtable` thresholds. + pub fn with_config( + backend: B, + params: Bm25Params, + memtable: MemtableConfig, + governor: Arc, + ) -> Self { + Self { + bm25_params: params, + memtable: Memtable::new(memtable), + ..Self::new(backend, governor) + } + } + + /// Access the underlying backend. + pub fn backend(&self) -> &B { + &self.backend + } + + /// Mutable access to the underlying backend. + pub fn backend_mut(&mut self) -> &mut B { + &mut self.backend + } + + /// Access the active memtable (for LSM query merging). + pub fn memtable(&self) -> &Memtable { + &self.memtable + } +} diff --git a/nodedb-fts/src/lib.rs b/nodedb-fts/src/lib.rs index b78ded283..f8b41b336 100644 --- a/nodedb-fts/src/lib.rs +++ b/nodedb-fts/src/lib.rs @@ -17,12 +17,14 @@ pub mod backend; pub mod block; pub mod bm25; pub mod codec; +pub mod document_text; pub mod fuzzy; pub mod highlight; pub mod index; pub mod lsm; mod mem_scope; pub mod posting; +pub mod scope; pub mod search; #[cfg(test)] mod test_support; @@ -33,9 +35,14 @@ pub use analyzer::{ }; pub use backend::FtsBackend; pub use block::{CompactPosting, PostingBlock}; +pub use document_text::DocumentText; pub use fuzzy::{fuzzy_discount, fuzzy_match, levenshtein, max_distance_for_length}; pub use index::{FtsIndex, FtsIndexError, MAX_INDEXABLE_SURROGATE, SynonymGroupRecord}; pub use nodedb_types::Surrogate; pub use posting::{Bm25Params, MatchOffset, Posting, QueryMode, TextSearchResult}; +pub use scope::IndexScope; pub use search::bm25_search::FtsSearchParams; +pub use search::doc_scorer::{DocScore, DocScorer}; pub use search::query_parser::{InvalidQuery, ParsedQuery}; +pub use search::query_terms::TextQuery; +pub use search::staged::{StagedDoc, StagedView}; diff --git a/nodedb-fts/src/lsm/compaction.rs b/nodedb-fts/src/lsm/compaction.rs index 9f6ea9082..8f9d4b73a 100644 --- a/nodedb-fts/src/lsm/compaction.rs +++ b/nodedb-fts/src/lsm/compaction.rs @@ -7,9 +7,13 @@ //! that level are merged into a single segment at the next level. use crate::backend::FtsBackend; +use crate::index::FtsIndexError; +use crate::scope::IndexScope; use super::merge; +use super::query::LiveSegment; use super::segment::{reader::SegmentReader, writer}; +use super::segment_deletes::SegmentDeletes; use std::sync::Arc; @@ -67,13 +71,17 @@ pub fn needs_compaction(segments: &[SegmentMeta], config: &CompactionConfig) -> /// Result of a compaction: new segment bytes and ids of merged (to-remove) segments. pub type CompactionResult = (Vec, Vec); -/// Errors from `compact_level` — wraps the backend error and budget exhaustion. +/// Errors from `compact_level` — wraps the backend error, budget exhaustion, +/// and index state that cannot be read or written. #[derive(Debug)] -pub enum CompactError { +pub enum CompactError { /// Underlying backend storage error. Backend(E), /// Memory budget exhausted. Budget(nodedb_mem::MemError), + /// A source segment or the index's delete sets are corrupt or missing, + /// or the merged segment cannot be encoded. No segment is replaced. + Index(FtsIndexError), } impl std::fmt::Display for CompactError { @@ -81,20 +89,24 @@ impl std::fmt::Display for CompactError { match self { CompactError::Backend(e) => write!(f, "compaction backend error: {e}"), CompactError::Budget(e) => write!(f, "compaction budget exhausted: {e}"), + CompactError::Index(e) => write!(f, "compaction index error: {e}"), } } } -/// Helper for converting a `CompactError` to `Backend` variant. -impl CompactError { - pub(crate) fn backend(e: E) -> Self { - CompactError::Backend(e) +impl From> for CompactError { + fn from(e: FtsIndexError) -> Self { + match e { + FtsIndexError::Backend(inner) => CompactError::Backend(inner), + FtsIndexError::BudgetExhausted(inner) => CompactError::Budget(inner), + other => CompactError::Index(other), + } } } /// Inputs to [`compact_level`]. /// -/// Groups the backend handle, the `(database_id, tid, collection)` scope, the +/// Groups the backend handle, the `(database_id, tid, index)` scope, the /// candidate segment list, the target level, and the optional memory governor. pub struct CompactLevelParams<'a, B: FtsBackend> { /// Backend the source segments are read from. @@ -103,9 +115,9 @@ pub struct CompactLevelParams<'a, B: FtsBackend> { pub database_id: u64, /// Owning tenant id. pub tid: u64, - /// Collection whose segments are being compacted. - pub collection: &'a str, - /// All known segments for the collection (filtered to `level` internally). + /// Index whose segments are being compacted. + pub index: IndexScope<'a>, + /// All known segments of the index (filtered to `level` internally). pub segments: &'a [SegmentMeta], /// Level whose segments are merged into `level + 1`. pub level: u32, @@ -117,6 +129,13 @@ pub struct CompactLevelParams<'a, B: FtsBackend> { /// /// Returns the merged segment bytes and the ids of segments that were merged /// (which should be removed from storage after the new segment is written). +/// The merge drops each source segment's postings of its deleted documents, +/// so the merged segment carries no delete set. Once a source segment is +/// removed, reads ignore its delete set and the next delete drops it. +/// +/// A source segment that is missing or fails validation fails the +/// compaction: merging without it and then removing it would lose its +/// postings. /// /// Each `Vec::with_capacity` allocation is budgeted via /// [`nodedb_mem::ScopedMemory::reserve`]. If the budget is exhausted the @@ -128,7 +147,7 @@ pub fn compact_level( backend, database_id, tid, - collection, + index, segments, level, governor, @@ -141,33 +160,44 @@ pub fn compact_level( let memory = fts_scope(governor, database_id, tid); let _readers_guard = memory - .reserve(to_merge.len() * size_of::()) + .reserve(to_merge.len() * size_of::()) .map_err(CompactError::Budget)?; - let mut readers = Vec::with_capacity(to_merge.len()); + let mut sources = Vec::with_capacity(to_merge.len()); let _ids_guard = memory .reserve(to_merge.len() * size_of::()) .map_err(CompactError::Budget)?; let mut merged_ids = Vec::with_capacity(to_merge.len()); - for meta in &to_merge { - if let Some(data) = backend - .read_segment(database_id, tid, collection, &meta.segment_id) - .map_err(CompactError::backend)? - && let Ok(reader) = SegmentReader::open(data) - { - readers.push(reader); - merged_ids.push(meta.segment_id.clone()); - } - } - - if readers.len() < 2 { - return Ok(None); + let merge_ids: Vec = to_merge.iter().map(|m| m.segment_id.clone()).collect(); + let mut deletes = SegmentDeletes::load(backend, database_id, tid, index, &merge_ids)?; + for segment_id in merge_ids { + let Some(data) = backend + .read_segment(database_id, tid, index, &segment_id) + .map_err(CompactError::Backend)? + else { + return Err(CompactError::Index(FtsIndexError::MissingSegment { + segment_id, + })); + }; + let reader = SegmentReader::open(data).map_err(|source| { + CompactError::Index(FtsIndexError::CorruptSegment { + segment_id: segment_id.clone(), + source, + }) + })?; + let deleted = deletes.take(&segment_id); + merged_ids.push(segment_id.clone()); + sources.push(LiveSegment { + segment_id, + reader, + deleted, + }); } - let merged_term_blocks = merge::merge_segments(&readers, &memory); + let merged_term_blocks = merge::merge_segments::(&sources, &memory)?; let new_segment = writer::build_from_blocks(&merged_term_blocks) - .expect("compaction produced a term longer than u16::MAX — data invariant violated"); + .map_err(|e| CompactError::Index(FtsIndexError::from(e)))?; Ok(Some((new_segment, merged_ids))) } diff --git a/nodedb-fts/src/lsm/memtable.rs b/nodedb-fts/src/lsm/memtable.rs deleted file mode 100644 index 2ba163cc7..000000000 --- a/nodedb-fts/src/lsm/memtable.rs +++ /dev/null @@ -1,386 +0,0 @@ -// SPDX-License-Identifier: Apache-2.0 - -//! In-memory memtable for the FTS LSM engine. -//! -//! Accumulates postings in a HashMap until a spill threshold is reached. -//! Serves queries from memory. Thread-safe via interior mutability. - -use std::cell::RefCell; -use std::collections::HashMap; - -use nodedb_types::Surrogate; - -use crate::block::CompactPosting; -use crate::codec::smallfloat; - -/// Spill threshold: flush memtable when total posting entries exceed this. -pub const DEFAULT_SPILL_POSTINGS: usize = 32 * 1024 * 1024; // 32M entries - -/// Spill threshold: flush when unique terms exceed this. -pub const DEFAULT_SPILL_TERMS: usize = 100_000; - -/// Configuration for memtable spill thresholds. -#[derive(Debug, Clone, Copy)] -pub struct MemtableConfig { - pub max_postings: usize, - pub max_terms: usize, -} - -impl Default for MemtableConfig { - fn default() -> Self { - Self { - max_postings: DEFAULT_SPILL_POSTINGS, - max_terms: DEFAULT_SPILL_TERMS, - } - } -} - -/// In-memory accumulator for FTS postings. -/// -/// Stores per-term posting lists keyed on `Surrogate` row identities. -/// Maintains incremental corpus stats. -pub struct Memtable { - /// term → sorted list of compact postings. - postings: RefCell>>, - /// Total number of posting entries across all terms. - total_postings: RefCell, - /// Incremental stats: (doc_count, total_token_sum). - stats: RefCell<(u32, u64)>, - /// Fieldnorm array: surrogate.0 → SmallFloat-encoded length (used by BM25). - fieldnorms: RefCell>, - /// Exact-length sidecar: surrogate.0 → original u32 token count. Used to - /// decrement `stats.total_token_sum` symmetrically with the exact - /// increment in `record_doc`. Storing the original length here keeps - /// the smallfloat array lossy (space-efficient for ranking) while - /// making corpus-stat accounting exact. - fieldnorms_exact: RefCell>, - /// Spill configuration. - config: MemtableConfig, -} - -impl Memtable { - pub fn new(config: MemtableConfig) -> Self { - Self { - postings: RefCell::new(HashMap::new()), - total_postings: RefCell::new(0), - stats: RefCell::new((0, 0)), - fieldnorms: RefCell::new(Vec::new()), - fieldnorms_exact: RefCell::new(Vec::new()), - config, - } - } - - /// Insert a posting for a term. Surrogate is supplied by the caller. - pub fn insert(&self, term: &str, posting: CompactPosting) { - let mut map = self.postings.borrow_mut(); - map.entry(term.to_string()).or_default().push(posting); - *self.total_postings.borrow_mut() += 1; - } - - /// Record a document's stats (call once per indexed document). - pub fn record_doc(&self, doc_id: Surrogate, doc_len: u32) { - let mut stats = self.stats.borrow_mut(); - stats.0 += 1; - stats.1 += doc_len as u64; - - let idx = doc_id.0 as usize; - { - let mut norms = self.fieldnorms.borrow_mut(); - if idx >= norms.len() { - norms.resize(idx + 1, 0); - } - norms[idx] = smallfloat::encode(doc_len); - } - let mut exact = self.fieldnorms_exact.borrow_mut(); - if idx >= exact.len() { - exact.resize(idx + 1, 0); - } - exact[idx] = doc_len; - } - - /// Remove a document's postings from all terms. - pub fn remove_doc(&self, doc_id: Surrogate) { - let mut map = self.postings.borrow_mut(); - let mut removed = 0usize; - map.retain(|_, postings| { - let before = postings.len(); - postings.retain(|p| p.doc_id != doc_id); - removed += before - postings.len(); - !postings.is_empty() - }); - *self.total_postings.borrow_mut() -= removed; - - // Decrement stats using the exact-length sidecar so the subtraction - // matches the exact-length increment in record_doc. Using the - // smallfloat-decoded length here would drift upward on churn. - let mut exact = self.fieldnorms_exact.borrow_mut(); - if let Some(slot) = exact.get_mut(doc_id.0 as usize) - && *slot > 0 - { - let doc_len = *slot; - *slot = 0; - let mut stats = self.stats.borrow_mut(); - stats.0 = stats.0.saturating_sub(1); - stats.1 = stats.1.saturating_sub(doc_len as u64); - } - } - - /// Check if the memtable should be flushed (spill threshold reached). - pub fn should_flush(&self) -> bool { - let tp = *self.total_postings.borrow(); - let terms = self.postings.borrow().len(); - tp >= self.config.max_postings || terms >= self.config.max_terms - } - - /// Get the posting list for a term. Returns empty vec if not found. - pub fn get_postings(&self, term: &str) -> Vec { - self.postings - .borrow() - .get(term) - .cloned() - .unwrap_or_default() - } - - /// Get all term names in the memtable. - pub fn terms(&self) -> Vec { - self.postings.borrow().keys().cloned().collect() - } - - /// Get corpus stats: (doc_count, total_token_sum). - pub fn stats(&self) -> (u32, u64) { - *self.stats.borrow() - } - - /// Get the fieldnorm array (SmallFloat-encoded doc lengths). - pub fn fieldnorms(&self) -> Vec { - self.fieldnorms.borrow().clone() - } - - /// Drain all postings from the memtable (for flush). - /// Returns the term→postings map and resets the memtable to empty, - /// including stats and fieldnorms. - pub fn drain(&self) -> HashMap> { - let mut map = self.postings.borrow_mut(); - *self.total_postings.borrow_mut() = 0; - *self.stats.borrow_mut() = (0, 0); - self.fieldnorms.borrow_mut().clear(); - self.fieldnorms_exact.borrow_mut().clear(); - std::mem::take(&mut *map) - } - - /// Drain only postings matching a key prefix. - /// The caller is responsible for providing the full prefix (including - /// any trailing separator). Resets stats/fieldnorms. - pub fn drain_collection(&self, prefix: &str) { - let mut map = self.postings.borrow_mut(); - let mut removed = 0usize; - map.retain(|k, v| { - if k.starts_with(prefix) { - removed += v.len(); - false - } else { - true - } - }); - *self.total_postings.borrow_mut() -= removed; - // Stats and fieldnorms are collection-scoped in the backend, - // but the memtable tracks them globally. Reset to be safe. - *self.stats.borrow_mut() = (0, 0); - self.fieldnorms.borrow_mut().clear(); - self.fieldnorms_exact.borrow_mut().clear(); - } - - /// Number of unique terms. - pub fn term_count(&self) -> usize { - self.postings.borrow().len() - } - - /// Total posting entries. - pub fn posting_count(&self) -> usize { - *self.total_postings.borrow() - } - - /// Whether the memtable is empty. - pub fn is_empty(&self) -> bool { - self.postings.borrow().is_empty() - } -} - -#[cfg(test)] -mod tests { - use super::*; - - fn make_posting(doc_id: u32, tf: u32) -> CompactPosting { - CompactPosting { - doc_id: Surrogate(doc_id), - term_freq: tf, - fieldnorm: smallfloat::encode(100), - positions: vec![0], - } - } - - #[test] - fn insert_and_query() { - let mt = Memtable::new(MemtableConfig::default()); - mt.insert("hello", make_posting(0, 1)); - mt.insert("hello", make_posting(1, 2)); - mt.insert("world", make_posting(0, 1)); - - assert_eq!(mt.get_postings("hello").len(), 2); - assert_eq!(mt.get_postings("world").len(), 1); - assert!(mt.get_postings("missing").is_empty()); - assert_eq!(mt.term_count(), 2); - assert_eq!(mt.posting_count(), 3); - } - - #[test] - fn remove_doc() { - let mt = Memtable::new(MemtableConfig::default()); - mt.insert("hello", make_posting(0, 1)); - mt.insert("hello", make_posting(1, 2)); - mt.record_doc(Surrogate(0), 100); - mt.record_doc(Surrogate(1), 50); - - mt.remove_doc(Surrogate(0)); - - assert_eq!(mt.get_postings("hello").len(), 1); - assert_eq!(mt.get_postings("hello")[0].doc_id, Surrogate(1)); - assert_eq!(mt.stats().0, 1); // Only doc 1 remains. - } - - #[test] - fn drain_resets_everything() { - let mt = Memtable::new(MemtableConfig::default()); - mt.insert("hello", make_posting(0, 1)); - mt.insert("world", make_posting(1, 1)); - mt.record_doc(Surrogate(0), 100); - mt.record_doc(Surrogate(1), 50); - - let drained = mt.drain(); - assert_eq!(drained.len(), 2); - assert!(mt.is_empty()); - assert_eq!(mt.posting_count(), 0); - assert_eq!(mt.stats(), (0, 0)); - assert!(mt.fieldnorms().is_empty()); - } - - #[test] - fn drain_collection_selective() { - let mt = Memtable::new(MemtableConfig::default()); - mt.insert("col_a:hello", make_posting(0, 1)); - mt.insert("col_a:world", make_posting(1, 1)); - mt.insert("col_b:rust", make_posting(2, 1)); - - mt.drain_collection("col_a:"); - - assert!(mt.get_postings("col_a:hello").is_empty()); - assert!(mt.get_postings("col_a:world").is_empty()); - assert_eq!(mt.get_postings("col_b:rust").len(), 1); - assert_eq!(mt.posting_count(), 1); - } - - #[test] - fn spill_threshold() { - let config = MemtableConfig { - max_postings: 5, - max_terms: 100, - }; - let mt = Memtable::new(config); - for i in 0..4 { - mt.insert("term", make_posting(i, 1)); - } - assert!(!mt.should_flush()); - mt.insert("term", make_posting(4, 1)); - assert!(mt.should_flush()); - } - - #[test] - fn stats_invariant_under_insert_delete_churn() { - // Spec: after any sequence of (record_doc, remove_doc) calls, the - // memtable's corpus stats — specifically `total_token_sum` that feeds - // BM25 `avg_doc_len` — must equal the stats of a memtable freshly - // populated with only the currently-live documents. - // - // Today `remove_doc` subtracts `smallfloat::decode(encoded_len)` from - // the exact-length sum. Because `decode(encode(L)) <= L`, every churn - // cycle leaves positive residue in `total_token_sum`. After N cycles - // of the same doc, the residue compounds and silently skews BM25. - let churn = Memtable::new(MemtableConfig::default()); - let doc_id = Surrogate(0); - let doc_len = 137u32; // not smallfloat-representable exactly - for _ in 0..500 { - churn.record_doc(doc_id, doc_len); - churn.remove_doc(doc_id); - } - // After equal numbers of record / remove, the memtable should be empty - // from a stats standpoint: zero docs, zero tokens. - let fresh = Memtable::new(MemtableConfig::default()); - assert_eq!( - churn.stats(), - fresh.stats(), - "stats drifted after insert/delete churn: churn={:?} fresh={:?}", - churn.stats(), - fresh.stats() - ); - } - - #[test] - fn stats_invariant_churn_then_single_live_doc() { - // Spec: churn on one doc, then leave a single live doc. Stats must - // exactly match a memtable that only ever saw that single live doc. - let churn = Memtable::new(MemtableConfig::default()); - for _ in 0..200 { - churn.record_doc(Surrogate(0), 400); - churn.remove_doc(Surrogate(0)); - } - churn.record_doc(Surrogate(0), 250); - - let fresh = Memtable::new(MemtableConfig::default()); - fresh.record_doc(Surrogate(0), 250); - - assert_eq!( - churn.stats(), - fresh.stats(), - "stats diverged from fresh baseline: churn={:?} fresh={:?}", - churn.stats(), - fresh.stats() - ); - } - - #[test] - fn stats_total_token_sum_never_exceeds_live_doc_length() { - // Regression guard on the specific failure mode: `total_token_sum` - // growing beyond the sum of live document lengths. If this fails, - // BM25 `avg_doc_len` is inflated and ranking is silently skewed. - let mt = Memtable::new(MemtableConfig::default()); - for _ in 0..100 { - mt.record_doc(Surrogate(0), 777); - mt.remove_doc(Surrogate(0)); - } - mt.record_doc(Surrogate(0), 777); - let (count, total) = mt.stats(); - assert_eq!(count, 1); - assert_eq!( - total, 777, - "total_token_sum drifted: got {total}, expected 777 for a single 777-token live doc" - ); - } - - #[test] - fn fieldnorms_recorded() { - let mt = Memtable::new(MemtableConfig::default()); - mt.record_doc(Surrogate(0), 100); - mt.record_doc(Surrogate(5), 50); - - let norms = mt.fieldnorms(); - assert_eq!(norms.len(), 6); // 0..=5 - assert_eq!( - smallfloat::decode(norms[0]), - smallfloat::decode(smallfloat::encode(100)) - ); - assert_eq!( - smallfloat::decode(norms[5]), - smallfloat::decode(smallfloat::encode(50)) - ); - } -} diff --git a/nodedb-fts/src/lsm/memtable/config.rs b/nodedb-fts/src/lsm/memtable/config.rs new file mode 100644 index 000000000..7137e1543 --- /dev/null +++ b/nodedb-fts/src/lsm/memtable/config.rs @@ -0,0 +1,26 @@ +// SPDX-License-Identifier: Apache-2.0 + +//! Memtable spill thresholds. + +/// Spill threshold: flush memtable when total posting entries exceed this. +pub const DEFAULT_SPILL_POSTINGS: usize = 32 * 1024 * 1024; // 32M entries + +/// Spill threshold: flush when unique terms exceed this. +pub const DEFAULT_SPILL_TERMS: usize = 100_000; + +/// Configuration for memtable spill thresholds. Both bound the whole +/// memtable: postings and terms are summed across every index it holds. +#[derive(Debug, Clone, Copy)] +pub struct MemtableConfig { + pub max_postings: usize, + pub max_terms: usize, +} + +impl Default for MemtableConfig { + fn default() -> Self { + Self { + max_postings: DEFAULT_SPILL_POSTINGS, + max_terms: DEFAULT_SPILL_TERMS, + } + } +} diff --git a/nodedb-fts/src/lsm/memtable/mod.rs b/nodedb-fts/src/lsm/memtable/mod.rs new file mode 100644 index 000000000..944b137b0 --- /dev/null +++ b/nodedb-fts/src/lsm/memtable/mod.rs @@ -0,0 +1,8 @@ +// SPDX-License-Identifier: Apache-2.0 + +pub mod config; +pub mod scope_state; +pub mod table; + +pub use config::{DEFAULT_SPILL_POSTINGS, DEFAULT_SPILL_TERMS, MemtableConfig}; +pub use table::{Memtable, MemtableScope}; diff --git a/nodedb-fts/src/lsm/memtable/scope_state.rs b/nodedb-fts/src/lsm/memtable/scope_state.rs new file mode 100644 index 000000000..040b1a7c5 --- /dev/null +++ b/nodedb-fts/src/lsm/memtable/scope_state.rs @@ -0,0 +1,117 @@ +// SPDX-License-Identifier: Apache-2.0 + +//! Memtable state of one index: its postings and its corpus stats. + +use std::collections::HashMap; + +use nodedb_types::Surrogate; + +use crate::block::CompactPosting; +use crate::codec::smallfloat; + +/// Postings, stats, and fieldnorms of one `(database, tenant, index)`. +#[derive(Default)] +pub(crate) struct ScopeState { + /// term → list of compact postings. + pub(super) postings: HashMap>, + /// Posting entries across every term of this index. + pub(super) total_postings: usize, + /// Incremental stats: (doc_count, total_token_sum). + pub(super) stats: (u32, u64), + /// Fieldnorm array: surrogate.0 → SmallFloat-encoded length (used by BM25). + pub(super) fieldnorms: Vec, + /// Exact-length sidecar: surrogate.0 → original u32 token count. Used to + /// decrement `stats.total_token_sum` symmetrically with the exact + /// increment in `record_doc`. Storing the original length here keeps + /// the smallfloat array lossy (space-efficient for ranking) while + /// making corpus-stat accounting exact. + pub(super) fieldnorms_exact: Vec, +} + +/// Postings and terms a mutation removed from one index. +#[derive(Debug, Default, Clone, Copy, PartialEq, Eq)] +pub(crate) struct Removed { + pub(crate) postings: usize, + pub(crate) terms: usize, +} + +impl ScopeState { + /// Insert a posting for a term. Returns whether the term is new here. + pub(super) fn insert(&mut self, term: &str, posting: CompactPosting) -> bool { + self.total_postings += 1; + match self.postings.get_mut(term) { + Some(list) => { + list.push(posting); + false + } + None => { + self.postings.insert(term.to_string(), vec![posting]); + true + } + } + } + + /// Record a document's stats (call once per indexed document). + pub(super) fn record_doc(&mut self, doc_id: Surrogate, doc_len: u32) { + self.stats.0 += 1; + self.stats.1 += doc_len as u64; + + let idx = doc_id.0 as usize; + if idx >= self.fieldnorms.len() { + self.fieldnorms.resize(idx + 1, 0); + } + self.fieldnorms[idx] = smallfloat::encode(doc_len); + if idx >= self.fieldnorms_exact.len() { + self.fieldnorms_exact.resize(idx + 1, 0); + } + self.fieldnorms_exact[idx] = doc_len; + } + + /// Remove a document's postings from every term of this index. + pub(super) fn remove_doc(&mut self, doc_id: Surrogate) -> Removed { + let mut removed = Removed::default(); + self.postings.retain(|_, postings| { + let before = postings.len(); + postings.retain(|p| p.doc_id != doc_id); + removed.postings += before - postings.len(); + if postings.len() < before + && !postings.is_empty() + && postings.capacity() > postings.len().saturating_mul(2).max(4) + { + postings.shrink_to(postings.len()); + } + let keep = !postings.is_empty(); + if !keep { + removed.terms += 1; + } + keep + }); + self.total_postings -= removed.postings; + + // Decrement stats using the exact-length sidecar so the subtraction + // matches the exact-length increment in record_doc. Using the + // smallfloat-decoded length here would drift upward on churn. + if let Some(slot) = self.fieldnorms_exact.get_mut(doc_id.0 as usize) + && *slot > 0 + { + let doc_len = *slot; + *slot = 0; + self.stats.0 = self.stats.0.saturating_sub(1); + self.stats.1 = self.stats.1.saturating_sub(doc_len as u64); + } + removed + } + + /// Whether the index holds neither postings nor counted documents. + pub(super) fn is_vacant(&self) -> bool { + self.postings.is_empty() && self.stats.0 == 0 + } + + /// What dropping this whole state removes. + pub(super) fn footprint(&self) -> Removed { + Removed { + postings: self.total_postings, + terms: self.postings.len(), + } + } +} diff --git a/nodedb-fts/src/lsm/memtable/table.rs b/nodedb-fts/src/lsm/memtable/table.rs new file mode 100644 index 000000000..c71219411 --- /dev/null +++ b/nodedb-fts/src/lsm/memtable/table.rs @@ -0,0 +1,649 @@ +// SPDX-License-Identifier: Apache-2.0 + +//! In-memory memtable for the FTS LSM engine. +//! +//! Accumulates postings per index until a spill threshold is reached. +//! Serves queries from memory. Thread-safe via interior mutability. +//! +//! Every index `(database, tenant, collection, field)` owns its postings, +//! corpus stats, and fieldnorms. Indexes are keyed structurally, so no +//! collection or field name can collide with another's. The spill +//! thresholds bound the sum over every index. + +use std::cell::{Cell, RefCell}; +use std::collections::HashMap; + +use nodedb_types::Surrogate; + +use super::config::MemtableConfig; +use super::scope_state::{Removed, ScopeState}; +use crate::block::CompactPosting; +use crate::scope::IndexScope; + +/// collection → field key → index state. +type CollectionScopes = HashMap>; +/// `(database_id, tid)` → that tenant's indexes. +type ScopeMap = HashMap<(u64, u64), CollectionScopes>; + +/// One index the memtable holds state for. +#[derive(Debug, Clone, PartialEq, Eq, Hash)] +pub struct MemtableScope { + pub database_id: u64, + pub tid: u64, + pub collection: String, + /// Empty for the whole-document index. + pub field: String, +} + +impl MemtableScope { + /// The index this scope addresses. + pub fn index(&self) -> IndexScope<'_> { + IndexScope::from_key(&self.collection, &self.field) + } +} + +/// In-memory accumulator for FTS postings. +/// +/// Stores per-index, per-term posting lists keyed on `Surrogate` row +/// identities. Maintains incremental corpus stats per index. +pub struct Memtable { + scopes: RefCell, + /// Posting entries across every index. + total_postings: Cell, + /// Distinct `(index, term)` pairs across every index. + total_terms: Cell, + /// Spill configuration. + config: MemtableConfig, +} + +impl Memtable { + pub fn new(config: MemtableConfig) -> Self { + Self { + scopes: RefCell::new(HashMap::new()), + total_postings: Cell::new(0), + total_terms: Cell::new(0), + config, + } + } + + /// Insert a posting for a term of one index. + pub fn insert( + &self, + database_id: u64, + tid: u64, + index: IndexScope<'_>, + term: &str, + posting: CompactPosting, + ) { + let mut scopes = self.scopes.borrow_mut(); + let state = state_entry(&mut scopes, database_id, tid, index); + let new_term = state.insert(term, posting); + self.add(1, usize::from(new_term)); + } + + /// Insert every posting of one document and record its length, in one + /// index lookup. + pub fn insert_doc<'t>( + &self, + database_id: u64, + tid: u64, + index: IndexScope<'_>, + doc_id: Surrogate, + doc_len: u32, + postings: impl IntoIterator, + ) { + let mut scopes = self.scopes.borrow_mut(); + let state = state_entry(&mut scopes, database_id, tid, index); + let (mut added, mut new_terms) = (0usize, 0usize); + for (term, posting) in postings { + added += 1; + new_terms += usize::from(state.insert(term, posting)); + } + state.record_doc(doc_id, doc_len); + self.add(added, new_terms); + } + + /// Record a document's stats in one index (call once per indexed document). + pub fn record_doc( + &self, + database_id: u64, + tid: u64, + index: IndexScope<'_>, + doc_id: Surrogate, + doc_len: u32, + ) { + let mut scopes = self.scopes.borrow_mut(); + state_entry(&mut scopes, database_id, tid, index).record_doc(doc_id, doc_len); + } + + /// Remove a document's postings from every term of one index. + pub fn remove_doc(&self, database_id: u64, tid: u64, index: IndexScope<'_>, doc_id: Surrogate) { + let mut scopes = self.scopes.borrow_mut(); + let Some(fields) = scopes + .get_mut(&(database_id, tid)) + .and_then(|collections| collections.get_mut(index.collection())) + else { + return; + }; + let Some(state) = fields.get_mut(index.field_key()) else { + return; + }; + let removed = state.remove_doc(doc_id); + if state.is_vacant() { + fields.remove(index.field_key()); + } + prune_empty(&mut scopes, database_id, tid, index.collection()); + drop(scopes); + self.subtract(removed); + } + + /// Check if the memtable should be flushed (spill threshold reached). + /// The thresholds bound the sum over every index. + pub fn should_flush(&self) -> bool { + self.total_postings.get() >= self.config.max_postings + || self.total_terms.get() >= self.config.max_terms + } + + /// Get the posting list for a term of one index. Empty if not found. + pub fn get_postings( + &self, + database_id: u64, + tid: u64, + index: IndexScope<'_>, + term: &str, + ) -> Vec { + self.read(database_id, tid, index, |s| { + s.postings.get(term).cloned().unwrap_or_default() + }) + .unwrap_or_default() + } + + /// Get all term names of one index. + pub fn terms(&self, database_id: u64, tid: u64, index: IndexScope<'_>) -> Vec { + self.read(database_id, tid, index, |s| { + s.postings.keys().cloned().collect() + }) + .unwrap_or_default() + } + + /// Get corpus stats of one index: (doc_count, total_token_sum). + pub fn stats(&self, database_id: u64, tid: u64, index: IndexScope<'_>) -> (u32, u64) { + self.read(database_id, tid, index, |s| s.stats) + .unwrap_or((0, 0)) + } + + /// Get the fieldnorm array (SmallFloat-encoded doc lengths) of one index. + pub fn fieldnorms(&self, database_id: u64, tid: u64, index: IndexScope<'_>) -> Vec { + self.read(database_id, tid, index, |s| s.fieldnorms.clone()) + .unwrap_or_default() + } + + /// Run `f` over one index's term→postings map without removing it. + /// `None` when the memtable holds no state for the index. + pub fn with_scope_postings( + &self, + database_id: u64, + tid: u64, + index: IndexScope<'_>, + f: impl FnOnce(&HashMap>) -> R, + ) -> Option { + self.read(database_id, tid, index, |s| f(&s.postings)) + } + + /// Every index the memtable holds state for. + pub fn scopes(&self) -> Vec { + let scopes = self.scopes.borrow(); + let mut out = Vec::new(); + for (&(database_id, tid), collections) in scopes.iter() { + for (collection, fields) in collections { + for field in fields.keys() { + out.push(MemtableScope { + database_id, + tid, + collection: collection.clone(), + field: field.clone(), + }); + } + } + } + out + } + + /// Drain one index's postings (for flush). Returns its term→postings + /// map and drops its state, including stats and fieldnorms. Every + /// other index is untouched. + pub fn drain_scope( + &self, + database_id: u64, + tid: u64, + index: IndexScope<'_>, + ) -> HashMap> { + let mut scopes = self.scopes.borrow_mut(); + let Some(state) = scopes + .get_mut(&(database_id, tid)) + .and_then(|collections| collections.get_mut(index.collection())) + .and_then(|fields| fields.remove(index.field_key())) + else { + return HashMap::new(); + }; + prune_empty(&mut scopes, database_id, tid, index.collection()); + drop(scopes); + self.subtract(state.footprint()); + state.postings + } + + /// Drop every index of one collection. + pub fn drain_collection(&self, database_id: u64, tid: u64, collection: &str) { + let mut scopes = self.scopes.borrow_mut(); + let Some(fields) = scopes + .get_mut(&(database_id, tid)) + .and_then(|collections| collections.remove(collection)) + else { + return; + }; + prune_empty(&mut scopes, database_id, tid, collection); + drop(scopes); + for state in fields.values() { + self.subtract(state.footprint()); + } + } + + /// Drop every index of one `(database_id, tid)`. + pub fn drain_tenant(&self, database_id: u64, tid: u64) { + let Some(collections) = self.scopes.borrow_mut().remove(&(database_id, tid)) else { + return; + }; + for state in collections.values().flat_map(HashMap::values) { + self.subtract(state.footprint()); + } + } + + /// Distinct `(index, term)` pairs across every index. + pub fn term_count(&self) -> usize { + self.total_terms.get() + } + + /// Posting entries across every index. + pub fn posting_count(&self) -> usize { + self.total_postings.get() + } + + /// Whether no index holds a posting. + pub fn is_empty(&self) -> bool { + self.total_terms.get() == 0 + } + + fn read( + &self, + database_id: u64, + tid: u64, + index: IndexScope<'_>, + f: impl FnOnce(&ScopeState) -> R, + ) -> Option { + self.scopes + .borrow() + .get(&(database_id, tid)) + .and_then(|collections| collections.get(index.collection())) + .and_then(|fields| fields.get(index.field_key())) + .map(f) + } + + fn add(&self, postings: usize, terms: usize) { + self.total_postings + .set(self.total_postings.get() + postings); + self.total_terms.set(self.total_terms.get() + terms); + } + + fn subtract(&self, removed: Removed) { + self.total_postings + .set(self.total_postings.get().saturating_sub(removed.postings)); + self.total_terms + .set(self.total_terms.get().saturating_sub(removed.terms)); + } +} + +/// The state of one index, created empty on first use. +fn state_entry<'m>( + scopes: &'m mut ScopeMap, + database_id: u64, + tid: u64, + index: IndexScope<'_>, +) -> &'m mut ScopeState { + scopes + .entry((database_id, tid)) + .or_default() + .entry(index.collection().to_string()) + .or_default() + .entry(index.field_key().to_string()) + .or_default() +} + +/// Drop the collection and tenant maps that no longer hold any index. +fn prune_empty(scopes: &mut ScopeMap, database_id: u64, tid: u64, collection: &str) { + let Some(collections) = scopes.get_mut(&(database_id, tid)) else { + return; + }; + if collections.get(collection).is_some_and(HashMap::is_empty) { + collections.remove(collection); + } + if collections.is_empty() { + scopes.remove(&(database_id, tid)); + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::codec::smallfloat; + + const DB: u64 = 0; + const T: u64 = 1; + const S: IndexScope<'static> = IndexScope::document("docs"); + + fn make_posting(doc_id: u32, tf: u32) -> CompactPosting { + CompactPosting { + doc_id: Surrogate(doc_id), + term_freq: tf, + fieldnorm: smallfloat::encode(100), + positions: vec![0], + } + } + + #[test] + fn anchored_posting_lists_release_sparse_migration_capacity() { + let mt = Memtable::new(MemtableConfig::default()); + let terms = 12u32; + let movers = 64u32; + for term in 0..terms { + mt.insert(DB, T, S, &format!("term{term}"), make_posting(term + 1, 1)); + mt.record_doc(DB, T, S, Surrogate(term + 1), 1); + } + for term in 0..terms { + let key = format!("term{term}"); + for mover in 0..movers { + let id = terms + mover + 1; + mt.insert(DB, T, S, &key, make_posting(id, 1)); + mt.record_doc(DB, T, S, Surrogate(id), 1); + } + assert_eq!(mt.posting_count(), (terms + movers) as usize); + assert_eq!( + mt.stats(DB, T, S), + (terms + movers, (terms + movers) as u64) + ); + for mover in 0..movers { + mt.remove_doc(DB, T, S, Surrogate(terms + mover + 1)); + let capacities_ok = mt + .read(DB, T, S, |s| { + s.postings.values().all(|postings| { + postings.capacity() <= postings.len().saturating_mul(2).max(4) + }) + }) + .unwrap_or(true); + assert!(capacities_ok); + } + let anchor = mt.get_postings(DB, T, S, &key); + assert_eq!(anchor.len(), 1); + assert_eq!(anchor[0].doc_id, Surrogate(term + 1)); + assert_eq!(anchor[0].positions, vec![0]); + assert_eq!(mt.posting_count(), terms as usize); + assert_eq!(mt.stats(DB, T, S), (terms, terms as u64)); + } + for term in 0..terms { + assert_eq!( + mt.get_postings(DB, T, S, &format!("term{term}"))[0].doc_id, + Surrogate(term + 1) + ); + } + } + + #[test] + fn insert_and_query() { + let mt = Memtable::new(MemtableConfig::default()); + mt.insert(DB, T, S, "hello", make_posting(0, 1)); + mt.insert(DB, T, S, "hello", make_posting(1, 2)); + mt.insert(DB, T, S, "world", make_posting(0, 1)); + + assert_eq!(mt.get_postings(DB, T, S, "hello").len(), 2); + assert_eq!(mt.get_postings(DB, T, S, "world").len(), 1); + assert!(mt.get_postings(DB, T, S, "missing").is_empty()); + assert_eq!(mt.term_count(), 2); + assert_eq!(mt.posting_count(), 3); + } + + #[test] + fn remove_doc() { + let mt = Memtable::new(MemtableConfig::default()); + mt.insert(DB, T, S, "hello", make_posting(0, 1)); + mt.insert(DB, T, S, "hello", make_posting(1, 2)); + mt.record_doc(DB, T, S, Surrogate(0), 100); + mt.record_doc(DB, T, S, Surrogate(1), 50); + + mt.remove_doc(DB, T, S, Surrogate(0)); + + assert_eq!(mt.get_postings(DB, T, S, "hello").len(), 1); + assert_eq!(mt.get_postings(DB, T, S, "hello")[0].doc_id, Surrogate(1)); + assert_eq!(mt.stats(DB, T, S).0, 1); // Only doc 1 remains. + } + + #[test] + fn remove_doc_touches_only_its_index() { + let mt = Memtable::new(MemtableConfig::default()); + let title = IndexScope::field("docs", "title").unwrap(); + mt.insert_doc(DB, T, S, Surrogate(1), 2, [("rust", make_posting(1, 1))]); + mt.insert_doc( + DB, + T, + title, + Surrogate(1), + 1, + [("rust", make_posting(1, 1))], + ); + + mt.remove_doc(DB, T, title, Surrogate(1)); + + assert!(mt.get_postings(DB, T, title, "rust").is_empty()); + assert_eq!(mt.stats(DB, T, title), (0, 0)); + assert_eq!(mt.get_postings(DB, T, S, "rust").len(), 1); + assert_eq!(mt.stats(DB, T, S), (1, 2)); + assert_eq!(mt.posting_count(), 1); + assert_eq!(mt.term_count(), 1); + } + + #[test] + fn drain_scope_resets_only_that_index() { + let mt = Memtable::new(MemtableConfig::default()); + let other = IndexScope::document("other"); + mt.insert_doc( + DB, + T, + S, + Surrogate(0), + 100, + [("hello", make_posting(0, 1)), ("world", make_posting(0, 1))], + ); + mt.insert_doc( + DB, + T, + other, + Surrogate(1), + 50, + [("rust", make_posting(1, 1))], + ); + + let drained = mt.drain_scope(DB, T, S); + assert_eq!(drained.len(), 2); + assert_eq!(mt.stats(DB, T, S), (0, 0)); + assert!(mt.fieldnorms(DB, T, S).is_empty()); + assert!(mt.terms(DB, T, S).is_empty()); + + assert_eq!(mt.stats(DB, T, other), (1, 50)); + assert_eq!(mt.get_postings(DB, T, other, "rust").len(), 1); + assert_eq!(mt.posting_count(), 1); + assert_eq!(mt.term_count(), 1); + assert!(!mt.is_empty()); + + mt.drain_scope(DB, T, other); + assert!(mt.is_empty()); + assert_eq!(mt.posting_count(), 0); + assert!(mt.scopes().is_empty()); + } + + #[test] + fn drain_collection_selective() { + let mt = Memtable::new(MemtableConfig::default()); + let a = IndexScope::document("col_a"); + let a_title = IndexScope::field("col_a", "title").unwrap(); + let b = IndexScope::document("col_b"); + mt.insert(DB, T, a, "hello", make_posting(0, 1)); + mt.insert(DB, T, a_title, "world", make_posting(1, 1)); + mt.insert(DB, T, b, "rust", make_posting(2, 1)); + mt.record_doc(DB, T, b, Surrogate(2), 1); + + mt.drain_collection(DB, T, "col_a"); + + assert!(mt.get_postings(DB, T, a, "hello").is_empty()); + assert!(mt.get_postings(DB, T, a_title, "world").is_empty()); + assert_eq!(mt.get_postings(DB, T, b, "rust").len(), 1); + assert_eq!(mt.stats(DB, T, b), (1, 1)); + assert_eq!(mt.posting_count(), 1); + } + + #[test] + fn drain_tenant_keeps_other_tenants() { + let mt = Memtable::new(MemtableConfig::default()); + mt.insert(DB, 1, S, "hello", make_posting(0, 1)); + mt.insert(DB, 2, S, "hello", make_posting(0, 1)); + + mt.drain_tenant(DB, 1); + + assert!(mt.get_postings(DB, 1, S, "hello").is_empty()); + assert_eq!(mt.get_postings(DB, 2, S, "hello").len(), 1); + assert_eq!(mt.posting_count(), 1); + } + + /// Collection `a:b` with field `c` and collection `a` with field `b:c` + /// are distinct indexes: no separator in a name can merge them. + #[test] + fn names_with_separators_never_collide() { + let mt = Memtable::new(MemtableConfig::default()); + let left = IndexScope::field("a:b", "c").unwrap(); + let right = IndexScope::field("a", "b:c").unwrap(); + mt.insert(DB, T, left, "term", make_posting(1, 1)); + + assert_eq!(mt.get_postings(DB, T, left, "term").len(), 1); + assert!(mt.get_postings(DB, T, right, "term").is_empty()); + assert_eq!(mt.scopes().len(), 1); + } + + #[test] + fn spill_threshold_sums_every_index() { + let config = MemtableConfig { + max_postings: 5, + max_terms: 100, + }; + let mt = Memtable::new(config); + let title = IndexScope::field("docs", "title").unwrap(); + for i in 0..2 { + mt.insert(DB, T, S, "term", make_posting(i, 1)); + } + for i in 0..2 { + mt.insert(DB, T, title, "term", make_posting(i, 1)); + } + assert!(!mt.should_flush()); + mt.insert(DB, T, title, "term", make_posting(4, 1)); + assert!(mt.should_flush()); + } + + #[test] + fn stats_invariant_under_insert_delete_churn() { + // Spec: after any sequence of (record_doc, remove_doc) calls, the + // memtable's corpus stats — specifically `total_token_sum` that feeds + // BM25 `avg_doc_len` — must equal the stats of a memtable freshly + // populated with only the currently-live documents. + let churn = Memtable::new(MemtableConfig::default()); + let doc_id = Surrogate(0); + let doc_len = 137u32; // not smallfloat-representable exactly + for _ in 0..500 { + churn.record_doc(DB, T, S, doc_id, doc_len); + churn.remove_doc(DB, T, S, doc_id); + } + // After equal numbers of record / remove, the memtable should be empty + // from a stats standpoint: zero docs, zero tokens. + let fresh = Memtable::new(MemtableConfig::default()); + assert_eq!( + churn.stats(DB, T, S), + fresh.stats(DB, T, S), + "stats drifted after insert/delete churn" + ); + } + + #[test] + fn stats_invariant_churn_then_single_live_doc() { + // Spec: churn on one doc, then leave a single live doc. Stats must + // exactly match a memtable that only ever saw that single live doc. + let churn = Memtable::new(MemtableConfig::default()); + for _ in 0..200 { + churn.record_doc(DB, T, S, Surrogate(0), 400); + churn.remove_doc(DB, T, S, Surrogate(0)); + } + churn.record_doc(DB, T, S, Surrogate(0), 250); + + let fresh = Memtable::new(MemtableConfig::default()); + fresh.record_doc(DB, T, S, Surrogate(0), 250); + + assert_eq!( + churn.stats(DB, T, S), + fresh.stats(DB, T, S), + "stats diverged from fresh baseline" + ); + } + + #[test] + fn stats_total_token_sum_never_exceeds_live_doc_length() { + // Regression guard on the specific failure mode: `total_token_sum` + // growing beyond the sum of live document lengths. If this fails, + // BM25 `avg_doc_len` is inflated and ranking is silently skewed. + let mt = Memtable::new(MemtableConfig::default()); + for _ in 0..100 { + mt.record_doc(DB, T, S, Surrogate(0), 777); + mt.remove_doc(DB, T, S, Surrogate(0)); + } + mt.record_doc(DB, T, S, Surrogate(0), 777); + let (count, total) = mt.stats(DB, T, S); + assert_eq!(count, 1); + assert_eq!( + total, 777, + "total_token_sum drifted: got {total}, expected 777 for a single 777-token live doc" + ); + } + + #[test] + fn stats_are_per_index() { + let mt = Memtable::new(MemtableConfig::default()); + let title = IndexScope::field("docs", "title").unwrap(); + mt.record_doc(DB, T, S, Surrogate(1), 10); + mt.record_doc(DB, T, S, Surrogate(2), 20); + mt.record_doc(DB, T, title, Surrogate(1), 3); + + assert_eq!(mt.stats(DB, T, S), (2, 30)); + assert_eq!(mt.stats(DB, T, title), (1, 3)); + } + + #[test] + fn fieldnorms_recorded() { + let mt = Memtable::new(MemtableConfig::default()); + mt.record_doc(DB, T, S, Surrogate(0), 100); + mt.record_doc(DB, T, S, Surrogate(5), 50); + + let norms = mt.fieldnorms(DB, T, S); + assert_eq!(norms.len(), 6); // 0..=5 + assert_eq!( + smallfloat::decode(norms[0]), + smallfloat::decode(smallfloat::encode(100)) + ); + assert_eq!( + smallfloat::decode(norms[5]), + smallfloat::decode(smallfloat::encode(50)) + ); + } +} diff --git a/nodedb-fts/src/lsm/merge.rs b/nodedb-fts/src/lsm/merge.rs index 99cf2a51f..9b94ca6dd 100644 --- a/nodedb-fts/src/lsm/merge.rs +++ b/nodedb-fts/src/lsm/merge.rs @@ -10,13 +10,19 @@ use std::collections::{BTreeSet, HashMap}; use nodedb_types::Surrogate; use crate::block::{CompactPosting, PostingBlock, into_blocks}; +use crate::index::FtsIndexError; -use super::segment::reader::SegmentReader; +use super::query::{LiveSegment, live_segment_postings}; use nodedb_mem::ScopedMemory; /// Merge multiple segments into a single set of per-term PostingBlocks. /// +/// Each segment's postings of its deleted documents are dropped, so the +/// merged segment holds live postings only. A term left with no live +/// posting is not written. Posting data that does not decode is a +/// [`FtsIndexError::CorruptSegment`]. +/// /// The result is a sorted list of `(term, blocks)` suitable for /// `segment::writer::build_from_blocks`. /// @@ -25,14 +31,14 @@ use nodedb_mem::ScopedMemory; /// exceeded the allocation still proceeds — the scope serves as an /// accounting and backpressure signal; callers that need hard rejection /// check pressure before dispatching the operation. -pub fn merge_segments( - segments: &[SegmentReader], +pub fn merge_segments( + segments: &[LiveSegment], memory: &ScopedMemory, -) -> Vec<(String, Vec)> { +) -> Result)>, FtsIndexError> { // Collect all unique terms across all segments. let mut all_terms = BTreeSet::new(); for seg in segments { - for entry in seg.term_dict() { + for entry in seg.reader.term_dict() { all_terms.insert(entry.term.clone()); } } @@ -44,21 +50,10 @@ pub fn merge_segments( let mut result = Vec::with_capacity(all_terms.len()); for term in &all_terms { - // Gather all postings for this term across segments. + // Gather the live postings for this term across segments. let mut merged_postings: Vec = Vec::new(); - for seg in segments { - let blocks = seg.read_postings(term); - for block in blocks { - for i in 0..block.doc_ids.len() { - merged_postings.push(CompactPosting { - doc_id: block.doc_ids[i], - term_freq: block.term_freqs[i], - fieldnorm: block.fieldnorms[i], - positions: block.positions[i].clone(), - }); - } - } + merged_postings.extend(live_segment_postings::(seg, term)?); } if merged_postings.is_empty() { @@ -73,7 +68,7 @@ pub fn merge_segments( result.push((term.clone(), blocks)); } - result + Ok(result) } /// Merge posting lists from a memtable and multiple segments for a single term. @@ -117,8 +112,11 @@ pub fn dedup_postings(postings: &mut Vec) { #[cfg(test)] mod tests { + use std::convert::Infallible; + use super::*; use crate::codec::smallfloat; + use crate::lsm::segment::reader::SegmentReader; use crate::lsm::segment::writer; use crate::test_support::test_memory; @@ -146,10 +144,11 @@ mod tests { writer::flush_to_segment(m).unwrap() }; - let r1 = SegmentReader::open(seg1).expect("seg1 must be valid"); - let r2 = SegmentReader::open(seg2).expect("seg2 must be valid"); - - let merged = merge_segments(&[r1, r2], &test_memory()); + let merged = merge_segments::( + &[live("s1", seg1, None), live("s2", seg2, None)], + &test_memory(), + ) + .unwrap(); let terms: Vec<&str> = merged.iter().map(|(t, _)| t.as_str()).collect(); assert!(terms.contains(&"hello")); assert!(terms.contains(&"world")); @@ -161,6 +160,44 @@ mod tests { assert_eq!(total_docs, 4); } + fn live(id: &str, data: Vec, deleted: Option<&[u32]>) -> LiveSegment { + LiveSegment { + segment_id: id.to_string(), + reader: SegmentReader::open(data).expect("segment must be valid"), + deleted: deleted.map(|ids| ids.iter().map(|d| Surrogate(*d)).collect()), + } + } + + /// A segment's deleted documents leave the merge, and a term holding + /// only deleted documents is not written. + #[test] + fn merge_drops_deleted_postings() { + let seg1 = { + let mut m = std::collections::HashMap::new(); + m.insert("hello".to_string(), vec![cp(1, 1), cp(2, 1)]); + m.insert("gone".to_string(), vec![cp(2, 1)]); + writer::flush_to_segment(m).unwrap() + }; + let seg2 = { + let mut m = std::collections::HashMap::new(); + m.insert("hello".to_string(), vec![cp(3, 1)]); + writer::flush_to_segment(m).unwrap() + }; + let merged = merge_segments::( + &[live("s1", seg1, Some(&[2])), live("s2", seg2, None)], + &test_memory(), + ) + .unwrap(); + let terms: Vec<&str> = merged.iter().map(|(t, _)| t.as_str()).collect(); + assert_eq!(terms, vec!["hello"]); + let docs: Vec = merged[0] + .1 + .iter() + .flat_map(|b| b.doc_ids.iter().copied()) + .collect(); + assert_eq!(docs, vec![Surrogate(1), Surrogate(3)]); + } + #[test] fn merge_term_dedup() { let mt_posts = vec![cp(0, 5)]; // Memtable version (newer, tf=5). diff --git a/nodedb-fts/src/lsm/mod.rs b/nodedb-fts/src/lsm/mod.rs index 2884d6c49..62c083cd0 100644 --- a/nodedb-fts/src/lsm/mod.rs +++ b/nodedb-fts/src/lsm/mod.rs @@ -6,3 +6,4 @@ pub mod merge; pub mod parallel_build; pub mod query; pub mod segment; +pub mod segment_deletes; diff --git a/nodedb-fts/src/lsm/parallel_build.rs b/nodedb-fts/src/lsm/parallel_build.rs index 065e63772..60dc10d32 100644 --- a/nodedb-fts/src/lsm/parallel_build.rs +++ b/nodedb-fts/src/lsm/parallel_build.rs @@ -8,15 +8,18 @@ //! is responsible for partitioning documents and spawning workers. use std::collections::HashMap; +use std::convert::Infallible; use nodedb_mem::ScopedMemory; use nodedb_types::Surrogate; use crate::block::CompactPosting; use crate::codec::smallfloat; +use crate::index::FtsIndexError; use super::merge; -use super::segment::{reader::SegmentReader, writer}; +use super::query::LiveSegment; +use super::segment::{error::SegmentError, reader::SegmentReader, writer}; /// A worker's accumulated result: per-term postings ready to flush. pub struct WorkerResult { @@ -39,14 +42,10 @@ impl WorkerResult { .push(posting); } - /// Flush this worker's result to a temporary segment. - /// - /// Panics if any term exceeds `MAX_TERM_LEN` — callers must validate - /// term lengths before insertion. - pub fn flush_to_segment(self) -> Vec { - writer::flush_to_segment(self.term_postings).expect( - "worker result contained a term exceeding u16::MAX bytes — caller invariant violated", - ) + /// Flush this worker's result to a temporary segment. A term longer + /// than `MAX_TERM_LEN` is [`SegmentError::TermTooLong`]. + pub fn flush_to_segment(self) -> Result, SegmentError> { + writer::flush_to_segment(self.term_postings) } } @@ -59,21 +58,29 @@ impl Default for WorkerResult { /// Merge multiple worker segments into a single compacted segment. /// /// This is the "leader" step: takes the flushed segments from all workers -/// and performs N-way merge into one final segment. -pub fn merge_worker_segments(worker_segments: Vec>, memory: &ScopedMemory) -> Vec { - let readers: Vec = worker_segments - .into_iter() - .filter_map(|data| SegmentReader::open(data).ok()) - .collect(); - - if readers.is_empty() { - return writer::build_from_blocks(&[]) - .expect("build_from_blocks on empty input must not fail"); +/// and performs N-way merge into one final segment. A worker segment that +/// fails validation fails the merge: merging without it would lose that +/// worker's postings. +pub fn merge_worker_segments( + worker_segments: Vec>, + memory: &ScopedMemory, +) -> Result, FtsIndexError> { + let mut sources = Vec::with_capacity(worker_segments.len()); + for (worker, data) in worker_segments.into_iter().enumerate() { + let segment_id = format!("worker:{worker}"); + let reader = SegmentReader::open(data).map_err(|source| FtsIndexError::CorruptSegment { + segment_id: segment_id.clone(), + source, + })?; + sources.push(LiveSegment { + segment_id, + reader, + deleted: None, + }); } - let merged_term_blocks = merge::merge_segments(&readers, memory); - writer::build_from_blocks(&merged_term_blocks) - .expect("merge produced a term exceeding u16::MAX bytes — data invariant violated") + let merged_term_blocks = merge::merge_segments::(&sources, memory)?; + Ok(writer::build_from_blocks(&merged_term_blocks)?) } /// Partition a document range into `num_workers` disjoint sub-ranges. @@ -155,11 +162,11 @@ mod tests { make_compact_posting(Surrogate(3), 3, 120, vec![0, 2, 7]), ); - let seg1 = w1.flush_to_segment(); - let seg2 = w2.flush_to_segment(); + let seg1 = w1.flush_to_segment().unwrap(); + let seg2 = w2.flush_to_segment().unwrap(); // Leader merge. - let merged = merge_worker_segments(vec![seg1, seg2], &test_memory()); + let merged = merge_worker_segments(vec![seg1, seg2], &test_memory()).unwrap(); // Verify merged segment. let reader = SegmentReader::open(merged).expect("merged segment must be valid"); @@ -171,8 +178,22 @@ mod tests { #[test] fn merge_empty_workers() { - let merged = merge_worker_segments(Vec::new(), &test_memory()); + let merged = merge_worker_segments(Vec::new(), &test_memory()).unwrap(); let reader = SegmentReader::open(merged).expect("merged segment must be valid"); assert_eq!(reader.num_terms(), 0); } + + #[test] + fn a_corrupt_worker_segment_fails_the_merge() { + let mut w1 = WorkerResult::new(); + w1.insert("hello", make_compact_posting(Surrogate(1), 1, 50, vec![0])); + let mut seg = w1.flush_to_segment().unwrap(); + let last = seg.len() - 1; + seg[last] ^= 0xFF; + let err = merge_worker_segments(vec![seg], &test_memory()).unwrap_err(); + assert!( + matches!(err, FtsIndexError::CorruptSegment { ref segment_id, .. } if segment_id == "worker:0"), + "{err}" + ); + } } diff --git a/nodedb-fts/src/lsm/query.rs b/nodedb-fts/src/lsm/query.rs index 339113163..bf6435bd2 100644 --- a/nodedb-fts/src/lsm/query.rs +++ b/nodedb-fts/src/lsm/query.rs @@ -5,21 +5,105 @@ use crate::backend::FtsBackend; use crate::block::{CompactPosting, into_blocks}; -use crate::index::writer::{memtable_collection_prefix, memtable_key}; +use crate::index::FtsIndexError; +use crate::scope::IndexScope; use crate::search::bmw::skip_index::TermBlocks; use super::memtable::Memtable; use super::merge::merge_term_postings; use super::segment::reader::SegmentReader; +use super::segment_deletes::SegmentDeletes; use std::sync::Arc; use nodedb_mem::MemoryGovernor; +use nodedb_types::SurrogateBitmap; use crate::mem_scope::fts_scope; +/// One opened segment of an index with the documents removed after it was +/// written. +pub struct LiveSegment { + pub segment_id: String, + pub reader: SegmentReader, + pub deleted: Option, +} + +/// Open every segment of one index, each with its delete set. A segment +/// that fails validation is a [`FtsIndexError::CorruptSegment`], and a +/// listed segment the backend does not hold is a +/// [`FtsIndexError::MissingSegment`]: a read that skipped either would +/// answer from part of the index. +pub(crate) fn open_segments( + backend: &B, + database_id: u64, + tid: u64, + index: IndexScope<'_>, +) -> Result, FtsIndexError> { + let seg_ids = backend + .list_segments(database_id, tid, index) + .map_err(FtsIndexError::Backend)?; + if seg_ids.is_empty() { + return Ok(Vec::new()); + } + let mut deletes = SegmentDeletes::load(backend, database_id, tid, index, &seg_ids)?; + let mut segments: Vec = Vec::with_capacity(seg_ids.len()); + for segment_id in seg_ids { + let Some(data) = backend + .read_segment(database_id, tid, index, &segment_id) + .map_err(FtsIndexError::Backend)? + else { + return Err(FtsIndexError::MissingSegment { segment_id }); + }; + let reader = SegmentReader::open(data).map_err(|source| FtsIndexError::CorruptSegment { + segment_id: segment_id.clone(), + source, + })?; + let deleted = deletes.take(&segment_id); + segments.push(LiveSegment { + segment_id, + reader, + deleted, + }); + } + Ok(segments) +} + +/// The postings of `token` in one segment, without its deleted documents. +pub(crate) fn live_segment_postings( + segment: &LiveSegment, + token: &str, +) -> Result, FtsIndexError> { + let blocks = + segment + .reader + .read_postings(token) + .map_err(|source| FtsIndexError::CorruptSegment { + segment_id: segment.segment_id.clone(), + source, + })?; + let mut postings = Vec::new(); + for block in blocks { + for i in 0..block.doc_ids.len() { + let doc_id = block.doc_ids[i]; + if segment.deleted.as_ref().is_some_and(|d| d.contains(doc_id)) { + continue; + } + postings.push(CompactPosting { + doc_id, + term_freq: block.term_freqs[i], + fieldnorm: block.fieldnorms[i], + positions: block.positions[i].clone(), + }); + } + } + Ok(postings) +} + /// Collect posting lists for a set of query tokens by merging across -/// the active memtable and all immutable segments. +/// the active memtable and all immutable segments of one index. A +/// segment's postings of the documents removed after it was written are +/// skipped, so document frequencies count only live documents. /// /// Returns per-term `TermBlocks` ready for BMW scoring. /// @@ -31,20 +115,12 @@ pub fn collect_merged_term_blocks( backend: &B, database_id: u64, tid: u64, - collection: &str, + index: IndexScope<'_>, memtable: &Memtable, query_tokens: &[String], governor: &Arc, -) -> Result, B::Error> { - let seg_ids = backend.list_segments(database_id, tid, collection)?; - let mut readers: Vec = Vec::new(); - for id in &seg_ids { - if let Some(data) = backend.read_segment(database_id, tid, collection, id)? - && let Ok(reader) = SegmentReader::open(data) - { - readers.push(reader); - } - } +) -> Result, FtsIndexError> { + let segments = open_segments(backend, database_id, tid, index)?; let memory = fts_scope(governor, database_id, tid); let bytes = query_tokens.len() * std::mem::size_of::(); @@ -53,27 +129,12 @@ pub fn collect_merged_term_blocks( let mut term_blocks_list = Vec::with_capacity(query_tokens.len()); for token in query_tokens { - let scoped_term = memtable_key(database_id, tid, collection, token); - let mt_postings = memtable.get_postings(&scoped_term); + let mt_postings = memtable.get_postings(database_id, tid, index, token); - let seg_postings: Vec> = readers + let seg_postings: Vec> = segments .iter() - .map(|reader| { - let blocks = reader.read_postings(token); - let mut postings = Vec::new(); - for block in blocks { - for i in 0..block.doc_ids.len() { - postings.push(CompactPosting { - doc_id: block.doc_ids[i], - term_freq: block.term_freqs[i], - fieldnorm: block.fieldnorms[i], - positions: block.positions[i].clone(), - }); - } - } - postings - }) - .collect(); + .map(|segment| live_segment_postings::(segment, token)) + .collect::>()?; let merged = merge_term_postings(&mt_postings, &seg_postings); if merged.is_empty() { @@ -88,47 +149,38 @@ pub fn collect_merged_term_blocks( Ok(term_blocks_list) } -/// Collect all unique term names across memtable + segments for a collection. +/// Collect all unique term names across memtable + segments of one index. /// /// Used by fuzzy matching to scan available terms. pub fn collect_all_terms( backend: &B, database_id: u64, tid: u64, - collection: &str, + index: IndexScope<'_>, memtable: &Memtable, -) -> Result, B::Error> { - let prefix = memtable_collection_prefix(database_id, tid, collection); - let mut terms: std::collections::HashSet = std::collections::HashSet::new(); - - for key in memtable.terms() { - if let Some(term) = key.strip_prefix(&prefix) { - terms.insert(term.to_string()); - } - } - - let seg_ids = backend.list_segments(database_id, tid, collection)?; - for id in &seg_ids { - if let Some(data) = backend.read_segment(database_id, tid, collection, id)? - && let Ok(reader) = SegmentReader::open(data) - { - for term in reader.terms() { - terms.insert(term); - } +) -> Result, FtsIndexError> { + let mut terms: std::collections::HashSet = memtable + .terms(database_id, tid, index) + .into_iter() + .collect(); + + for segment in open_segments(backend, database_id, tid, index)? { + for term in segment.reader.terms() { + terms.insert(term); } } Ok(terms.into_iter().collect()) } -/// Compute merged corpus stats from memtable + all segments. +/// Corpus stats of one index: `(doc_count, avg_doc_len)`. pub fn merged_collection_stats( backend: &B, database_id: u64, tid: u64, - collection: &str, + index: IndexScope<'_>, ) -> Result<(u32, f32), B::Error> { - let (count, total) = backend.collection_stats(database_id, tid, collection)?; + let (count, total) = backend.collection_stats(database_id, tid, index)?; let avg = if count > 0 { total as f32 / count as f32 } else { @@ -149,6 +201,7 @@ mod tests { const DB: u64 = 0; const T: u64 = 1; + const COL: IndexScope<'static> = IndexScope::document("col"); fn cp(doc_id: u32, tf: u32) -> CompactPosting { CompactPosting { @@ -163,12 +216,12 @@ mod tests { fn memtable_only() { let backend = MemoryBackend::new(); let mt = Memtable::new(MemtableConfig::default()); - mt.insert(&memtable_key(DB, T, "col", "hello"), cp(0, 2)); - mt.insert(&memtable_key(DB, T, "col", "hello"), cp(1, 1)); + mt.insert(DB, T, COL, "hello", cp(0, 2)); + mt.insert(DB, T, COL, "hello", cp(1, 1)); let tokens = vec!["hello".to_string()]; let term_blocks = - collect_merged_term_blocks(&backend, DB, T, "col", &mt, &tokens, &test_governor()) + collect_merged_term_blocks(&backend, DB, T, COL, &mt, &tokens, &test_governor()) .unwrap(); assert_eq!(term_blocks.len(), 1); @@ -182,13 +235,13 @@ mod tests { postings.insert("hello".to_string(), vec![cp(0, 1), cp(5, 2)]); let seg_bytes = writer::flush_to_segment(postings).unwrap(); backend - .write_segment(DB, T, "col", "L0:0000000000000001", &seg_bytes) + .write_segment(DB, T, COL, "L0:0000000000000001", &seg_bytes) .unwrap(); let mt = Memtable::new(MemtableConfig::default()); let tokens = vec!["hello".to_string()]; let term_blocks = - collect_merged_term_blocks(&backend, DB, T, "col", &mt, &tokens, &test_governor()) + collect_merged_term_blocks(&backend, DB, T, COL, &mt, &tokens, &test_governor()) .unwrap(); assert_eq!(term_blocks.len(), 1); @@ -203,22 +256,103 @@ mod tests { seg_postings.insert("hello".to_string(), vec![cp(0, 1), cp(5, 2)]); let seg_bytes = writer::flush_to_segment(seg_postings).unwrap(); backend - .write_segment(DB, T, "col", "L0:0000000000000001", &seg_bytes) + .write_segment(DB, T, COL, "L0:0000000000000001", &seg_bytes) .unwrap(); let mt = Memtable::new(MemtableConfig::default()); - mt.insert(&memtable_key(DB, T, "col", "hello"), cp(0, 10)); - mt.insert(&memtable_key(DB, T, "col", "hello"), cp(3, 1)); + mt.insert(DB, T, COL, "hello", cp(0, 10)); + mt.insert(DB, T, COL, "hello", cp(3, 1)); let tokens = vec!["hello".to_string()]; let term_blocks = - collect_merged_term_blocks(&backend, DB, T, "col", &mt, &tokens, &test_governor()) + collect_merged_term_blocks(&backend, DB, T, COL, &mt, &tokens, &test_governor()) .unwrap(); assert_eq!(term_blocks.len(), 1); assert_eq!(term_blocks[0].df, 3); } + #[test] + fn other_index_postings_are_invisible() { + let backend = MemoryBackend::new(); + let title = IndexScope::field("col", "title").unwrap(); + let mt = Memtable::new(MemtableConfig::default()); + mt.insert(DB, T, title, "hello", cp(1, 1)); + + let tokens = vec!["hello".to_string()]; + let term_blocks = + collect_merged_term_blocks(&backend, DB, T, COL, &mt, &tokens, &test_governor()) + .unwrap(); + assert_eq!(term_blocks[0].df, 0); + assert!( + collect_all_terms(&backend, DB, T, COL, &mt) + .unwrap() + .is_empty() + ); + assert_eq!( + collect_all_terms(&backend, DB, T, title, &mt).unwrap(), + vec!["hello".to_string()] + ); + } + + /// A segment that fails validation fails the read instead of being + /// skipped. + #[test] + fn a_corrupt_segment_fails_the_read() { + let backend = MemoryBackend::new(); + let mut postings = HashMap::new(); + postings.insert("hello".to_string(), vec![cp(1, 1)]); + let mut seg_bytes = writer::flush_to_segment(postings).unwrap(); + let last = seg_bytes.len() - 1; + seg_bytes[last] ^= 0xFF; + backend + .write_segment(DB, T, COL, "L0:0000000000000001", &seg_bytes) + .unwrap(); + + let mt = Memtable::new(MemtableConfig::default()); + let tokens = vec!["hello".to_string()]; + let err = collect_merged_term_blocks(&backend, DB, T, COL, &mt, &tokens, &test_governor()) + .unwrap_err(); + assert!( + matches!( + err, + FtsIndexError::CorruptSegment { ref segment_id, .. } + if segment_id == "L0:0000000000000001" + ), + "{err}" + ); + assert!(matches!( + collect_all_terms(&backend, DB, T, COL, &mt).unwrap_err(), + FtsIndexError::CorruptSegment { .. } + )); + } + + /// A segment's deleted documents are skipped and leave the document + /// frequency. + #[test] + fn deleted_documents_of_a_segment_are_skipped() { + let backend = MemoryBackend::new(); + let mut postings = HashMap::new(); + postings.insert("hello".to_string(), vec![cp(1, 1), cp(2, 1)]); + let seg_bytes = writer::flush_to_segment(postings).unwrap(); + let id = "L0:0000000000000001".to_string(); + backend.write_segment(DB, T, COL, &id, &seg_bytes).unwrap(); + let mut deletes = SegmentDeletes::default(); + deletes.mark_removed(std::slice::from_ref(&id), nodedb_types::Surrogate(1)); + deletes.store(&backend, DB, T, COL).unwrap(); + + let mt = Memtable::new(MemtableConfig::default()); + let tokens = vec!["hello".to_string()]; + let blocks = + collect_merged_term_blocks(&backend, DB, T, COL, &mt, &tokens, &test_governor()) + .unwrap(); + assert_eq!(blocks[0].df, 1); + assert_eq!( + blocks[0].doc_ids().collect::>(), + vec![nodedb_types::Surrogate(2)] + ); + } + #[test] fn missing_term() { let backend = MemoryBackend::new(); @@ -226,7 +360,7 @@ mod tests { let tokens = vec!["nonexistent".to_string()]; let term_blocks = - collect_merged_term_blocks(&backend, DB, T, "col", &mt, &tokens, &test_governor()) + collect_merged_term_blocks(&backend, DB, T, COL, &mt, &tokens, &test_governor()) .unwrap(); assert_eq!(term_blocks.len(), 1); diff --git a/nodedb-fts/src/lsm/segment/error.rs b/nodedb-fts/src/lsm/segment/error.rs index db9d179cd..de2a8f58c 100644 --- a/nodedb-fts/src/lsm/segment/error.rs +++ b/nodedb-fts/src/lsm/segment/error.rs @@ -24,6 +24,11 @@ pub enum SegmentError { #[error("segment data is truncated")] Truncated, + /// A term's posting data lies outside the segment body or does not + /// decode, although the segment checksum matched. + #[error("posting data of term '{term}' is corrupt")] + CorruptPostings { term: String }, + /// A term exceeds the maximum encodable length (u16::MAX bytes). #[error("term length {term_len} exceeds maximum {max}")] TermTooLong { term_len: usize, max: usize }, diff --git a/nodedb-fts/src/lsm/segment/reader.rs b/nodedb-fts/src/lsm/segment/reader.rs index 3b8c75562..1f8395660 100644 --- a/nodedb-fts/src/lsm/segment/reader.rs +++ b/nodedb-fts/src/lsm/segment/reader.rs @@ -73,37 +73,39 @@ impl SegmentReader { /// Read and decode posting blocks for a term. /// - /// Returns empty vec if the term is not in this segment. - pub fn read_postings(&self, term: &str) -> Vec { + /// Returns an empty vec if the term is not in this segment. Posting data + /// that lies outside the segment body or does not decode is + /// [`SegmentError::CorruptPostings`]. + pub fn read_postings(&self, term: &str) -> Result, SegmentError> { let Some(entry) = self.find_term(term) else { - return Vec::new(); + return Ok(Vec::new()); + }; + let corrupt = || SegmentError::CorruptPostings { + term: term.to_string(), }; - let Some(start) = usize::try_from(self.header.posting_data_offset) + let start = usize::try_from(self.header.posting_data_offset) .ok() .and_then(|base| { usize::try_from(entry.posting_offset) .ok() .and_then(|offset| base.checked_add(offset)) }) - else { - return Vec::new(); - }; - let Some(end) = usize::try_from(entry.posting_len) + .ok_or_else(corrupt)?; + let end = usize::try_from(entry.posting_len) .ok() .and_then(|len| start.checked_add(len)) - else { - return Vec::new(); - }; - let Some(body_end) = self.data.len().checked_sub(format::FOOTER_SIZE) else { - return Vec::new(); - }; + .ok_or_else(corrupt)?; + let body_end = self + .data + .len() + .checked_sub(format::FOOTER_SIZE) + .ok_or_else(corrupt)?; if end > body_end { - return Vec::new(); + return Err(corrupt()); } - let buf = &self.data[start..end]; - decode_term_blocks(buf) + decode_term_blocks(&self.data[start..end]).ok_or_else(corrupt) } /// Get all unique terms in this segment. @@ -120,50 +122,47 @@ impl SegmentReader { /// Decode posting blocks from the term's posting data bytes. /// /// Format: [num_blocks: u32 LE][for each: block_len: u32 LE, block_bytes] -fn decode_term_blocks(buf: &[u8]) -> Vec { +/// +/// `None` when the bytes do not hold every block the count names, or a +/// block does not decode. +fn decode_term_blocks(buf: &[u8]) -> Option> { if buf.len() < 4 { - return Vec::new(); + return None; } - let num_blocks = match usize::try_from(u32::from_le_bytes([buf[0], buf[1], buf[2], buf[3]])) { - Ok(count) => count, - Err(_) => return Vec::new(), - }; + let num_blocks = usize::try_from(u32::from_le_bytes([buf[0], buf[1], buf[2], buf[3]])).ok()?; // Each block has at least its four-byte length prefix, so the remaining // bytes prove a safe upper bound before reserving the result vector. if num_blocks > (buf.len() - 4) / 4 { - return Vec::new(); + return None; } const MAX_POSTING_BLOCK_ALLOCATION_BYTES: usize = 64 * 1024 * 1024; - let Some(blocks_capacity) = checked_decode_capacity( + let blocks_capacity = checked_decode_capacity( num_blocks, size_of::(), buf.len() - 4, 4, (buf.len() - 4) / 4, MAX_POSTING_BLOCK_ALLOCATION_BYTES, - ) else { - return Vec::new(); - }; - let mut pos = 4; + )?; + let mut pos: usize = 4; let mut blocks = Vec::with_capacity(blocks_capacity); for _ in 0..num_blocks { - if pos + 4 > buf.len() { - break; - } - let block_len = - u32::from_le_bytes([buf[pos], buf[pos + 1], buf[pos + 2], buf[pos + 3]]) as usize; - pos += 4; - if pos + block_len > buf.len() { - break; - } - if let Some(block) = PostingBlock::from_bytes(&buf[pos..pos + block_len]) { - blocks.push(block); - } - pos += block_len; + let len_end = pos.checked_add(4).filter(|end| *end <= buf.len())?; + let block_len = usize::try_from(u32::from_le_bytes([ + buf[pos], + buf[pos + 1], + buf[pos + 2], + buf[pos + 3], + ])) + .ok()?; + pos = len_end; + let block_end = pos.checked_add(block_len).filter(|end| *end <= buf.len())?; + blocks.push(PostingBlock::from_bytes(&buf[pos..block_end])?); + pos = block_end; } - blocks + Some(blocks) } #[cfg(test)] @@ -207,7 +206,16 @@ mod tests { #[test] fn rejects_huge_block_count_with_tiny_payload_before_allocation() { - assert!(decode_term_blocks(&u32::MAX.to_le_bytes()).is_empty()); + assert!(decode_term_blocks(&u32::MAX.to_le_bytes()).is_none()); + } + + #[test] + fn a_truncated_block_list_does_not_decode() { + // Two blocks named, one empty block present. + let mut buf = 2u32.to_le_bytes().to_vec(); + buf.extend_from_slice(&0u32.to_le_bytes()); + buf.extend_from_slice(&0u32.to_le_bytes()); + assert!(decode_term_blocks(&buf).is_none()); } #[test] @@ -216,7 +224,7 @@ mod tests { let reader = SegmentReader::open(seg_data).unwrap(); assert_eq!(reader.num_terms(), 2); - let blocks = reader.read_postings("alpha"); + let blocks = reader.read_postings("alpha").unwrap(); assert_eq!(blocks.len(), 1); // 2 docs fit in 1 block. assert_eq!( blocks[0].doc_ids, @@ -226,11 +234,14 @@ mod tests { } #[test] - fn overflowing_posting_offset_returns_empty() { + fn overflowing_posting_offset_is_corrupt() { let seg_data = make_segment(); let mut reader = SegmentReader::open(seg_data).unwrap(); reader.term_dict[0].posting_offset = u64::MAX; - assert!(reader.read_postings("alpha").is_empty()); + assert!(matches!( + reader.read_postings("alpha"), + Err(SegmentError::CorruptPostings { ref term }) if term == "alpha" + )); } #[test] @@ -249,7 +260,7 @@ mod tests { fn missing_term_returns_empty() { let seg_data = make_segment(); let reader = SegmentReader::open(seg_data).unwrap(); - assert!(reader.read_postings("nonexistent").is_empty()); + assert!(reader.read_postings("nonexistent").unwrap().is_empty()); } #[test] diff --git a/nodedb-fts/src/lsm/segment/writer.rs b/nodedb-fts/src/lsm/segment/writer.rs index 6ee3dca55..2952bb74b 100644 --- a/nodedb-fts/src/lsm/segment/writer.rs +++ b/nodedb-fts/src/lsm/segment/writer.rs @@ -14,15 +14,23 @@ use super::format::{self, TermDictEntry}; /// Flush a memtable's postings into an immutable segment byte buffer. /// -/// `term_postings` is the drained HashMap from `Memtable::drain()`. +/// `term_postings` is one index's drained map from `Memtable::drain_scope()`. /// Returns the serialized segment bytes or a `SegmentError` if any term /// exceeds `MAX_TERM_LEN`. pub fn flush_to_segment( term_postings: HashMap>, +) -> Result, SegmentError> { + flush_postings_to_segment(&term_postings) +} + +/// Build a segment from a borrowed term→postings map. The map stays intact, +/// so a caller can drop it only after the segment is durably written. +pub fn flush_postings_to_segment( + term_postings: &HashMap>, ) -> Result, SegmentError> { // Sort terms for binary-searchable term dictionary. - let mut sorted_terms: Vec<(String, Vec)> = term_postings.into_iter().collect(); - sorted_terms.sort_by(|(a, _), (b, _)| a.cmp(b)); + let mut sorted_terms: Vec<(&String, &Vec)> = term_postings.iter().collect(); + sorted_terms.sort_by_key(|(term, _)| *term); // Phase 1: Encode posting blocks for each term, collect byte offsets. let mut posting_data = Vec::new(); @@ -41,7 +49,7 @@ pub fn flush_to_segment( let df = postings.len() as u32; // Split into 128-doc blocks and serialize each. - let blocks = into_blocks(postings.clone()); + let blocks = into_blocks(postings.to_vec()); let mut term_bytes = Vec::new(); // Write number of blocks. @@ -57,7 +65,7 @@ pub fn flush_to_segment( posting_data.extend_from_slice(&term_bytes); dict_entries.push(TermDictEntry { - term: term.clone(), + term: term.to_string(), posting_offset: offset, posting_len, df, @@ -124,7 +132,7 @@ pub fn build_from_blocks( posting_data.extend_from_slice(&term_bytes); dict_entries.push(TermDictEntry { - term: term.clone(), + term: term.to_string(), posting_offset: offset, posting_len, df, diff --git a/nodedb-fts/src/lsm/segment_deletes.rs b/nodedb-fts/src/lsm/segment_deletes.rs new file mode 100644 index 000000000..c108da05e --- /dev/null +++ b/nodedb-fts/src/lsm/segment_deletes.rs @@ -0,0 +1,150 @@ +// SPDX-License-Identifier: Apache-2.0 + +//! Per-segment delete sets of one index. +//! +//! A segment is immutable. Removing a document after its postings reached a +//! segment records the document in the delete set of every segment the index +//! holds at that moment. A query skips each segment's postings of its deleted +//! documents. A merge drops them, and the merged segment starts with no +//! deletes. A document indexed again after its removal lands in the memtable +//! and then in a newer segment, which no older delete set names. +//! +//! The sets persist in the index's backend metadata under +//! [`SEGMENT_DELETES_META_KEY`]. A set whose segment no longer exists is +//! ignored on read and dropped on the next write. + +use std::collections::BTreeMap; + +use nodedb_types::{Surrogate, SurrogateBitmap}; + +use crate::backend::FtsBackend; +use crate::index::FtsIndexError; +use crate::scope::IndexScope; + +/// Metadata sub-key of an index's segment delete sets. +pub const SEGMENT_DELETES_META_KEY: &str = "segment_deletes"; + +/// Segment id → the documents removed after that segment was written. +#[derive(Debug, Clone, Default, PartialEq)] +pub struct SegmentDeletes { + sets: BTreeMap, +} + +impl SegmentDeletes { + /// The stored delete sets of `index`, restricted to `live_segments`. + /// A stored blob that does not decode is a typed corruption error. + pub fn load( + backend: &B, + database_id: u64, + tid: u64, + index: IndexScope<'_>, + live_segments: &[String], + ) -> Result> { + let Some(bytes) = backend + .read_meta(database_id, tid, index, SEGMENT_DELETES_META_KEY) + .map_err(FtsIndexError::Backend)? + else { + return Ok(Self::default()); + }; + let entries: Vec<(String, SurrogateBitmap)> = + zerompk::from_msgpack(&bytes).map_err(|e| FtsIndexError::CorruptState { + subkey: SEGMENT_DELETES_META_KEY, + detail: e.to_string(), + })?; + let sets = entries + .into_iter() + .filter(|(segment_id, _)| live_segments.contains(segment_id)) + .collect(); + Ok(Self { sets }) + } + + /// Persist these delete sets as the delete sets of `index`. + pub fn store( + &self, + backend: &B, + database_id: u64, + tid: u64, + index: IndexScope<'_>, + ) -> Result<(), FtsIndexError> { + let entries: Vec<(&String, &SurrogateBitmap)> = self.sets.iter().collect(); + let bytes = zerompk::to_msgpack_vec(&entries).map_err(|e| FtsIndexError::StateEncode { + subkey: SEGMENT_DELETES_META_KEY, + detail: e.to_string(), + })?; + backend + .write_meta(database_id, tid, index, SEGMENT_DELETES_META_KEY, &bytes) + .map_err(FtsIndexError::Backend) + } + + /// Record `doc_id` as removed from each of `segments`. + pub fn mark_removed(&mut self, segments: &[String], doc_id: Surrogate) { + for segment_id in segments { + self.sets + .entry(segment_id.clone()) + .or_default() + .insert(doc_id); + } + } + + /// The documents removed from `segment_id`. `None` when it has none. + pub fn deleted(&self, segment_id: &str) -> Option<&SurrogateBitmap> { + self.sets.get(segment_id).filter(|set| !set.is_empty()) + } + + /// Move out the documents removed from `segment_id`. `None` when it has + /// none. + pub fn take(&mut self, segment_id: &str) -> Option { + self.sets.remove(segment_id).filter(|set| !set.is_empty()) + } + + /// Whether no segment has a removed document. + pub fn is_empty(&self) -> bool { + self.sets.values().all(SurrogateBitmap::is_empty) + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::backend::memory::MemoryBackend; + + const DB: u64 = 0; + const T: u64 = 1; + const DOCS: IndexScope<'static> = IndexScope::document("docs"); + + fn ids(names: &[&str]) -> Vec { + names.iter().map(|s| (*s).to_string()).collect() + } + + #[test] + fn delete_sets_round_trip_and_drop_removed_segments() { + let backend = MemoryBackend::new(); + let mut deletes = SegmentDeletes::default(); + deletes.mark_removed(&ids(&["L0:1", "L0:2"]), Surrogate(7)); + deletes.store(&backend, DB, T, DOCS).unwrap(); + + let loaded = SegmentDeletes::load(&backend, DB, T, DOCS, &ids(&["L0:1", "L0:2"])).unwrap(); + assert_eq!(loaded, deletes); + assert!( + loaded + .deleted("L0:1") + .is_some_and(|s| s.contains(Surrogate(7))) + ); + + // A merged-away segment's set is not read back. + let loaded = SegmentDeletes::load(&backend, DB, T, DOCS, &ids(&["L0:2"])).unwrap(); + assert!(loaded.deleted("L0:1").is_none()); + assert!(loaded.deleted("L0:2").is_some()); + assert!(loaded.deleted("L1:3").is_none()); + } + + #[test] + fn an_undecodable_blob_is_a_corruption_error() { + let backend = MemoryBackend::new(); + backend + .write_meta(DB, T, DOCS, SEGMENT_DELETES_META_KEY, &[0xc1, 0xff]) + .unwrap(); + let err = SegmentDeletes::load(&backend, DB, T, DOCS, &ids(&["L0:1"])).unwrap_err(); + assert!(matches!(err, FtsIndexError::CorruptState { .. }), "{err}"); + } +} diff --git a/nodedb-fts/src/posting.rs b/nodedb-fts/src/posting.rs index 0d5af1aa9..897ea12c9 100644 --- a/nodedb-fts/src/posting.rs +++ b/nodedb-fts/src/posting.rs @@ -36,6 +36,15 @@ pub enum QueryMode { Or, } +impl From for QueryMode { + fn from(mode: nodedb_types::text_search::QueryMode) -> Self { + match mode { + nodedb_types::text_search::QueryMode::And => Self::And, + nodedb_types::text_search::QueryMode::Or => Self::Or, + } + } +} + /// A scored search result from the inverted index. #[derive(Debug, Clone)] pub struct TextSearchResult { @@ -77,6 +86,13 @@ mod tests { assert_eq!(QueryMode::default(), QueryMode::And); } + #[test] + fn public_query_mode_maps_to_the_same_mode() { + use nodedb_types::text_search::QueryMode as Public; + assert_eq!(QueryMode::from(Public::And), QueryMode::And); + assert_eq!(QueryMode::from(Public::Or), QueryMode::Or); + } + #[test] fn default_bm25_params() { let p = Bm25Params::default(); diff --git a/nodedb-fts/src/scope.rs b/nodedb-fts/src/scope.rs new file mode 100644 index 000000000..2e74a6450 --- /dev/null +++ b/nodedb-fts/src/scope.rs @@ -0,0 +1,98 @@ +// SPDX-License-Identifier: Apache-2.0 + +//! The inverted index a read or write addresses: a collection's +//! whole-document text, or one of its fields. + +/// One inverted index of a collection. `field` is empty for the +/// whole-document index. A field index always has a non-empty name, so no +/// field can address the whole-document index. +#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)] +pub struct IndexScope<'a> { + collection: &'a str, + field: &'a str, +} + +impl<'a> IndexScope<'a> { + /// The whole-document index of `collection`. + pub const fn document(collection: &'a str) -> Self { + Self { + collection, + field: "", + } + } + + /// The index of `field`. `None` for an empty name: it has no index of its own. + pub fn field(collection: &'a str, field: &'a str) -> Option { + (!field.is_empty()).then_some(Self { collection, field }) + } + + /// The collection that owns this index. + pub const fn collection(&self) -> &'a str { + self.collection + } + + /// The field name. `None` for the whole-document index. + pub fn field_name(&self) -> Option<&'a str> { + (!self.field.is_empty()).then_some(self.field) + } + + /// Storage key component: empty for the whole-document index. + pub const fn field_key(&self) -> &'a str { + self.field + } + + /// The scope a stored `(collection, field_key)` pair names. + pub const fn from_key(collection: &'a str, field_key: &'a str) -> Self { + Self { + collection, + field: field_key, + } + } +} + +impl<'a> From<&'a str> for IndexScope<'a> { + fn from(collection: &'a str) -> Self { + Self::document(collection) + } +} + +impl<'a> From<&'a String> for IndexScope<'a> { + fn from(collection: &'a String) -> Self { + Self::document(collection.as_str()) + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn empty_field_name_has_no_scope() { + assert_eq!(IndexScope::field("docs", ""), None); + } + + #[test] + fn document_and_field_scopes_never_alias() { + let doc = IndexScope::document("docs"); + let field = IndexScope::field("docs", "title").expect("non-empty field"); + assert_ne!(doc, field); + assert_eq!(doc.field_name(), None); + assert_eq!(doc.field_key(), ""); + assert_eq!(field.field_name(), Some("title")); + assert_eq!(field.collection(), "docs"); + } + + #[test] + fn from_key_round_trips_both_kinds() { + let doc = IndexScope::document("docs"); + assert_eq!(IndexScope::from_key("docs", doc.field_key()), doc); + let field = IndexScope::field("docs", "body").expect("non-empty field"); + assert_eq!(IndexScope::from_key("docs", field.field_key()), field); + } + + #[test] + fn str_converts_to_document_scope() { + let scope: IndexScope<'_> = "docs".into(); + assert_eq!(scope, IndexScope::document("docs")); + } +} diff --git a/nodedb-fts/src/search/bm25_search.rs b/nodedb-fts/src/search/bm25_search.rs index 2a5e27541..1dac21089 100644 --- a/nodedb-fts/src/search/bm25_search.rs +++ b/nodedb-fts/src/search/bm25_search.rs @@ -1,41 +1,44 @@ // SPDX-License-Identifier: Apache-2.0 -//! BM25 search over the FtsIndex with AND-first OR-fallback, phrase boost, -//! and NOT-term exclusion. +//! BM25 search over the FtsIndex with AND-first OR-fallback and NOT-term +//! exclusion. //! //! ## AND-first / OR-fallback //! -//! A multi-term query `rust programming` first attempts AND (all terms must -//! match). If no document satisfies all terms, results fall back to OR with -//! coverage-penalised scores. +//! A multi-word query `rust programming` in AND mode matches documents +//! holding every word (a word matches through any of its synonyms). The +//! top-k is the true top-k of those documents, whatever `top_k` is. When no +//! admitted document holds every word, the query falls back to OR with +//! coverage-scaled scores. //! //! ## NOT operator //! //! `rust NOT python` and `rust -python` are equivalent. The query parser //! splits the input into positive and negative term lists. BM25 scoring runs -//! on positive terms only; a bitmap of doc IDs that match **any** negative -//! term is built separately and used to filter final results. Negative terms -//! do not affect BM25 scores. +//! on positive terms only. A document holding any negative term is excluded +//! before scoring, so the top-k cut counts only surviving documents. +//! Negative terms do not affect BM25 scores. //! //! Synonym expansion applies to both positive and negative term lists, so //! `rust NOT db` also excludes documents that contain synonym expansions of //! `db` (e.g. `database`, `datastore`). -use std::collections::HashMap; - -use nodedb_types::{Surrogate, SurrogateBitmap}; +use nodedb_types::SurrogateBitmap; use crate::backend::FtsBackend; -use crate::bm25::bm25_score; use crate::index::FtsIndex; use crate::index::error::FtsIndexError; -use crate::posting::{Posting, QueryMode, TextSearchResult}; -use crate::search::phrase; +use crate::posting::{QueryMode, TextSearchResult}; +use crate::scope::IndexScope; +use crate::search::bmw::scorer::{BmwInput, bmw_score}; +use crate::search::match_mode::staged_candidates; use crate::search::query_parser::parse_query; +use crate::search::query_terms::TextQuery; +use crate::search::staged::StagedView; /// Query and tuning parameters for a BM25 search. /// -/// The `(database_id, tid, collection)` scope is passed separately so the +/// The `(database_id, tid, index)` scope is passed separately so the /// same struct can be shared by callers that hold the tenant id as either a /// raw `u64` (this crate) or a strongly-typed `TenantId` (the Origin wrapper). pub struct FtsSearchParams<'a> { @@ -51,30 +54,40 @@ pub struct FtsSearchParams<'a> { pub prefilter: Option<&'a SurrogateBitmap>, } -/// Inputs to the AND-mode post-filter that drops BMW candidates which do not -/// match at least `num_terms` of the analyzed query tokens. -struct FilterAndModeParams<'a> { - database_id: u64, - tid: u64, - collection: &'a str, - query_tokens: &'a [String], - candidates: &'a [TextSearchResult], - num_terms: usize, -} - impl FtsIndex { - /// Search the index with explicit boolean mode, fuzzy, and optional prefilter. + /// Search one index with explicit boolean mode, fuzzy, and optional prefilter. + /// + /// Analyzer, fuzzy default, and synonyms come from the index's + /// collection. Postings and BM25 stats come from the index itself. /// /// Supports `NOT ` and `-` negation in the query string. /// Returns `Err(FtsIndexError::InvalidQuery)` for ill-formed queries such /// as NOT-only queries or unsupported parenthesised groups. - pub fn search( + pub fn search<'a>( &self, database_id: u64, tid: u64, - collection: &str, + index: impl Into>, params: FtsSearchParams<'_>, ) -> Result, FtsIndexError> { + self.search_staged(database_id, tid, index, params, None) + } + + /// [`Self::search`] inside an open transaction: indexed documents the + /// transaction hides do not match, and its staged documents compete in + /// the same ranking under the same query semantics. `prefilter` bounds + /// staged documents too. + /// + /// Results are ordered by score descending, then surrogate ascending. + pub fn search_staged<'a>( + &self, + database_id: u64, + tid: u64, + index: impl Into>, + params: FtsSearchParams<'_>, + staged: Option<&StagedView>, + ) -> Result, FtsIndexError> { + let index = index.into(); let FtsSearchParams { query, top_k, @@ -82,392 +95,71 @@ impl FtsIndex { mode, prefilter, } = params; - // A collection configured with `FUZZY true` falls back to fuzzy - // matching even when the query did not ask for it — that is what - // makes it an index property rather than a per-query flag. Resolved - // here, at the one point every search path funnels through, so no - // caller can be wired up without it. - let fuzzy_enabled = fuzzy_enabled - || self - .get_collection_fuzzy(database_id, tid, collection) - .map_err(FtsIndexError::backend)?; - - // Parse the query for NOT / - negation operators before analysis. - let parsed = parse_query(query)?; - - // Reconstruct the positive-only query string for the existing analyzer path. - // Each raw positive token is passed to the analyzer individually rather than - // joining them, because some analyzers are sensitive to token boundaries. - // Joining with a space is safe for the standard/language analyzers. - let positive_raw = parsed.positive.join(" "); - let negative_raw_terms = parsed.negative; - - let base_tokens = self - .analyze_for_collection(database_id, tid, collection, &positive_raw) - .map_err(FtsIndexError::backend)?; - if base_tokens.is_empty() { + if top_k == 0 { + // The query is still parsed: an ill-formed one is an error at + // any limit. + parse_query(query)?; return Ok(Vec::new()); } - - let base_token_count = base_tokens.len(); - let query_tokens = self - .expand_query_with_synonyms(database_id, tid, base_tokens) - .map_err(FtsIndexError::backend)?; - let num_query_terms = query_tokens.len(); - let and_threshold = base_token_count; - - let raw_tokens = if fuzzy_enabled { - self.tokenize_raw_for_collection(database_id, tid, collection, &positive_raw) - .map_err(FtsIndexError::backend)? - } else { - Vec::new() + let text_query = TextQuery { + query, + fuzzy_enabled, + mode, }; - - let (total_docs, avg_doc_len) = self - .index_stats(database_id, tid, collection) - .map_err(FtsIndexError::backend)?; - if total_docs == 0 { + let Some(resolved) = self.resolve_query(database_id, tid, index, text_query, staged)? + else { return Ok(Vec::new()); - } - - // Build the negative-term exclusion set before scoring. - // Negative terms are analyzed and synonym-expanded just like positive - // terms. The result is a set of doc IDs that match any negative term. - let negative_set = - self.build_negative_set(database_id, tid, collection, &negative_raw_terms)?; - - let bmw_params = super::bmw::query::BmwParams { - query_tokens: &query_tokens, - raw_tokens: &raw_tokens, - fuzzy_enabled, - top_k: if mode == QueryMode::And && and_threshold > 1 { - top_k.saturating_mul(3).max(20) - } else { - top_k - }, - total_docs, - avg_doc_len, - bm25: &self.bm25_params, - prefilter, }; - if let Ok(Some(bmw_results)) = - super::bmw::query::bmw_search(self, database_id, tid, collection, &bmw_params) - { - if mode == QueryMode::Or || and_threshold == 1 { - let mut results: Vec = bmw_results - .into_iter() - .filter(|r| !negative_set.contains(&r.doc_id)) - .take(top_k) - .collect(); - results.truncate(top_k); - return Ok(results); - } - - let and_results = self - .filter_and_mode(FilterAndModeParams { - database_id, - tid, - collection, - query_tokens: &query_tokens, - candidates: &bmw_results, - num_terms: and_threshold, - }) - .map_err(FtsIndexError::backend)?; - - if !and_results.is_empty() { - let filtered: Vec = and_results - .into_iter() - .filter(|r| !negative_set.contains(&r.doc_id)) - .take(top_k) - .collect(); - return Ok(filtered); - } - - let penalized: Vec = bmw_results - .into_iter() - .filter(|r| !negative_set.contains(&r.doc_id)) - .map(|mut r| { - let matched = self.count_term_matches( - database_id, - tid, - collection, - &query_tokens, - r.doc_id, - ); - let coverage = matched as f32 / and_threshold as f32; - r.score *= coverage; - r - }) - .collect(); - let mut sorted = penalized; - sorted.sort_by(|a, b| { - b.score - .partial_cmp(&a.score) - .unwrap_or(std::cmp::Ordering::Equal) - }); - sorted.truncate(top_k); - return Ok(sorted); - } - - // Fallback: exhaustive BM25 scoring reading directly from the backend. - let term_postings_memory = crate::mem_scope::fts_scope(&self.governor, database_id, tid); - let term_postings_bytes = - num_query_terms * (std::mem::size_of::>() + std::mem::size_of::()); - let _term_postings_guard = term_postings_memory.reserve(term_postings_bytes).ok(); - let mut term_postings: Vec<(Vec, bool)> = Vec::with_capacity(num_query_terms); - for (i, token) in query_tokens.iter().enumerate() { - let postings = self - .backend - .read_postings(database_id, tid, collection, token) - .map_err(FtsIndexError::backend)?; - if !postings.is_empty() { - term_postings.push((postings, false)); - } else if fuzzy_enabled { - let raw = raw_tokens - .get(i) - .map(String::as_str) - .unwrap_or(token.as_str()); - let (fuzzy_posts, is_fuzzy) = self - .fuzzy_lookup(database_id, tid, collection, raw) - .map_err(FtsIndexError::backend)?; - term_postings.push((fuzzy_posts, is_fuzzy)); - } else { - term_postings.push((Vec::new(), false)); - } - } - - let mut doc_scores: HashMap = HashMap::new(); - - for (token_idx, (postings, is_fuzzy)) in term_postings.iter().enumerate() { - if postings.is_empty() { - continue; - } - let df = postings.len() as u32; - - for posting in postings { - // Prefilter: skip surrogates not present in the bitmap. - if let Some(bm) = prefilter - && !bm.contains(posting.doc_id) - { - continue; - } - - let doc_len = self - .backend - .read_doc_length(database_id, tid, collection, posting.doc_id) - .map_err(FtsIndexError::backend)? - .unwrap_or(1); - - let mut score = bm25_score( - posting.term_freq, - df, - doc_len, - total_docs, - avg_doc_len, - &self.bm25_params, - ); - - if *is_fuzzy { - score *= crate::fuzzy::fuzzy_discount(1); - } - - let entry = doc_scores.entry(posting.doc_id).or_insert((0.0, false, 0)); - entry.0 += score; - if *is_fuzzy { - entry.1 = true; - } - entry.2 += 1; - } - let _ = token_idx; - } - if num_query_terms >= 2 { - let doc_postings_map = phrase::collect_doc_postings(&query_tokens, &term_postings); - for (doc_id, token_postings) in &doc_postings_map { - if let Some(entry) = doc_scores.get_mut(doc_id) { - let boost = phrase::phrase_boost(&query_tokens, token_postings); - entry.0 *= boost; - } - } + let mut deny = resolved.negated.clone(); + if let Some(view) = staged { + deny.union_in_place(view.hidden()); } + let index_visible = staged.is_none_or(|view| !view.hides_all()); + let staged_docs = staged_candidates(&resolved, staged, prefilter, &self.bm25_params); + let match_mode = resolved.match_mode(index_visible, prefilter, &deny, &staged_docs); + let allow = match &match_mode { + super::match_mode::MatchMode::All(docs) => Some(docs), + super::match_mode::MatchMode::Coverage | super::match_mode::MatchMode::Any => prefilter, + }; - if mode == QueryMode::And && and_threshold > 1 { - let and_results: HashMap = doc_scores - .iter() - .filter(|(_, (_, _, match_count))| *match_count >= and_threshold) - .map(|(k, v)| (*k, *v)) - .collect(); - - if !and_results.is_empty() { - let filtered = and_results - .into_iter() - .filter(|(doc_id, _)| !negative_set.contains(doc_id)) - .collect(); - return Ok(Self::to_sorted_results(filtered, top_k)); - } - - for (score, _, match_count) in doc_scores.values_mut() { - let coverage = *match_count as f32 / and_threshold as f32; - *score *= coverage; - } - } - - // Apply negative filter to final fallback results. - let filtered: HashMap = doc_scores - .into_iter() - .filter(|(doc_id, _)| !negative_set.contains(doc_id)) - .collect(); - - Ok(Self::to_sorted_results(filtered, top_k)) - } - - /// Build a set of doc IDs that match any of the given raw negative terms. - /// - /// Each raw negative term is analyzed and synonym-expanded before posting - /// lookup, matching the same pipeline as positive terms. - fn build_negative_set( - &self, - database_id: u64, - tid: u64, - collection: &str, - raw_negative_terms: &[String], - ) -> Result, FtsIndexError> { - if raw_negative_terms.is_empty() { - return Ok(std::collections::HashSet::new()); - } - - // Analyze all negative terms together (join is safe for standard analyzer). - let neg_raw = raw_negative_terms.join(" "); - let neg_base_tokens = self - .analyze_for_collection(database_id, tid, collection, &neg_raw) - .map_err(FtsIndexError::backend)?; - - if neg_base_tokens.is_empty() { - return Ok(std::collections::HashSet::new()); - } - - // Synonym-expand negative tokens so negating 'db' also excludes 'database'. - let neg_tokens = self - .expand_query_with_synonyms(database_id, tid, neg_base_tokens) - .map_err(FtsIndexError::backend)?; - - let mut excluded: std::collections::HashSet = std::collections::HashSet::new(); - - // Collect postings from memtable + segments for each negative token. - let term_blocks = crate::lsm::query::collect_merged_term_blocks( - &self.backend, - database_id, - tid, - collection, - self.memtable(), - &neg_tokens, - &self.governor, - ) - .map_err(FtsIndexError::backend)?; - - for tb in &term_blocks { - for block in &tb.blocks { - for doc_id in &block.doc_ids { - excluded.insert(*doc_id); - } - } + let mut hits: Vec = Vec::new(); + if index_visible { + let heap = bmw_score(&BmwInput { + terms: &resolved.blocks, + term_groups: &resolved.term_groups, + groups: resolved.groups, + total_docs: resolved.total_docs, + avg_doc_len: resolved.avg_doc_len, + params: &self.bm25_params, + top_k, + allow, + deny: Some(&deny), + combine: match_mode.combine(), + }); + hits.extend(heap.into_sorted().into_iter().map(|doc| TextSearchResult { + doc_id: doc.doc_id, + score: doc.score, + fuzzy: resolved.fuzzy, + })); } - - // Also check the backend postings directly (covers the exhaustive path). - for token in &neg_tokens { - let postings = self - .backend - .read_postings(database_id, tid, collection, token) - .map_err(FtsIndexError::backend)?; - for posting in postings { - excluded.insert(posting.doc_id); + for (doc_id, contributions) in &staged_docs { + if let Some(score) = match_mode.score(&resolved, contributions) { + hits.push(TextSearchResult { + doc_id: *doc_id, + score, + fuzzy: resolved.fuzzy, + }); } } - - Ok(excluded) - } - - fn filter_and_mode( - &self, - params: FilterAndModeParams<'_>, - ) -> Result, B::Error> { - let FilterAndModeParams { - database_id, - tid, - collection, - query_tokens, - candidates, - num_terms, - } = params; - let term_blocks = crate::lsm::query::collect_merged_term_blocks( - &self.backend, - database_id, - tid, - collection, - self.memtable(), - query_tokens, - &self.governor, - )?; - - let mut results = Vec::new(); - for candidate in candidates { - let surrogate = candidate.doc_id; - let matched = term_blocks - .iter() - .filter(|tb| tb.blocks.iter().any(|b| b.doc_ids.contains(&surrogate))) - .count(); - if matched >= num_terms { - results.push(candidate.clone()); - } - } - Ok(results) - } - - fn count_term_matches( - &self, - database_id: u64, - tid: u64, - collection: &str, - query_tokens: &[String], - doc_id: Surrogate, - ) -> usize { - let term_blocks = match crate::lsm::query::collect_merged_term_blocks( - &self.backend, - database_id, - tid, - collection, - self.memtable(), - query_tokens, - &self.governor, - ) { - Ok(tb) => tb, - Err(_) => return 0, - }; - term_blocks - .iter() - .filter(|tb| tb.blocks.iter().any(|b| b.doc_ids.contains(&doc_id))) - .count() - } - - fn to_sorted_results( - doc_scores: HashMap, - top_k: usize, - ) -> Vec { - let mut results: Vec = doc_scores - .into_iter() - .map(|(doc_id, (score, fuzzy_flag, _))| TextSearchResult { - doc_id, - score, - fuzzy: fuzzy_flag, - }) - .collect(); - results.sort_by(|a, b| { + hits.sort_by(|a, b| { b.score .partial_cmp(&a.score) .unwrap_or(std::cmp::Ordering::Equal) + .then(a.doc_id.cmp(&b.doc_id)) }); - results.truncate(top_k); - results + hits.truncate(top_k); + Ok(hits) } } @@ -481,6 +173,7 @@ mod tests { use crate::index::error::FtsIndexError; use crate::posting::QueryMode; use crate::search::query_parser::InvalidQuery; + use crate::search::staged::StagedDoc; use crate::test_support::test_governor; const DB: u64 = 0; @@ -1125,4 +818,192 @@ mod tests { "expected InvalidQuery(ParenthesesNotSupported), got {err:?}" ); } + + // ── top-k semantics ─────────────────────────────────────────────────────── + + fn ranked( + idx: &FtsIndex, + query: &str, + top_k: usize, + mode: QueryMode, + prefilter: Option<&SurrogateBitmap>, + staged: Option<&super::StagedView>, + ) -> Vec { + idx.search_staged( + DB, + T, + "docs", + FtsSearchParams { + query, + top_k, + fuzzy_enabled: false, + mode, + prefilter, + }, + staged, + ) + .unwrap() + .into_iter() + .map(|r| r.doc_id) + .collect() + } + + /// The best `rust` documents hold `python`. Negation drops them before + /// the cut, so `LIMIT 1` still returns the surviving document. + #[test] + fn not_terms_are_excluded_before_the_limit() { + let idx = FtsIndex::new(MemoryBackend::new(), test_governor()); + idx.index_document(DB, T, "docs", D1, "rust rust rust python") + .unwrap(); + idx.index_document(DB, T, "docs", D2, "rust rust rust python") + .unwrap(); + idx.index_document(DB, T, "docs", D3, "rust golang compiler toolchain") + .unwrap(); + assert_eq!( + ranked(&idx, "rust -python", 1, QueryMode::And, None, None), + vec![D3] + ); + assert_eq!( + ranked(&idx, "rust -python", 2, QueryMode::Or, None, None), + vec![D3] + ); + } + + /// Many documents out-score the single AND match on one word. The AND + /// match is found whatever the limit, and the query does not fall back. + #[test] + fn and_match_is_found_past_many_single_word_out_scorers() { + let idx = FtsIndex::new(MemoryBackend::new(), test_governor()); + for i in 1..=40u32 { + idx.index_document(DB, T, "docs", Surrogate(i), "alpha alpha alpha alpha") + .unwrap(); + } + for i in 41..=80u32 { + idx.index_document(DB, T, "docs", Surrogate(i), "bravo bravo bravo bravo") + .unwrap(); + } + let both = Surrogate(81); + idx.index_document(DB, T, "docs", both, "alpha bravo filler words here") + .unwrap(); + for limit in [1, 3, 10, usize::MAX] { + assert_eq!( + ranked(&idx, "alpha bravo", limit, QueryMode::And, None, None), + vec![both], + "limit {limit}" + ); + } + } + + /// The AND decision reads only the admitted rows. An AND match outside + /// the prefilter does not stop the fallback inside it, and an AND match + /// inside it keeps AND semantics. + #[test] + fn and_mode_inside_a_prefilter() { + let idx = FtsIndex::new(MemoryBackend::new(), test_governor()); + idx.index_document(DB, T, "docs", D1, "alpha bravo") + .unwrap(); + idx.index_document(DB, T, "docs", D2, "alpha charlie") + .unwrap(); + idx.index_document(DB, T, "docs", D3, "bravo delta") + .unwrap(); + + let mut without_match = SurrogateBitmap::new(); + without_match.insert(D2); + without_match.insert(D3); + let mut fallback = ranked( + &idx, + "alpha bravo", + 10, + QueryMode::And, + Some(&without_match), + None, + ); + fallback.sort(); + assert_eq!(fallback, vec![D2, D3], "no admitted AND match: OR fallback"); + + let mut with_match = without_match.clone(); + with_match.insert(D1); + assert_eq!( + ranked( + &idx, + "alpha bravo", + 10, + QueryMode::And, + Some(&with_match), + None + ), + vec![D1], + "an admitted AND match keeps AND semantics" + ); + } + + /// A staged update that removes the best hit leaves the limit filled by + /// the next document, and a staged document competes in the ranking. + #[test] + fn staged_rows_are_ranked_before_the_cut() { + let idx = FtsIndex::new(MemoryBackend::new(), test_governor()); + idx.index_document(DB, T, "docs", D1, "rust rust rust rust") + .unwrap(); + idx.index_document(DB, T, "docs", D2, "rust rust lang") + .unwrap(); + idx.index_document(DB, T, "docs", D3, "rust tooling compiler words") + .unwrap(); + + let mut hidden = SurrogateBitmap::new(); + hidden.insert(D1); + let removed_top = super::StagedView::new( + hidden, + false, + vec![StagedDoc { + doc_id: D1, + tokens: vec!["python".into()], + }], + ); + assert_eq!( + ranked(&idx, "rust", 2, QueryMode::And, None, Some(&removed_top)), + vec![D2, D3] + ); + + let new_best = super::StagedView::new( + SurrogateBitmap::new(), + false, + vec![StagedDoc { + doc_id: Surrogate(9), + tokens: vec!["rust".into(); 6], + }], + ); + assert_eq!( + ranked(&idx, "rust", 1, QueryMode::And, None, Some(&new_best))[0], + Surrogate(9) + ); + } + + /// A staged document is scored with AND semantics: holding one of two + /// words does not match while an AND match exists. + #[test] + fn staged_rows_use_and_semantics() { + let idx = FtsIndex::new(MemoryBackend::new(), test_governor()); + idx.index_document(DB, T, "docs", D1, "alpha bravo") + .unwrap(); + let view = super::StagedView::new( + SurrogateBitmap::new(), + false, + vec![ + StagedDoc { + doc_id: Surrogate(8), + tokens: vec!["alpha".into(), "alpha".into()], + }, + StagedDoc { + doc_id: Surrogate(9), + tokens: vec!["alpha".into(), "bravo".into()], + }, + ], + ); + let mut hits = ranked(&idx, "alpha bravo", 10, QueryMode::And, None, Some(&view)); + hits.sort(); + assert_eq!(hits, vec![D1, Surrogate(9)]); + + let negated = ranked(&idx, "alpha -bravo", 10, QueryMode::And, None, Some(&view)); + assert_eq!(negated, vec![Surrogate(8)]); + } } diff --git a/nodedb-fts/src/search/bmw/heap.rs b/nodedb-fts/src/search/bmw/heap.rs index 4bb6c3c97..0f81c6a0d 100644 --- a/nodedb-fts/src/search/bmw/heap.rs +++ b/nodedb-fts/src/search/bmw/heap.rs @@ -4,6 +4,12 @@ //! //! Maintains the k best `(score, surrogate)` pairs. The threshold (minimum //! score to enter the heap) is the root's score when full, 0.0 when filling. +//! +//! Rank order is score descending, then surrogate ascending. The kept set +//! is the first k documents in that order, so a cut at k is a prefix of the +//! uncut result, ties included. + +use std::cmp::Ordering; use nodedb_types::Surrogate; @@ -14,20 +20,31 @@ pub struct ScoredDoc { pub doc_id: Surrogate, } -/// Fixed-capacity min-heap: smallest score is at the root. +/// `Greater` when `a` ranks above `b`: a higher score, or an equal score and +/// a lower surrogate. +fn rank_cmp(a: &ScoredDoc, b: &ScoredDoc) -> Ordering { + a.score + .total_cmp(&b.score) + .then_with(|| b.doc_id.cmp(&a.doc_id)) +} + +/// Fixed-capacity min-heap: the lowest-ranked candidate is at the root. /// -/// When full, only candidates exceeding the root's score are admitted -/// (the root is replaced and the heap is sifted down). +/// When full, only a candidate that ranks above the root is admitted (the +/// root is replaced and the heap is sifted down). pub struct TopKHeap { data: Vec, capacity: usize, } +/// Initial heap allocation. A larger `k` grows on demand; `usize::MAX` means every match. +const INITIAL_HEAP_CAPACITY: usize = 1024; + impl TopKHeap { - /// Create a new heap with the given capacity (k). + /// Create a new heap that keeps the best `k` candidates. pub fn new(k: usize) -> Self { Self { - data: Vec::with_capacity(k), + data: Vec::with_capacity(k.min(INITIAL_HEAP_CAPACITY)), capacity: k, } } @@ -44,28 +61,26 @@ impl TopKHeap { /// Try to insert a scored document. /// - /// If the heap is not full, always inserts. If full, only inserts - /// if `score > threshold()`, replacing the root. + /// If the heap is not full, always inserts. If full, inserts only when + /// the candidate ranks above the root, replacing the root. pub fn insert(&mut self, score: f32, doc_id: Surrogate) { + let candidate = ScoredDoc { score, doc_id }; if self.data.len() < self.capacity { - self.data.push(ScoredDoc { score, doc_id }); + self.data.push(candidate); if self.data.len() == self.capacity { // Build the min-heap once full. self.build_heap(); } - } else if score > self.data[0].score { - self.data[0] = ScoredDoc { score, doc_id }; + } else if rank_cmp(&candidate, &self.data[0]) == Ordering::Greater { + self.data[0] = candidate; self.sift_down(0); } } - /// Drain the heap into a sorted vec (descending by score). + /// Drain the heap into a vec in rank order: score descending, then + /// surrogate ascending. pub fn into_sorted(mut self) -> Vec { - self.data.sort_by(|a, b| { - b.score - .partial_cmp(&a.score) - .unwrap_or(std::cmp::Ordering::Equal) - }); + self.data.sort_by(|a, b| rank_cmp(b, a)); self.data } @@ -93,10 +108,10 @@ impl TopKHeap { let right = 2 * pos + 2; let mut smallest = pos; - if left < n && self.data[left].score < self.data[smallest].score { + if left < n && rank_cmp(&self.data[left], &self.data[smallest]) == Ordering::Less { smallest = left; } - if right < n && self.data[right].score < self.data[smallest].score { + if right < n && rank_cmp(&self.data[right], &self.data[smallest]) == Ordering::Less { smallest = right; } @@ -161,6 +176,36 @@ mod tests { assert_eq!(heap.threshold(), 0.0); } + #[test] + fn unbounded_k_allocates_at_most_the_initial_capacity() { + let mut heap = TopKHeap::new(usize::MAX); + assert!(heap.data.capacity() <= INITIAL_HEAP_CAPACITY); + for i in 1..=3000u32 { + heap.insert(i as f32, Surrogate(i)); + } + assert_eq!(heap.len(), 3000, "an unbounded heap keeps every match"); + assert_eq!(heap.into_sorted()[0].doc_id, Surrogate(3000)); + } + + /// A full heap of tied scores evicts the tie with the highest surrogate, + /// so the kept set is the first k rows of the uncut order. + #[test] + fn eviction_among_ties_keeps_the_lowest_surrogates() { + let mut heap = TopKHeap::new(2); + heap.insert(1.0, Surrogate(1)); + heap.insert(1.0, Surrogate(2)); + heap.insert(2.0, Surrogate(3)); + let kept: Vec = heap.into_sorted().iter().map(|d| d.doc_id).collect(); + assert_eq!(kept, vec![Surrogate(3), Surrogate(1)]); + + let mut heap = TopKHeap::new(3); + for id in [7, 4, 9, 2] { + heap.insert(1.0, Surrogate(id)); + } + let kept: Vec = heap.into_sorted().iter().map(|d| d.doc_id).collect(); + assert_eq!(kept, vec![Surrogate(2), Surrogate(4), Surrogate(7)]); + } + #[test] fn single_element() { let mut heap = TopKHeap::new(1); diff --git a/nodedb-fts/src/search/bmw/mod.rs b/nodedb-fts/src/search/bmw/mod.rs index f42b8e3fc..1ca9c39ef 100644 --- a/nodedb-fts/src/search/bmw/mod.rs +++ b/nodedb-fts/src/search/bmw/mod.rs @@ -1,6 +1,5 @@ // SPDX-License-Identifier: Apache-2.0 pub mod heap; -pub mod query; pub mod scorer; pub mod skip_index; diff --git a/nodedb-fts/src/search/bmw/query.rs b/nodedb-fts/src/search/bmw/query.rs deleted file mode 100644 index cb5c8ffa8..000000000 --- a/nodedb-fts/src/search/bmw/query.rs +++ /dev/null @@ -1,288 +0,0 @@ -// SPDX-License-Identifier: Apache-2.0 - -//! BMW query entry point: merges memtable + segments via LSM layer, -//! runs BMW scoring on `Surrogate` row identities. - -use nodedb_types::SurrogateBitmap; - -use crate::backend::FtsBackend; -use crate::block::{CompactPosting, into_blocks}; -use crate::codec::smallfloat; -use crate::index::FtsIndex; -use crate::lsm::query as lsm_query; -use crate::posting::{Bm25Params, Posting, TextSearchResult}; -use crate::search::bmw::skip_index::TermBlocks; - -use super::scorer::bmw_score; - -/// Corpus-level parameters for BMW search. -pub struct BmwParams<'a> { - pub query_tokens: &'a [String], - pub raw_tokens: &'a [String], - pub fuzzy_enabled: bool, - pub top_k: usize, - pub total_docs: u32, - pub avg_doc_len: f32, - pub bm25: &'a Bm25Params, - /// Optional surrogate prefilter. Only surrogates present in the bitmap - /// will be scored; all others are skipped before BM25 computation. - pub prefilter: Option<&'a SurrogateBitmap>, -} - -/// Run BMW search over the FtsIndex. -pub fn bmw_search( - index: &FtsIndex, - database_id: u64, - tid: u64, - collection: &str, - p: &BmwParams<'_>, -) -> Result>, B::Error> { - let mut has_fuzzy = vec![false; p.query_tokens.len()]; - - let mut lsm_term_blocks = lsm_query::collect_merged_term_blocks( - &index.backend, - database_id, - tid, - collection, - index.memtable(), - p.query_tokens, - &index.governor, - )?; - - let all_empty = lsm_term_blocks.iter().all(|tb| tb.df == 0); - // When the LSM has no entries for any token and fuzzy is disabled, skip the - // backend exact-lookup pass and return None immediately so the caller falls - // back to the non-BMW scoring path, which reads from the backend directly. - // When fuzzy is enabled (or when we may have backend-resident postings that - // are worth probing), we always enter the per-token resolution loop below. - if all_empty && !p.fuzzy_enabled { - // Even with no LSM postings, Origin may have postings stored directly in - // the redb backend (bypassing LSM). Check at least one token to decide. - let has_any_backend = p.query_tokens.iter().any(|tok| { - index - .backend - .read_postings(database_id, tid, collection, tok) - .ok() - .is_some_and(|v| !v.is_empty()) - }); - if !has_any_backend { - return Ok(None); - } - } - - for (i, token) in p.query_tokens.iter().enumerate() { - if lsm_term_blocks[i].df == 0 { - // First, try an exact backend lookup (covers Origin's redb-direct indexing - // path which bypasses the LSM memtable). - let backend_posts = index - .backend - .read_postings(database_id, tid, collection, token)?; - if !backend_posts.is_empty() { - let compact = to_compact(&backend_posts, index, database_id, tid, collection)?; - let blocks = into_blocks(compact); - lsm_term_blocks[i] = TermBlocks::from_blocks(blocks); - continue; - } - // Fall back to fuzzy matching when no exact posting exists. - if p.fuzzy_enabled { - let raw = p.raw_tokens.get(i).unwrap_or(token); - let (posts, is_fuzzy) = index.fuzzy_lookup(database_id, tid, collection, raw)?; - has_fuzzy[i] = is_fuzzy; - if !posts.is_empty() { - let compact = to_compact(&posts, index, database_id, tid, collection)?; - let blocks = into_blocks(compact); - lsm_term_blocks[i] = TermBlocks::from_blocks(blocks); - } - } - } - } - - if lsm_term_blocks.iter().all(|tb| tb.df == 0) { - return Ok(None); - } - - let all_term_blocks = lsm_term_blocks; - - let heap = bmw_score( - &all_term_blocks, - p.total_docs, - p.avg_doc_len, - p.bm25, - p.top_k, - p.prefilter, - ); - let scored = heap.into_sorted(); - - let is_fuzzy = has_fuzzy.iter().any(|&f| f); - let results: Vec = scored - .iter() - .map(|doc| TextSearchResult { - doc_id: doc.doc_id, - score: doc.score, - fuzzy: is_fuzzy, - }) - .collect(); - - Ok(Some(results)) -} - -/// Convert `Vec` → `Vec`, reading fieldnorms from the index. -fn to_compact( - postings: &[Posting], - index: &FtsIndex, - database_id: u64, - tid: u64, - collection: &str, -) -> Result, B::Error> { - let mut compact = Vec::with_capacity(postings.len()); - for p in postings { - let fieldnorm = index - .read_fieldnorm(database_id, tid, collection, p.doc_id)? - .map(smallfloat::encode) - .unwrap_or_else(|| smallfloat::encode(p.term_freq)); - compact.push(CompactPosting { - doc_id: p.doc_id, - term_freq: p.term_freq, - fieldnorm, - positions: p.positions.clone(), - }); - } - Ok(compact) -} - -#[cfg(test)] -mod tests { - use nodedb_types::Surrogate; - - use super::*; - use crate::backend::memory::MemoryBackend; - use crate::index::FtsIndex; - use crate::test_support::test_governor; - - const DB: u64 = 0; - const T: u64 = 1; - const D1: Surrogate = Surrogate(1); - const D2: Surrogate = Surrogate(2); - const D3: Surrogate = Surrogate(3); - - fn make_params<'a>( - tokens: &'a [String], - total: u32, - avg: f32, - top_k: usize, - fuzzy: bool, - bm25: &'a Bm25Params, - ) -> BmwParams<'a> { - BmwParams { - query_tokens: tokens, - raw_tokens: tokens, - fuzzy_enabled: fuzzy, - top_k, - total_docs: total, - avg_doc_len: avg, - bm25, - prefilter: None, - } - } - - #[test] - fn bmw_query_basic() { - let idx = FtsIndex::new(MemoryBackend::new(), test_governor()); - idx.index_document( - DB, - T, - "docs", - D1, - "The quick brown fox jumps over the lazy dog", - ) - .unwrap(); - idx.index_document(DB, T, "docs", D2, "A fast brown dog runs across the field") - .unwrap(); - idx.index_document(DB, T, "docs", D3, "Rust programming language for systems") - .unwrap(); - - let tokens = crate::analyze("brown fox"); - let (total, avg) = idx.index_stats(DB, T, "docs").unwrap(); - let bm25 = Bm25Params::default(); - let p = make_params(&tokens, total, avg, 10, false, &bm25); - - let results = bmw_search(&idx, DB, T, "docs", &p).unwrap().unwrap(); - assert!(!results.is_empty()); - assert_eq!(results[0].doc_id, D1); - } - - #[test] - fn bmw_query_empty_collection() { - let idx = FtsIndex::new(MemoryBackend::new(), test_governor()); - let tokens = crate::analyze("hello"); - let bm25 = Bm25Params::default(); - let p = make_params(&tokens, 0, 1.0, 10, false, &bm25); - - let result = bmw_search(&idx, DB, T, "empty", &p).unwrap(); - assert!(result.is_none()); - } - - #[test] - fn bmw_query_respects_top_k() { - let idx = FtsIndex::new(MemoryBackend::new(), test_governor()); - for i in 1..=50u32 { - idx.index_document(DB, T, "docs", Surrogate(i), &format!("common term word{i}")) - .unwrap(); - } - - let tokens = crate::analyze("common term"); - let (total, avg) = idx.index_stats(DB, T, "docs").unwrap(); - let bm25 = Bm25Params::default(); - let p = make_params(&tokens, total, avg, 5, false, &bm25); - - let results = bmw_search(&idx, DB, T, "docs", &p).unwrap().unwrap(); - assert_eq!(results.len(), 5); - } - - #[test] - fn bmw_query_with_fuzzy() { - let idx = FtsIndex::new(MemoryBackend::new(), test_governor()); - idx.index_document(DB, T, "docs", D1, "distributed database systems") - .unwrap(); - - let stemmed = crate::analyze("databse"); - let raw = crate::analyzer::pipeline::tokenize_no_stem("databse"); - let (total, avg) = idx.index_stats(DB, T, "docs").unwrap(); - let bm25 = Bm25Params::default(); - let p = BmwParams { - query_tokens: &stemmed, - raw_tokens: &raw, - fuzzy_enabled: true, - top_k: 10, - total_docs: total, - avg_doc_len: avg, - bm25: &bm25, - prefilter: None, - }; - - let result = bmw_search(&idx, DB, T, "docs", &p); - match &result { - Ok(Some(r)) => assert!(!r.is_empty(), "BMW returned empty results"), - Ok(None) => panic!("BMW returned None (no term blocks)"), - Err(e) => panic!("BMW returned Err: {e}"), - } - } - - #[test] - fn bmw_query_uses_memtable() { - let idx = FtsIndex::new(MemoryBackend::new(), test_governor()); - idx.index_document(DB, T, "docs", D1, "hello world greeting") - .unwrap(); - - assert!(!idx.memtable().is_empty()); - - let tokens = crate::analyze("hello"); - let (total, avg) = idx.index_stats(DB, T, "docs").unwrap(); - let bm25 = Bm25Params::default(); - let p = make_params(&tokens, total, avg, 10, false, &bm25); - - let results = bmw_search(&idx, DB, T, "docs", &p).unwrap().unwrap(); - assert!(!results.is_empty()); - assert_eq!(results[0].doc_id, D1); - } -} diff --git a/nodedb-fts/src/search/bmw/scorer.rs b/nodedb-fts/src/search/bmw/scorer.rs index d9596ecc1..d9f37f954 100644 --- a/nodedb-fts/src/search/bmw/scorer.rs +++ b/nodedb-fts/src/search/bmw/scorer.rs @@ -5,12 +5,17 @@ //! Two-level skip: WAND pivot selection across terms, then block-level //! pruning within each term's posting list. Operates on `Surrogate` row //! identities for zero-allocation scoring. +//! +//! Admission (the allow and deny sets) is checked before a document is +//! scored, so the top-k cut counts only admitted documents. Every term score +//! comes from [`term_score`] and is combined by [`combine`] in term order, the +//! same functions the per-document scorer uses. use nodedb_types::{Surrogate, SurrogateBitmap}; use crate::bm25; -use crate::codec::smallfloat; use crate::posting::Bm25Params; +use crate::search::doc_score::{Combine, combine, term_score}; use super::heap::TopKHeap; use super::skip_index::TermBlocks; @@ -18,10 +23,33 @@ use super::skip_index::TermBlocks; /// Sentinel returned by an exhausted cursor — sorts after every real surrogate. const EXHAUSTED: Surrogate = Surrogate(u32::MAX); +/// Inputs of one BMW run. +pub struct BmwInput<'a> { + /// Posting blocks of each query term. + pub terms: &'a [TermBlocks], + /// Query word of each term, parallel to `terms`. + pub term_groups: &'a [usize], + /// Number of query words. + pub groups: usize, + pub total_docs: u32, + pub avg_doc_len: f32, + pub params: &'a Bm25Params, + pub top_k: usize, + /// Only these documents are scored. `None` admits every document. + pub allow: Option<&'a SurrogateBitmap>, + /// These documents are never scored. + pub deny: Option<&'a SurrogateBitmap>, + /// How term contributions combine into a document score. Every mode + /// scores at most the plain sum, so the BMW upper bounds stay valid. + pub combine: Combine, +} + /// Per-term iterator state during BMW traversal. struct TermCursor<'a> { /// The term's posting blocks + skip index. term: &'a TermBlocks, + /// Position of the term in the query's term list. + term_idx: usize, /// Current block index within `term.blocks`. block_idx: usize, /// Current position within the current block's doc_ids. @@ -33,7 +61,13 @@ struct TermCursor<'a> { } impl<'a> TermCursor<'a> { - fn new(term: &'a TermBlocks, total_docs: u32, avg_doc_len: f32, params: &Bm25Params) -> Self { + fn new( + term: &'a TermBlocks, + term_idx: usize, + total_docs: u32, + avg_doc_len: f32, + params: &Bm25Params, + ) -> Self { let idf = bm25::idf(term.df, total_docs); let max_score = bm25::term_max_score( term.global_max_tf, @@ -45,6 +79,7 @@ impl<'a> TermCursor<'a> { ); Self { term, + term_idx, block_idx: 0, pos_in_block: 0, idf, @@ -136,50 +171,54 @@ impl<'a> TermCursor<'a> { return 0.0; } let block = &self.term.blocks[self.block_idx]; - let tf = block.term_freqs[self.pos_in_block]; - let fieldnorm = block.fieldnorms[self.pos_in_block]; - let doc_len = smallfloat::decode(fieldnorm).max(1); - - let tf_f = tf as f32; - let dl = doc_len as f32; - - let tf_norm = (tf_f * (params.k1 + 1.0)) - / (tf_f + params.k1 * (1.0 - params.b + params.b * dl / avg_doc_len)); - - self.idf * tf_norm + term_score( + self.idf, + block.term_freqs[self.pos_in_block], + block.fieldnorms[self.pos_in_block], + avg_doc_len, + params, + ) } } +/// Whether `doc_id` may be scored. +fn admitted(input: &BmwInput<'_>, doc_id: Surrogate) -> bool { + input.allow.is_none_or(|bm| bm.contains(doc_id)) + && !input.deny.is_some_and(|bm| bm.contains(doc_id)) +} + /// Run BMW scoring across multiple term posting lists. /// -/// Returns the top-k `(score, doc_id_u32)` results. -/// -/// When `prefilter` is `Some`, only surrogates present in the bitmap are -/// scored; all others are skipped before any BM25 computation. -pub fn bmw_score( - term_blocks: &[TermBlocks], - total_docs: u32, - avg_doc_len: f32, - params: &Bm25Params, - top_k: usize, - prefilter: Option<&SurrogateBitmap>, -) -> TopKHeap { - let mut heap = TopKHeap::new(top_k); - - if term_blocks.is_empty() { +/// Returns the top-k admitted documents by score. +pub fn bmw_score(input: &BmwInput<'_>) -> TopKHeap { + let mut heap = TopKHeap::new(input.top_k); + let BmwInput { + terms, + term_groups, + groups, + total_docs, + avg_doc_len, + params, + .. + } = *input; + + if input.top_k == 0 || terms.is_empty() || input.allow.is_some_and(|bm| bm.is_empty()) { return heap; } - let mut cursors: Vec = term_blocks + let mut cursors: Vec = terms .iter() - .filter(|tb| tb.df > 0) - .map(|tb| TermCursor::new(tb, total_docs, avg_doc_len, params)) + .enumerate() + .filter(|(_, tb)| tb.df > 0) + .map(|(idx, tb)| TermCursor::new(tb, idx, total_docs, avg_doc_len, params)) .collect(); if cursors.is_empty() { return heap; } + let mut contributions: Vec> = vec![None; terms.len()]; + loop { // Sort cursors by current doc_id (ascending). Exhausted cursors go to the end. cursors.sort_by_key(|c| c.current_doc_id()); @@ -218,34 +257,30 @@ pub fn bmw_score( // Check if all cursors [0..=pivot_idx] point to the same doc_id. let first_doc_id = cursors[0].current_doc_id(); if first_doc_id == pivot_doc_id { - // Prefilter: skip surrogates not present in the bitmap. - if let Some(bm) = prefilter - && !bm.contains(pivot_doc_id) - { - for cursor in &mut cursors { - if cursor.current_doc_id() == pivot_doc_id { - cursor.next(); + if admitted(input, pivot_doc_id) { + // Block-level pruning over every cursor on the pivot doc: a + // cursor past the pivot index can sit on the same doc and + // adds to its score. + let block_upper: f32 = cursors + .iter() + .filter(|c| c.current_doc_id() == pivot_doc_id) + .map(|c| c.block_upper_bound(total_docs, avg_doc_len, params)) + .sum(); + + if block_upper > threshold { + contributions.fill(None); + for cursor in cursors.iter() { + if cursor.current_doc_id() == pivot_doc_id { + contributions[cursor.term_idx] = + Some(cursor.score_current(avg_doc_len, params)); + } } - } - continue; - } - - // All essential terms are at the pivot doc — score it. - // But first: block-level pruning. Sum block upper bounds. - let mut block_upper = 0.0f32; - for cursor in cursors.iter().take(pivot_idx + 1) { - block_upper += cursor.block_upper_bound(total_docs, avg_doc_len, params); - } - - if block_upper > threshold { - // Actually score the document. - let mut doc_score = 0.0f32; - for cursor in cursors.iter() { - if cursor.current_doc_id() == pivot_doc_id { - doc_score += cursor.score_current(avg_doc_len, params); + if let Some((score, _)) = + combine(&contributions, term_groups, groups, input.combine) + { + heap.insert(score, pivot_doc_id); } } - heap.insert(doc_score, pivot_doc_id); } // Advance all cursors that were at pivot_doc_id. @@ -287,18 +322,40 @@ mod tests { TermBlocks::from_blocks(blocks) } + fn run( + terms: &[TermBlocks], + total_docs: u32, + top_k: usize, + allow: Option<&SurrogateBitmap>, + deny: Option<&SurrogateBitmap>, + ) -> Vec { + let params = Bm25Params::default(); + let term_groups: Vec = (0..terms.len()).collect(); + bmw_score(&BmwInput { + terms, + term_groups: &term_groups, + groups: terms.len(), + total_docs, + avg_doc_len: 100.0, + params: ¶ms, + top_k, + allow, + deny, + combine: Combine::Sum, + }) + .into_sorted() + .iter() + .map(|d| d.doc_id) + .collect() + } + #[test] fn bmw_basic() { let term_a = make_term(&[0, 1, 2, 3, 4], 2); let term_b = make_term(&[2, 3, 5, 6], 3); - - let params = Bm25Params::default(); - let heap = bmw_score(&[term_a, term_b], 100, 100.0, ¶ms, 3, None); - - let results = heap.into_sorted(); - assert!(!results.is_empty()); + let top_ids = run(&[term_a, term_b], 100, 3, None, None); + assert!(!top_ids.is_empty()); // Docs 2 and 3 match both terms — should score highest. - let top_ids: Vec = results.iter().map(|r| r.doc_id).collect(); assert!(top_ids.contains(&Surrogate(2))); assert!(top_ids.contains(&Surrogate(3))); } @@ -306,30 +363,19 @@ mod tests { #[test] fn bmw_single_term() { let term = make_term(&[10, 20, 30, 40, 50], 1); - let params = Bm25Params::default(); - let heap = bmw_score(&[term], 1000, 100.0, ¶ms, 3, None); - - let results = heap.into_sorted(); - assert_eq!(results.len(), 3); - // All have the same score (same tf, same doc_len). + assert_eq!(run(&[term], 1000, 3, None, None).len(), 3); } #[test] fn bmw_empty_terms() { - let params = Bm25Params::default(); - let heap = bmw_score(&[], 1000, 100.0, ¶ms, 10, None); - assert!(heap.is_empty()); + assert!(run(&[], 1000, 10, None, None).is_empty()); } #[test] fn bmw_respects_top_k() { let ids: Vec = (0..500).collect(); let term = make_term(&ids, 1); - let params = Bm25Params::default(); - let heap = bmw_score(&[term], 1000, 100.0, ¶ms, 5, None); - - let results = heap.into_sorted(); - assert_eq!(results.len(), 5); + assert_eq!(run(&[term], 1000, 5, None, None).len(), 5); } #[test] @@ -341,15 +387,44 @@ mod tests { let term_common = make_term(&common_ids, 1); let term_rare = make_term(&rare_ids, 5); - let params = Bm25Params::default(); - let heap = bmw_score(&[term_common, term_rare], 10_000, 100.0, ¶ms, 3, None); - - let results = heap.into_sorted(); - assert_eq!(results.len(), 3); + let top_ids = run(&[term_common, term_rare], 10_000, 3, None, None); + assert_eq!(top_ids.len(), 3); // The rare-term docs (50, 200, 500) should dominate due to high IDF + tf. - let top_ids: Vec = results.iter().map(|r| r.doc_id).collect(); assert!(top_ids.contains(&Surrogate(50))); assert!(top_ids.contains(&Surrogate(200))); assert!(top_ids.contains(&Surrogate(500))); } + + #[test] + fn deny_set_is_excluded_before_the_cut() { + let term = make_term(&[1, 2, 3, 4, 5], 1); + let mut deny = SurrogateBitmap::new(); + deny.insert(Surrogate(1)); + deny.insert(Surrogate(2)); + let top_ids = run(&[term], 100, 2, None, Some(&deny)); + assert_eq!(top_ids.len(), 2, "the cut counts only admitted documents"); + assert!(!top_ids.contains(&Surrogate(1))); + assert!(!top_ids.contains(&Surrogate(2))); + } + + #[test] + fn allow_set_restricts_candidates() { + let term = make_term(&[1, 2, 3, 4, 5], 1); + let mut allow = SurrogateBitmap::new(); + allow.insert(Surrogate(4)); + assert_eq!( + run(&[term], 100, 10, Some(&allow), None), + vec![Surrogate(4)] + ); + } + + /// The document both terms occur in outranks every single-term document, + /// whichever cursor the pivot lands on. + #[test] + fn two_term_document_wins_top_one() { + let common: Vec = (0..300).collect(); + let a = make_term(&common, 1); + let b = make_term(&[150], 1); + assert_eq!(run(&[a, b], 300, 1, None, None), vec![Surrogate(150)]); + } } diff --git a/nodedb-fts/src/search/bmw/skip_index.rs b/nodedb-fts/src/search/bmw/skip_index.rs index ae08ea256..58b672f7c 100644 --- a/nodedb-fts/src/search/bmw/skip_index.rs +++ b/nodedb-fts/src/search/bmw/skip_index.rs @@ -93,6 +93,19 @@ impl TermBlocks { None } } + + /// The `(term_freq, fieldnorm)` of `doc_id`'s posting, or `None` when the + /// term does not occur in the document. + pub fn lookup(&self, doc_id: Surrogate) -> Option<(u32, u8)> { + let block = &self.blocks[self.advance_to_block(doc_id)?]; + let pos = block.doc_ids.binary_search(&doc_id).ok()?; + Some((block.term_freqs[pos], block.fieldnorms[pos])) + } + + /// Every document the term occurs in. + pub fn doc_ids(&self) -> impl Iterator + '_ { + self.blocks.iter().flat_map(|b| b.doc_ids.iter().copied()) + } } #[cfg(test)] @@ -152,5 +165,20 @@ mod tests { assert_eq!(tb.df, 0); assert_eq!(tb.num_blocks(), 0); assert_eq!(tb.advance_to_block(Surrogate(0)), None); + assert_eq!(tb.lookup(Surrogate(0)), None); + } + + #[test] + fn lookup_finds_postings_across_blocks() { + let ids: Vec = (0..300).map(|i| i * 2).collect(); + let tb = make_term_blocks(&ids, 3); + assert_eq!(tb.lookup(Surrogate(0)), Some((3, smallfloat::encode(100)))); + assert_eq!( + tb.lookup(Surrogate(400)), + Some((3, smallfloat::encode(100))) + ); + assert_eq!(tb.lookup(Surrogate(401)), None); + assert_eq!(tb.lookup(Surrogate(1000)), None); + assert_eq!(tb.doc_ids().count(), 300); } } diff --git a/nodedb-fts/src/search/doc_score.rs b/nodedb-fts/src/search/doc_score.rs new file mode 100644 index 000000000..648dd62d1 --- /dev/null +++ b/nodedb-fts/src/search/doc_score.rs @@ -0,0 +1,108 @@ +// SPDX-License-Identifier: Apache-2.0 + +//! The one BM25 formula every search path scores a document with. +//! +//! The BMW scorer, the per-document scorer, and the scorer of a +//! transaction's staged documents all call these functions. A document +//! therefore scores the same value whichever path reads it. + +use crate::codec::smallfloat; +use crate::posting::Bm25Params; + +/// BM25 contribution of one term to one document: `idf` of the term, the +/// term's frequency in the document, and the document's fieldnorm byte. +pub(crate) fn term_score( + idf: f32, + tf: u32, + fieldnorm: u8, + avg_doc_len: f32, + params: &Bm25Params, +) -> f32 { + let doc_len = smallfloat::decode(fieldnorm).max(1); + let tf_f = tf as f32; + let dl = doc_len as f32; + let tf_norm = (tf_f * (params.k1 + 1.0)) + / (tf_f + params.k1 * (1.0 - params.b + params.b * dl / avg_doc_len)); + idf * tf_norm +} + +/// How a document's term contributions turn into its score. +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum Combine { + /// The plain sum. + Sum, + /// The sum scaled by the share of query words the document matches. + /// The AND-mode fallback ranks partial matches this way. + Coverage, +} + +/// A document's score from its per-term contributions, summed in term order. +/// +/// `contributions` and `term_groups` are parallel: one entry per query term, +/// `None` where the term does not occur in the document. Returns the score and +/// the number of distinct query words matched, or `None` when no term occurs. +pub(crate) fn combine( + contributions: &[Option], + term_groups: &[usize], + groups: usize, + mode: Combine, +) -> Option<(f32, usize)> { + let mut score = 0.0f32; + let mut matched_any = false; + let mut matched = vec![false; groups]; + for (contribution, group) in contributions.iter().zip(term_groups) { + if let Some(value) = contribution { + score += value; + matched_any = true; + if let Some(slot) = matched.get_mut(*group) { + *slot = true; + } + } + } + if !matched_any { + return None; + } + let matched_groups = matched.iter().filter(|m| **m).count(); + let score = match mode { + Combine::Sum => score, + Combine::Coverage => score * (matched_groups as f32 / groups.max(1) as f32), + }; + Some((score, matched_groups)) +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn combine_counts_distinct_words() { + let contributions = [Some(1.0), Some(2.0), None]; + let (score, matched) = combine(&contributions, &[0, 0, 1], 2, Combine::Sum).unwrap(); + assert_eq!(score, 3.0); + assert_eq!(matched, 1); + } + + #[test] + fn coverage_scales_by_matched_share() { + let contributions = [Some(2.0), None]; + let (score, matched) = combine(&contributions, &[0, 1], 2, Combine::Coverage).unwrap(); + assert_eq!(score, 1.0); + assert_eq!(matched, 1); + } + + #[test] + fn no_contribution_is_no_match() { + assert!(combine(&[None, None], &[0, 1], 2, Combine::Sum).is_none()); + } + + #[test] + fn term_score_matches_bm25_formula() { + let params = Bm25Params::default(); + let fieldnorm = smallfloat::encode(10); + let idf = crate::bm25::idf(3, 100); + let expected = + crate::bm25::bm25_score(2, 3, smallfloat::decode(fieldnorm), 100, 10.0, ¶ms); + let got = term_score(idf, 2, fieldnorm, 10.0, ¶ms); + assert!((expected - got).abs() < 1e-6); + } +} diff --git a/nodedb-fts/src/search/doc_scorer.rs b/nodedb-fts/src/search/doc_scorer.rs new file mode 100644 index 000000000..226ceff54 --- /dev/null +++ b/nodedb-fts/src/search/doc_scorer.rs @@ -0,0 +1,271 @@ +// SPDX-License-Identifier: Apache-2.0 + +//! Per-document `bm25_score(field, query)` values. +//! +//! A [`DocScorer`] resolves its query once, then scores any document by +//! point lookups: the document's postings of the query terms, and its +//! recorded length for index membership. It reads no corpus-wide score map, +//! so scoring the rows a query emits costs memory in the query terms' +//! postings, never in the collection. +//! +//! The score is the one a search of the same query returns for the +//! document, under the same match-mode decision over the same admitted rows. + +use nodedb_types::{Surrogate, SurrogateBitmap}; + +use crate::backend::FtsBackend; +use crate::index::FtsIndex; +use crate::index::error::FtsIndexError; +use crate::scope::IndexScope; +use crate::search::doc_score::term_score; +use crate::search::match_mode::{MatchMode, staged_candidates}; +use crate::search::query_terms::{ResolvedQuery, TextQuery}; +use crate::search::staged::{StagedView, staged_contributions}; + +/// One document's score against one index. +#[derive(Debug, Clone, Copy, PartialEq)] +pub enum DocScore { + /// The document matches the query with this score. + Match(f32), + /// The index holds the document, and the query does not match it. + Miss, + /// The index does not hold the document. + Absent, +} + +/// Scores documents against one resolved query. +pub struct DocScorer<'a, B: FtsBackend> { + fts: &'a FtsIndex, + database_id: u64, + tid: u64, + index: IndexScope<'a>, + staged: Option, + /// `None` when the query has no positive term: it matches nothing. + resolved: Option<(ResolvedQuery, MatchMode)>, +} + +impl FtsIndex { + /// A scorer of `query` against `index`. `eligible` is the set of rows the + /// reading query admits: the AND-mode fallback is decided over it, the + /// same way a search with that prefilter decides it. `staged` is the + /// open transaction's view of the index. + pub fn doc_scorer<'a>( + &'a self, + database_id: u64, + tid: u64, + index: impl Into>, + query: TextQuery<'_>, + eligible: Option<&SurrogateBitmap>, + staged: Option, + ) -> Result, FtsIndexError> { + let index = index.into(); + let resolved = match self.resolve_query(database_id, tid, index, query, staged.as_ref())? { + Some(resolved) => { + let mut deny = resolved.negated.clone(); + if let Some(view) = staged.as_ref() { + deny.union_in_place(view.hidden()); + } + let index_visible = staged.as_ref().is_none_or(|view| !view.hides_all()); + let candidates = + staged_candidates(&resolved, staged.as_ref(), eligible, &self.bm25_params); + let mode = resolved.match_mode(index_visible, eligible, &deny, &candidates); + Some((resolved, mode)) + } + None => None, + }; + Ok(DocScorer { + fts: self, + database_id, + tid, + index, + staged, + resolved, + }) + } +} + +impl DocScorer<'_, B> { + /// The score of each of `docs`, parallel to `docs`. Index membership of + /// the documents that do not match is read in one batch. + pub fn score(&self, docs: &[Surrogate]) -> Result, FtsIndexError> { + let mut scores = vec![DocScore::Absent; docs.len()]; + let mut membership: Vec<(usize, Surrogate)> = Vec::new(); + for (slot, doc_id) in docs.iter().enumerate() { + if let Some(view) = self.staged.as_ref() { + if let Some(doc) = view.doc(*doc_id) { + scores[slot] = self + .staged_match(doc) + .map_or(DocScore::Miss, DocScore::Match); + continue; + } + if view.hides(*doc_id) { + continue; + } + } + match self.indexed_match(*doc_id) { + Some(score) => scores[slot] = DocScore::Match(score), + None => membership.push((slot, *doc_id)), + } + } + if membership.is_empty() { + return Ok(scores); + } + let ids: Vec = membership.iter().map(|(_, doc_id)| *doc_id).collect(); + let lengths = self + .fts + .backend + .read_doc_lengths(self.database_id, self.tid, self.index, &ids) + .map_err(FtsIndexError::backend)?; + for ((slot, _), length) in membership.into_iter().zip(lengths) { + if length.is_some() { + scores[slot] = DocScore::Miss; + } + } + Ok(scores) + } + + fn staged_match(&self, doc: &crate::search::staged::StagedDoc) -> Option { + let (query, mode) = self.resolved.as_ref()?; + let contributions = staged_contributions(query, doc, &self.fts.bm25_params)?; + mode.score(query, &contributions) + } + + fn indexed_match(&self, doc_id: Surrogate) -> Option { + let (query, mode) = self.resolved.as_ref()?; + if query.negated.contains(doc_id) { + return None; + } + let params = &self.fts.bm25_params; + let contributions: Vec> = query + .blocks + .iter() + .map(|blocks| { + blocks.lookup(doc_id).map(|(tf, fieldnorm)| { + term_score( + crate::bm25::idf(blocks.df, query.total_docs), + tf, + fieldnorm, + query.avg_doc_len, + params, + ) + }) + }) + .collect(); + mode.score(query, &contributions) + } +} + +#[cfg(test)] +mod tests { + use nodedb_types::{Surrogate, SurrogateBitmap}; + + use super::DocScore; + use crate::backend::memory::MemoryBackend; + use crate::index::FtsIndex; + use crate::posting::QueryMode; + use crate::search::bm25_search::FtsSearchParams; + use crate::search::query_terms::TextQuery; + use crate::search::staged::{StagedDoc, StagedView}; + use crate::test_support::test_governor; + + const DB: u64 = 0; + const T: u64 = 1; + + fn query(text: &str) -> TextQuery<'_> { + TextQuery { + query: text, + fuzzy_enabled: false, + mode: QueryMode::And, + } + } + + /// A document's point score equals the score a search returns for it. + #[test] + fn point_scores_equal_search_scores() { + let idx = FtsIndex::new(MemoryBackend::new(), test_governor()); + idx.index_document(DB, T, "docs", Surrogate(1), "rust systems language") + .unwrap(); + idx.index_document(DB, T, "docs", Surrogate(2), "rust rust web") + .unwrap(); + idx.index_document(DB, T, "docs", Surrogate(3), "python scripting") + .unwrap(); + + let hits = idx + .search( + DB, + T, + "docs", + FtsSearchParams { + query: "rust", + top_k: 10, + fuzzy_enabled: false, + mode: QueryMode::And, + prefilter: None, + }, + ) + .unwrap(); + let scorer = idx + .doc_scorer(DB, T, "docs", query("rust"), None, None) + .unwrap(); + let docs: Vec = hits.iter().map(|h| h.doc_id).collect(); + let scores = scorer.score(&docs).unwrap(); + for (hit, score) in hits.iter().zip(scores) { + assert_eq!(score, DocScore::Match(hit.score)); + } + assert_eq!( + scorer.score(&[Surrogate(3), Surrogate(99)]).unwrap(), + vec![DocScore::Miss, DocScore::Absent] + ); + } + + /// The fallback decision reads the eligible rows: with the only AND + /// match outside them, a one-word match scores under OR fallback. + #[test] + fn fallback_follows_the_eligible_rows() { + let idx = FtsIndex::new(MemoryBackend::new(), test_governor()); + idx.index_document(DB, T, "docs", Surrogate(1), "alpha bravo") + .unwrap(); + idx.index_document(DB, T, "docs", Surrogate(2), "alpha charlie") + .unwrap(); + + let all = idx + .doc_scorer(DB, T, "docs", query("alpha bravo"), None, None) + .unwrap(); + assert_eq!(all.score(&[Surrogate(2)]).unwrap(), vec![DocScore::Miss]); + + let mut eligible = SurrogateBitmap::new(); + eligible.insert(Surrogate(2)); + let scoped = idx + .doc_scorer(DB, T, "docs", query("alpha bravo"), Some(&eligible), None) + .unwrap(); + assert!(matches!( + scoped.score(&[Surrogate(2)]).unwrap()[0], + DocScore::Match(_) + )); + } + + /// A staged document scores from its tokens; a hidden indexed document + /// the transaction removed is absent. + #[test] + fn staged_documents_score_and_hidden_documents_are_absent() { + let idx = FtsIndex::new(MemoryBackend::new(), test_governor()); + idx.index_document(DB, T, "docs", Surrogate(1), "rust language") + .unwrap(); + let mut hidden = SurrogateBitmap::new(); + hidden.insert(Surrogate(1)); + let view = StagedView::new( + hidden, + false, + vec![StagedDoc { + doc_id: Surrogate(2), + tokens: vec!["rust".into()], + }], + ); + let scorer = idx + .doc_scorer(DB, T, "docs", query("rust"), None, Some(view)) + .unwrap(); + let scores = scorer.score(&[Surrogate(1), Surrogate(2)]).unwrap(); + assert_eq!(scores[0], DocScore::Absent); + assert!(matches!(scores[1], DocScore::Match(_))); + } +} diff --git a/nodedb-fts/src/search/fuzzy_search.rs b/nodedb-fts/src/search/fuzzy_search.rs index 9d30aeef7..b7496f5dc 100644 --- a/nodedb-fts/src/search/fuzzy_search.rs +++ b/nodedb-fts/src/search/fuzzy_search.rs @@ -1,77 +1,42 @@ // SPDX-License-Identifier: Apache-2.0 -//! Fuzzy term lookup for the FtsIndex. +//! Fuzzy term candidates for the FtsIndex. use crate::backend::FtsBackend; use crate::fuzzy; -use crate::index::FtsIndex; +use crate::index::{FtsIndex, FtsIndexError}; use crate::lsm::query as lsm_query; -use crate::posting::Posting; +use crate::scope::IndexScope; impl FtsIndex { - /// Find the best fuzzy-matching term and return its posting list. - pub(crate) fn fuzzy_lookup( + /// Terms within fuzzy distance of `query_term`, best first. The + /// vocabulary is every term of the index plus `extra_vocabulary` (the + /// tokens of a transaction's staged documents). + pub(crate) fn fuzzy_candidates( &self, database_id: u64, tid: u64, - collection: &str, + index: IndexScope<'_>, query_term: &str, - ) -> Result<(Vec, bool), B::Error> { - let mut all_terms = lsm_query::collect_all_terms( - &self.backend, - database_id, - tid, - collection, - self.memtable(), - )?; - - let backend_terms = self - .backend - .collection_terms(database_id, tid, collection)?; - for t in backend_terms { - all_terms.push(t); - } + extra_vocabulary: &[&str], + ) -> Result, FtsIndexError> { + let mut all_terms = + lsm_query::collect_all_terms(&self.backend, database_id, tid, index, self.memtable())?; + all_terms.extend( + self.backend + .collection_terms(database_id, tid, index) + .map_err(FtsIndexError::backend)?, + ); + all_terms.extend(extra_vocabulary.iter().map(|t| (*t).to_string())); all_terms.sort(); all_terms.dedup(); - let matches = fuzzy::fuzzy_match(query_term, all_terms.iter().map(String::as_str)); - - if let Some((best_term, _dist)) = matches.first() { - let postings = self - .backend - .read_postings(database_id, tid, collection, best_term)?; - if !postings.is_empty() { - return Ok((postings, true)); - } - - let tokens = vec![best_term.to_string()]; - let term_blocks = lsm_query::collect_merged_term_blocks( - &self.backend, - database_id, - tid, - collection, - self.memtable(), - &tokens, - &self.governor, - )?; - if !term_blocks.is_empty() && term_blocks[0].df > 0 { - let mut postings = Vec::new(); - for block in &term_blocks[0].blocks { - for i in 0..block.doc_ids.len() { - postings.push(Posting { - doc_id: block.doc_ids[i], - term_freq: block.term_freqs[i], - positions: block.positions[i].clone(), - }); - } - } - if !postings.is_empty() { - return Ok((postings, true)); - } - } - } - - Ok((Vec::new(), false)) + Ok( + fuzzy::fuzzy_match(query_term, all_terms.iter().map(String::as_str)) + .into_iter() + .map(|(term, _dist)| term.to_string()) + .collect(), + ) } } @@ -86,24 +51,35 @@ mod tests { const T: u64 = 1; #[test] - fn fuzzy_lookup_finds_close_term() { + fn fuzzy_candidates_find_close_term() { let idx = FtsIndex::new(MemoryBackend::new(), test_governor()); idx.index_document(DB, T, "docs", Surrogate(1), "distributed database systems") .unwrap(); - let (postings, is_fuzzy) = idx.fuzzy_lookup(DB, T, "docs", "databse").unwrap(); - assert!(is_fuzzy, "should find fuzzy match"); - assert!(!postings.is_empty(), "should return postings from LSM"); + let candidates = idx + .fuzzy_candidates(DB, T, "docs".into(), "databse", &[]) + .unwrap(); + assert!(!candidates.is_empty(), "should find a fuzzy match"); } #[test] - fn fuzzy_lookup_no_match() { + fn fuzzy_candidates_no_match() { let idx = FtsIndex::new(MemoryBackend::new(), test_governor()); idx.index_document(DB, T, "docs", Surrogate(1), "hello world") .unwrap(); - let (postings, is_fuzzy) = idx.fuzzy_lookup(DB, T, "docs", "zzzzzzz").unwrap(); - assert!(postings.is_empty()); - assert!(!is_fuzzy); + let candidates = idx + .fuzzy_candidates(DB, T, "docs".into(), "zzzzzzz", &[]) + .unwrap(); + assert!(candidates.is_empty()); + } + + #[test] + fn fuzzy_candidates_include_extra_vocabulary() { + let idx = FtsIndex::new(MemoryBackend::new(), test_governor()); + let candidates = idx + .fuzzy_candidates(DB, T, "docs".into(), "databse", &["databas"]) + .unwrap(); + assert_eq!(candidates, vec!["databas".to_string()]); } } diff --git a/nodedb-fts/src/search/match_mode.rs b/nodedb-fts/src/search/match_mode.rs new file mode 100644 index 000000000..c97b57082 --- /dev/null +++ b/nodedb-fts/src/search/match_mode.rs @@ -0,0 +1,147 @@ +// SPDX-License-Identifier: Apache-2.0 + +//! Which documents a resolved query matches, decided once per query. +//! +//! AND mode over several words matches documents holding every word. When no +//! admitted document holds every word, the query falls back to OR with +//! coverage-scaled scores. The decision reads only whether any admitted +//! document holds every word, never the top-k cut, so it is the same for +//! every LIMIT and for every reader of the same admitted set. + +use nodedb_types::{Surrogate, SurrogateBitmap}; + +use crate::posting::Bm25Params; +use crate::search::doc_score::{Combine, combine}; +use crate::search::query_terms::ResolvedQuery; +use crate::search::staged::{StagedView, staged_contributions}; + +/// How a resolved query matches documents. +pub(crate) enum MatchMode { + /// Every word must match. Holds the admitted indexed documents that hold + /// every word. + All(SurrogateBitmap), + /// AND-mode fallback: any word matches, scores scale by word coverage. + Coverage, + /// Any word matches. + Any, +} + +impl MatchMode { + /// How term contributions combine under this mode. + pub(crate) fn combine(&self) -> Combine { + match self { + MatchMode::Coverage => Combine::Coverage, + MatchMode::All(_) | MatchMode::Any => Combine::Sum, + } + } + + /// A document's score under this mode, or `None` when it does not match. + pub(crate) fn score( + &self, + query: &ResolvedQuery, + contributions: &[Option], + ) -> Option { + let (score, matched) = combine( + contributions, + &query.term_groups, + query.groups, + self.combine(), + )?; + match self { + MatchMode::All(_) if matched < query.groups => None, + MatchMode::All(_) | MatchMode::Coverage | MatchMode::Any => Some(score), + } + } +} + +/// The staged documents `allow` admits and no negative term excludes, each +/// with its per-term contributions. +pub(crate) fn staged_candidates( + query: &ResolvedQuery, + staged: Option<&StagedView>, + allow: Option<&SurrogateBitmap>, + params: &Bm25Params, +) -> Vec<(Surrogate, Vec>)> { + let Some(view) = staged else { + return Vec::new(); + }; + view.docs() + .iter() + .filter(|doc| allow.is_none_or(|bm| bm.contains(doc.doc_id))) + .filter_map(|doc| { + staged_contributions(query, doc, params) + .map(|contributions| (doc.doc_id, contributions)) + }) + .collect() +} + +impl ResolvedQuery { + /// Decide the match mode. `allow` and `deny` bound the indexed documents + /// the query reads. `index_visible` is `false` when the transaction + /// truncated the collection. `staged` are the admitted staged candidates. + pub(crate) fn match_mode( + &self, + index_visible: bool, + allow: Option<&SurrogateBitmap>, + deny: &SurrogateBitmap, + staged: &[(Surrogate, Vec>)], + ) -> MatchMode { + if !self.require_all { + return MatchMode::Any; + } + let all_words = if index_visible { + self.all_words_docs(allow, deny) + } else { + SurrogateBitmap::new() + }; + let staged_holds_all = staged.iter().any(|(_, contributions)| { + combine(contributions, &self.term_groups, self.groups, Combine::Sum) + .is_some_and(|(_, matched)| matched == self.groups) + }); + if all_words.is_empty() && !staged_holds_all { + MatchMode::Coverage + } else { + MatchMode::All(all_words) + } + } + + /// Indexed documents holding every query word, within `allow` and + /// outside `deny`. + fn all_words_docs( + &self, + allow: Option<&SurrogateBitmap>, + deny: &SurrogateBitmap, + ) -> SurrogateBitmap { + let mut result: Option = None; + for group in 0..self.groups { + let mut word = SurrogateBitmap::new(); + for (blocks, _) in self + .blocks + .iter() + .zip(&self.term_groups) + .filter(|(_, g)| **g == group) + { + for doc_id in blocks.doc_ids() { + word.insert(doc_id); + } + } + let next = match result { + None => word, + Some(mut acc) => { + acc.intersect_in_place(&word); + acc + } + }; + let empty = next.is_empty(); + result = Some(next); + if empty { + break; + } + } + let mut docs = result.unwrap_or_default(); + if let Some(allow) = allow { + docs.intersect_in_place(allow); + } + docs.difference(deny) + } +} diff --git a/nodedb-fts/src/search/mod.rs b/nodedb-fts/src/search/mod.rs index 0210a462e..9f4db34a3 100644 --- a/nodedb-fts/src/search/mod.rs +++ b/nodedb-fts/src/search/mod.rs @@ -2,7 +2,12 @@ pub mod bm25_search; pub mod bmw; +pub mod doc_score; +pub mod doc_scorer; pub mod field_scoring; pub mod fuzzy_search; +pub mod match_mode; pub mod phrase; pub mod query_parser; +pub mod query_terms; +pub mod staged; diff --git a/nodedb-fts/src/search/phrase.rs b/nodedb-fts/src/search/phrase.rs index 4e13f8671..0a8468996 100644 --- a/nodedb-fts/src/search/phrase.rs +++ b/nodedb-fts/src/search/phrase.rs @@ -11,8 +11,6 @@ use std::collections::HashMap; -use nodedb_types::Surrogate; - use crate::posting::Posting; /// Maximum phrase boost multiplier for a full consecutive phrase match. @@ -75,33 +73,6 @@ fn has_consecutive_positions(a: &[u32], b: &[u32]) -> bool { false } -/// Collect per-document postings for phrase boost computation. -/// -/// Given the query tokens and the posting lists retrieved during BM25 scoring, -/// returns a map from `doc_id` → (token_index → posting). -pub(crate) fn collect_doc_postings<'a>( - query_tokens: &[String], - term_postings: &'a [(Vec, bool)], -) -> HashMap> { - let mut doc_map: HashMap> = HashMap::new(); - - for (token_idx, (postings, _is_fuzzy)) in term_postings.iter().enumerate() { - for posting in postings { - doc_map - .entry(posting.doc_id) - .or_default() - .insert(token_idx, posting); - } - } - - // Only keep documents that have at least 2 matched terms (phrase boost is meaningless for 1). - if query_tokens.len() >= 2 { - doc_map.retain(|_, postings| postings.len() >= 2); - } - - doc_map -} - #[cfg(test)] mod tests { use super::*; diff --git a/nodedb-fts/src/search/query_terms.rs b/nodedb-fts/src/search/query_terms.rs new file mode 100644 index 000000000..8ea90eaea --- /dev/null +++ b/nodedb-fts/src/search/query_terms.rs @@ -0,0 +1,289 @@ +// SPDX-License-Identifier: Apache-2.0 + +//! A query string resolved against one index: its positive terms grouped by +//! query word, each with the posting blocks it scores against, and the +//! documents its negative terms exclude. +//! +//! Every search path resolves a query here, once, so the analyzer, synonym +//! expansion, fuzzy replacement, and negation are identical for indexed and +//! staged documents. + +use std::collections::HashSet; + +use nodedb_types::{Surrogate, SurrogateBitmap}; + +use crate::backend::FtsBackend; +use crate::block::{CompactPosting, into_blocks}; +use crate::codec::smallfloat; +use crate::index::FtsIndex; +use crate::index::error::FtsIndexError; +use crate::lsm::query as lsm_query; +use crate::posting::{Posting, QueryMode}; +use crate::scope::IndexScope; +use crate::search::bmw::skip_index::TermBlocks; +use crate::search::query_parser::parse_query; +use crate::search::staged::StagedView; + +/// The query inputs every search path resolves the same way. +#[derive(Debug, Clone, Copy)] +pub struct TextQuery<'a> { + /// Raw query string (may contain `NOT ` / `-` negation). + pub query: &'a str, + /// When `true`, a term with no posting falls back to fuzzy lookup. A + /// collection configured with `FUZZY true` falls back either way. + pub fuzzy_enabled: bool, + /// Boolean combination of the query words. + pub mode: QueryMode, +} + +/// A query resolved against one index. `texts`, `blocks`, and `term_groups` +/// are parallel: one entry per positive query term. +pub(crate) struct ResolvedQuery { + /// Each term looked up: the analyzed token, or its fuzzy replacement. + pub(crate) texts: Vec, + /// Posting blocks of each term. + pub(crate) blocks: Vec, + /// Query word of each term. + pub(crate) term_groups: Vec, + /// Whether any term is a fuzzy replacement. + pub(crate) fuzzy: bool, + /// Number of query words. + pub(crate) groups: usize, + /// AND mode over more than one word: a match holds every word. + pub(crate) require_all: bool, + /// Analyzed, synonym-expanded negative terms. + pub(crate) negative_terms: HashSet, + /// Indexed documents holding a negative term. + pub(crate) negated: SurrogateBitmap, + pub(crate) total_docs: u32, + pub(crate) avg_doc_len: f32, +} + +impl FtsIndex { + /// Resolve `query` against `index`. `None` when the query has no + /// positive term after analysis, so it matches nothing. + pub(crate) fn resolve_query( + &self, + database_id: u64, + tid: u64, + index: IndexScope<'_>, + query: TextQuery<'_>, + staged: Option<&StagedView>, + ) -> Result, FtsIndexError> { + let collection = index.collection(); + let fuzzy_enabled = query.fuzzy_enabled + || self + .get_collection_fuzzy(database_id, tid, collection) + .map_err(FtsIndexError::backend)?; + let parsed = parse_query(query.query)?; + let positive_raw = parsed.positive.join(" "); + let base_tokens = self + .analyze_for_collection(database_id, tid, collection, &positive_raw) + .map_err(FtsIndexError::backend)?; + if base_tokens.is_empty() { + return Ok(None); + } + let raw_tokens = if fuzzy_enabled { + self.tokenize_raw_for_collection(database_id, tid, collection, &positive_raw) + .map_err(FtsIndexError::backend)? + } else { + Vec::new() + }; + let words = self + .expand_query_groups(database_id, tid, base_tokens) + .map_err(FtsIndexError::backend)?; + let groups = words.len(); + + let mut texts: Vec = Vec::new(); + let mut term_groups: Vec = Vec::new(); + let mut is_word: Vec = Vec::new(); + for (group, word) in words.into_iter().enumerate() { + for (pos, term) in word.into_iter().enumerate() { + texts.push(term); + term_groups.push(group); + is_word.push(pos == 0); + } + } + + let mut blocks = self.term_blocks(database_id, tid, index, &texts)?; + let mut fuzzy = false; + let positions = is_word.iter().zip(&term_groups); + for ((text, term_blocks), (word, group)) in + texts.iter_mut().zip(blocks.iter_mut()).zip(positions) + { + if !fuzzy_enabled || term_is_live(text, term_blocks, staged) { + continue; + } + let raw = if *word { + raw_tokens.get(*group).map_or(text.as_str(), String::as_str) + } else { + text.as_str() + }; + if let Some((replacement, replacement_blocks)) = + self.fuzzy_replacement(database_id, tid, index, raw, staged)? + { + *text = replacement; + *term_blocks = replacement_blocks; + fuzzy = true; + } + } + + let (negative_terms, negated) = + self.negative_terms(database_id, tid, index, &parsed.negative)?; + let (total_docs, avg_doc_len) = self + .index_stats(database_id, tid, index) + .map_err(FtsIndexError::backend)?; + Ok(Some(ResolvedQuery { + texts, + blocks, + term_groups, + fuzzy, + groups, + require_all: query.mode == QueryMode::And && groups > 1, + negative_terms, + negated, + total_docs, + avg_doc_len, + })) + } + + /// The closest fuzzy term that matches a live document, with its blocks. + fn fuzzy_replacement( + &self, + database_id: u64, + tid: u64, + index: IndexScope<'_>, + raw: &str, + staged: Option<&StagedView>, + ) -> Result, FtsIndexError> { + let vocabulary = staged.map(StagedView::vocabulary).unwrap_or_default(); + let candidates = self.fuzzy_candidates(database_id, tid, index, raw, &vocabulary)?; + for candidate in candidates { + let mut blocks = + self.term_blocks(database_id, tid, index, std::slice::from_ref(&candidate))?; + let Some(blocks) = blocks.pop() else { + continue; + }; + if term_is_live(&candidate, &blocks, staged) { + return Ok(Some((candidate, blocks))); + } + } + Ok(None) + } + + /// The analyzed, synonym-expanded negative terms and the indexed + /// documents holding any of them. + fn negative_terms( + &self, + database_id: u64, + tid: u64, + index: IndexScope<'_>, + raw_negative: &[String], + ) -> Result<(HashSet, SurrogateBitmap), FtsIndexError> { + let mut negated = SurrogateBitmap::new(); + if raw_negative.is_empty() { + return Ok((HashSet::new(), negated)); + } + let base = self + .analyze_for_collection( + database_id, + tid, + index.collection(), + &raw_negative.join(" "), + ) + .map_err(FtsIndexError::backend)?; + if base.is_empty() { + return Ok((HashSet::new(), negated)); + } + let terms = self + .expand_query_with_synonyms(database_id, tid, base) + .map_err(FtsIndexError::backend)?; + for blocks in self.term_blocks(database_id, tid, index, &terms)? { + for doc_id in blocks.doc_ids() { + negated.insert(doc_id); + } + } + Ok((terms.into_iter().collect(), negated)) + } + + /// Posting blocks of each term: the LSM memtable and segments, or the + /// backend's posting table when the LSM holds none. + pub(crate) fn term_blocks( + &self, + database_id: u64, + tid: u64, + index: IndexScope<'_>, + texts: &[String], + ) -> Result, FtsIndexError> { + let mut blocks = lsm_query::collect_merged_term_blocks( + &self.backend, + database_id, + tid, + index, + self.memtable(), + texts, + &self.governor, + )?; + for (text, term) in texts.iter().zip(blocks.iter_mut()) { + if term.df > 0 { + continue; + } + let postings = self + .backend + .read_postings(database_id, tid, index, text) + .map_err(FtsIndexError::backend)?; + if !postings.is_empty() { + let compact = self.compact_postings(database_id, tid, index, &postings)?; + *term = TermBlocks::from_blocks(into_blocks(compact)); + } + } + Ok(blocks) + } + + /// Backend postings as compact postings, each with the fieldnorm of its + /// document's recorded length. Lengths are read in one batch. + fn compact_postings( + &self, + database_id: u64, + tid: u64, + index: IndexScope<'_>, + postings: &[Posting], + ) -> Result, FtsIndexError> { + let doc_ids: Vec = postings.iter().map(|p| p.doc_id).collect(); + let lengths = self + .backend + .read_doc_lengths(database_id, tid, index, &doc_ids) + .map_err(FtsIndexError::backend)?; + let mut compact = Vec::with_capacity(postings.len()); + for (posting, length) in postings.iter().zip(lengths) { + let doc_len = match length { + Some(len) => len, + // A posting outlives no length in a consistent index. The + // term frequency is the tightest length the posting proves. + None => self + .read_fieldnorm(database_id, tid, index, posting.doc_id) + .map_err(FtsIndexError::backend)? + .unwrap_or(posting.term_freq), + }; + compact.push(CompactPosting { + doc_id: posting.doc_id, + term_freq: posting.term_freq, + fieldnorm: smallfloat::encode(doc_len), + positions: posting.positions.clone(), + }); + } + Ok(compact) + } +} + +/// Whether `term` matches a document the reader sees: an indexed document +/// the transaction does not hide, or a staged document. +fn term_is_live(term: &str, blocks: &TermBlocks, staged: Option<&StagedView>) -> bool { + match staged { + None => blocks.df > 0, + Some(view) => { + (!view.hides_all() && blocks.doc_ids().any(|doc| !view.hides(doc))) + || view.holds_term(term) + } + } +} diff --git a/nodedb-fts/src/search/staged.rs b/nodedb-fts/src/search/staged.rs new file mode 100644 index 000000000..588501952 --- /dev/null +++ b/nodedb-fts/src/search/staged.rs @@ -0,0 +1,171 @@ +// SPDX-License-Identifier: Apache-2.0 + +//! An open transaction's view over one index. +//! +//! A transaction's writes are not in the index until it commits. A search +//! inside the transaction reads the index with the documents the transaction +//! replaced or removed hidden, and scores the transaction's own documents +//! from their analyzed tokens with the same query semantics and formula as +//! indexed documents. Staged documents score against the index's corpus +//! statistics. + +use std::collections::HashMap; + +use nodedb_types::{Surrogate, SurrogateBitmap}; + +use crate::codec::smallfloat; +use crate::posting::Bm25Params; +use crate::search::doc_score::term_score; +use crate::search::query_terms::ResolvedQuery; + +/// A document the transaction wrote, with the analyzed tokens of the text +/// it holds in the index, in order. +#[derive(Debug, Clone)] +pub struct StagedDoc { + pub doc_id: Surrogate, + pub tokens: Vec, +} + +/// The transaction's writes as one index sees them. +#[derive(Debug, Clone)] +pub struct StagedView { + hidden: SurrogateBitmap, + hide_all: bool, + docs: Vec, + by_id: HashMap, +} + +impl StagedView { + /// `hidden`: indexed documents the transaction replaced or removed. + /// `hide_all`: the transaction truncated the collection, so no indexed + /// document counts. `docs`: the transaction's documents that hold text + /// in this index. A document listed in `docs` is hidden in the index too. + pub fn new(mut hidden: SurrogateBitmap, hide_all: bool, docs: Vec) -> Self { + let by_id = docs + .iter() + .enumerate() + .map(|(idx, doc)| { + hidden.insert(doc.doc_id); + (doc.doc_id, idx) + }) + .collect(); + Self { + hidden, + hide_all, + docs, + by_id, + } + } + + /// Whether the transaction replaced or removed the indexed `doc_id`. + pub fn hides(&self, doc_id: Surrogate) -> bool { + self.hide_all || self.hidden.contains(doc_id) + } + + /// Whether every indexed document is hidden. + pub fn hides_all(&self) -> bool { + self.hide_all + } + + /// Indexed documents the transaction replaced or removed. + pub fn hidden(&self) -> &SurrogateBitmap { + &self.hidden + } + + /// The transaction's documents that hold text in this index. + pub fn docs(&self) -> &[StagedDoc] { + &self.docs + } + + /// The staged document `doc_id`, when the transaction wrote one that + /// holds text in this index. + pub fn doc(&self, doc_id: Surrogate) -> Option<&StagedDoc> { + self.by_id.get(&doc_id).and_then(|idx| self.docs.get(*idx)) + } + + /// Whether any staged document holds `term`. + pub(crate) fn holds_term(&self, term: &str) -> bool { + self.docs.iter().any(|d| d.tokens.iter().any(|t| t == term)) + } + + /// Every distinct token of the staged documents. + pub(crate) fn vocabulary(&self) -> Vec<&str> { + let mut terms: Vec<&str> = self + .docs + .iter() + .flat_map(|d| d.tokens.iter().map(String::as_str)) + .collect(); + terms.sort_unstable(); + terms.dedup(); + terms + } +} + +/// Per-term contributions of a staged document to `query`, parallel to the +/// query's terms. `None` when a negative term excludes the document. +pub(crate) fn staged_contributions( + query: &ResolvedQuery, + doc: &StagedDoc, + params: &Bm25Params, +) -> Option>> { + let mut tf: HashMap<&str, u32> = HashMap::new(); + for token in &doc.tokens { + *tf.entry(token.as_str()).or_insert(0) += 1; + } + if query + .negative_terms + .iter() + .any(|term| tf.contains_key(term.as_str())) + { + return None; + } + let doc_len = u32::try_from(doc.tokens.len()).unwrap_or(u32::MAX); + let fieldnorm = smallfloat::encode(doc_len); + Some( + query + .texts + .iter() + .zip(&query.blocks) + .map(|(text, blocks)| { + tf.get(text.as_str()).map(|freq| { + term_score( + crate::bm25::idf(blocks.df, query.total_docs), + *freq, + fieldnorm, + query.avg_doc_len, + params, + ) + }) + }) + .collect(), + ) +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn staged_docs_are_hidden_in_the_index() { + let view = StagedView::new( + SurrogateBitmap::new(), + false, + vec![StagedDoc { + doc_id: Surrogate(7), + tokens: vec!["rust".into()], + }], + ); + assert!(view.hides(Surrogate(7))); + assert!(!view.hides(Surrogate(8))); + assert!(view.doc(Surrogate(7)).is_some()); + assert!(view.holds_term("rust")); + assert!(!view.holds_term("python")); + } + + #[test] + fn truncate_hides_every_indexed_document() { + let view = StagedView::new(SurrogateBitmap::new(), true, Vec::new()); + assert!(view.hides(Surrogate(1))); + assert!(view.hides_all()); + } +} diff --git a/nodedb-graph/src/bfs_params.rs b/nodedb-graph/src/bfs_params.rs index ce045dfc5..6f5653491 100644 --- a/nodedb-graph/src/bfs_params.rs +++ b/nodedb-graph/src/bfs_params.rs @@ -10,7 +10,9 @@ use crate::csr::Direction; /// traversal semantics. pub struct BfsParams<'a> { pub start_nodes: &'a [&'a str], - pub label_filter: Option<&'a str>, + /// Empty keeps every edge. Otherwise an edge whose label is any listed + /// label passes. + pub label_filter: &'a [&'a str], pub direction: Direction, pub max_depth: usize, pub max_visited: usize, diff --git a/nodedb-graph/src/csr/index/label_filter.rs b/nodedb-graph/src/csr/index/label_filter.rs index 2b29cf757..78f5ce6e2 100644 --- a/nodedb-graph/src/csr/index/label_filter.rs +++ b/nodedb-graph/src/csr/index/label_filter.rs @@ -2,46 +2,71 @@ //! An edge-label filter resolved against one CSR partition. //! -//! A filter names a label as a string, and the CSR stores labels as dense -//! ids. A label this partition has never interned carries no edge here, so -//! the filter keeps no durable edge. Falling back to "no filter" instead -//! would return every other label's edges under the caller's label. In a -//! cluster that happens whenever the label lives only on other nodes. +//! A filter names labels as strings, and the CSR stores labels as dense ids. +//! An empty set keeps every edge. A non-empty set keeps an edge whose label is +//! any listed label. A listed label this partition has never interned carries +//! no edge here, so it adds nothing. A set whose labels are all unknown keeps +//! no edge. Widening it to "no filter" returns other labels' edges under the +//! caller's labels. In a cluster that happens whenever the labels live only on +//! other nodes. use super::types::CsrIndex; /// Which durable edges a label filter keeps in one partition. -#[derive(Debug, Clone, Copy, PartialEq, Eq)] +#[derive(Debug, Clone, PartialEq, Eq)] pub enum LabelFilter { /// No filter: every edge. Any, /// Edges with this dense label id. Only(u32), - /// A label this partition has never seen: no edge. + /// Edges with any of these dense label ids. Sorted, deduplicated, two or more. + AnyOf(Box<[u32]>), + /// Every listed label is unknown to this partition: no edge. Unknown, } impl LabelFilter { /// Whether an edge with dense label id `lid` passes the filter. #[inline] - pub fn keeps(self, lid: u32) -> bool { + pub fn keeps(&self, lid: u32) -> bool { match self { LabelFilter::Any => true, - LabelFilter::Only(id) => id == lid, + LabelFilter::Only(id) => *id == lid, + LabelFilter::AnyOf(ids) => ids.binary_search(&lid).is_ok(), LabelFilter::Unknown => false, } } + + /// Whether the filter keeps no edge. + #[inline] + pub fn keeps_none(&self) -> bool { + matches!(self, LabelFilter::Unknown) + } + + /// Whether a staged edge named `label` passes `labels`. Staged edges carry + /// label names, not partition ids. An empty set keeps every edge. + #[inline] + pub fn keeps_name(labels: &[&str], label: &str) -> bool { + labels.is_empty() || labels.contains(&label) + } } impl CsrIndex { - /// Resolve `filter` against this partition's interned labels. - pub fn label_filter(&self, filter: Option<&str>) -> LabelFilter { - match filter { - None => LabelFilter::Any, - Some(label) => match self.label_to_id.get(label) { - Some(&id) => LabelFilter::Only(id), - None => LabelFilter::Unknown, - }, + /// Resolve `labels` against this partition's interned labels. + pub fn label_filter(&self, labels: &[&str]) -> LabelFilter { + if labels.is_empty() { + return LabelFilter::Any; + } + let mut ids: Vec = labels + .iter() + .filter_map(|label| self.label_to_id.get(*label).copied()) + .collect(); + ids.sort_unstable(); + ids.dedup(); + match ids.as_slice() { + [] => LabelFilter::Unknown, + [id] => LabelFilter::Only(*id), + _ => LabelFilter::AnyOf(ids.into_boxed_slice()), } } } @@ -59,14 +84,29 @@ mod tests { csr } + fn csr_three_labels() -> CsrIndex { + let mut csr = CsrIndex::new(test_memory()); + for (label, dst) in [("knows", "b"), ("likes", "c"), ("hates", "d")] { + csr.add_edge("a", label, dst) + .unwrap_or_else(|e| panic!("seed edge: {e}")); + } + csr + } + + fn id(csr: &CsrIndex, label: &str) -> u32 { + csr.label_id(label) + .unwrap_or_else(|| panic!("label {label} interned")) + } + #[test] fn an_unknown_label_keeps_no_edge() { let csr = csr(); - assert_eq!(csr.label_filter(Some("likes")), LabelFilter::Unknown); - assert!(!csr.label_filter(Some("likes")).keeps(0)); - assert!(csr.neighbors("a", Some("likes"), Direction::Out).is_empty()); + assert_eq!(csr.label_filter(&["likes"]), LabelFilter::Unknown); + assert!(!csr.label_filter(&["likes"]).keeps(0)); + assert!(csr.label_filter(&["likes"]).keeps_none()); + assert!(csr.neighbors("a", &["likes"], Direction::Out).is_empty()); assert!( - csr.neighbors_in_collection("a", Some("likes"), Direction::Out, "people") + csr.neighbors_in_collection("a", &["likes"], Direction::Out, "people") .is_empty() ); } @@ -74,7 +114,81 @@ mod tests { #[test] fn no_filter_and_a_known_label_keep_their_edges() { let csr = csr(); - assert_eq!(csr.neighbors("a", None, Direction::Out).len(), 1); - assert_eq!(csr.neighbors("a", Some("knows"), Direction::Out).len(), 1); + assert_eq!(csr.neighbors("a", &[], Direction::Out).len(), 1); + assert_eq!(csr.neighbors("a", &["knows"], Direction::Out).len(), 1); + } + + #[test] + fn an_empty_set_resolves_to_any() { + let csr = csr(); + let filter = csr.label_filter(&[]); + assert_eq!(filter, LabelFilter::Any); + assert!(filter.keeps(0)); + assert!(!filter.keeps_none()); + } + + #[test] + fn one_known_label_resolves_to_only() { + let csr = csr(); + assert_eq!( + csr.label_filter(&["knows"]), + LabelFilter::Only(id(&csr, "knows")) + ); + } + + #[test] + fn a_known_and_an_unknown_label_resolve_to_the_known_one() { + let csr = csr(); + assert_eq!( + csr.label_filter(&["knows", "nope"]), + LabelFilter::Only(id(&csr, "knows")) + ); + } + + #[test] + fn two_known_labels_keep_both_and_drop_a_third() { + let csr = csr_three_labels(); + let filter = csr.label_filter(&["likes", "knows"]); + assert!(matches!(filter, LabelFilter::AnyOf(_))); + assert!(filter.keeps(id(&csr, "knows"))); + assert!(filter.keeps(id(&csr, "likes"))); + assert!(!filter.keeps(id(&csr, "hates"))); + assert!(!filter.keeps_none()); + } + + #[test] + fn all_unknown_labels_resolve_to_unknown() { + let csr = csr(); + assert_eq!(csr.label_filter(&["nope", "never"]), LabelFilter::Unknown); + } + + #[test] + fn duplicate_labels_resolve_to_only() { + let csr = csr(); + assert_eq!( + csr.label_filter(&["knows", "knows"]), + LabelFilter::Only(id(&csr, "knows")) + ); + } + + #[test] + fn keeps_name_matches_staged_labels() { + assert!(LabelFilter::keeps_name(&[], "knows")); + assert!(LabelFilter::keeps_name(&["likes", "knows"], "knows")); + assert!(!LabelFilter::keeps_name(&["likes"], "knows")); + } + + #[test] + fn neighbors_over_a_set_returns_the_union() { + let csr = csr_three_labels(); + let mut got = csr.neighbors("a", &["knows", "likes"], Direction::Out); + got.sort(); + assert_eq!( + got, + vec![ + ("knows".to_string(), "b".to_string()), + ("likes".to_string(), "c".to_string()), + ] + ); } } diff --git a/nodedb-graph/src/csr/index/lookup.rs b/nodedb-graph/src/csr/index/lookup.rs index cce17b63e..b96507d69 100644 --- a/nodedb-graph/src/csr/index/lookup.rs +++ b/nodedb-graph/src/csr/index/lookup.rs @@ -25,8 +25,7 @@ pub(crate) struct DenseAdjacency { } impl CsrIndex { - /// Partition tag assigned at construction. Embedded in every - /// `LocalNodeId` this index produces. + /// Partition tag assigned at construction. Embedded in every `LocalNodeId` this index produces. #[inline] pub fn partition_tag(&self) -> u32 { self.partition_tag @@ -40,11 +39,12 @@ impl CsrIndex { LocalNodeId::new(id, self.partition_tag) } - /// Get immediate neighbors by string name. + /// Get immediate neighbors by string name. An empty `label_filter` keeps + /// every edge. Otherwise an edge whose label is any listed label passes. pub fn neighbors( &self, node: &str, - label_filter: Option<&str>, + label_filter: &[&str], direction: Direction, ) -> Vec<(String, String)> { let Some(&node_id) = self.node_to_id.get(node) else { @@ -79,51 +79,6 @@ impl CsrIndex { result } - /// Get neighbors with multi-label filter. Empty labels = all edges. - pub fn neighbors_multi( - &self, - node: &str, - label_filters: &[&str], - direction: Direction, - ) -> Vec<(String, String)> { - let Some(&node_id) = self.node_to_id.get(node) else { - return Vec::new(); - }; - self.record_access(node_id); - let label_ids: Vec = label_filters - .iter() - .filter_map(|l| self.label_to_id.get(*l).copied()) - .collect(); - // Filters this partition has never seen match no edge here; they must - // not widen the filter to every edge. - let match_label = |lid: u32| label_filters.is_empty() || label_ids.contains(&lid); - - let mut result = Vec::new(); - - if matches!(direction, Direction::Out | Direction::Both) { - for (lid, dst) in self.dense_iter_out(node_id) { - if match_label(lid) { - result.push(( - self.id_to_label[lid as usize].clone(), - self.id_to_node[dst as usize].clone(), - )); - } - } - } - if matches!(direction, Direction::In | Direction::Both) { - for (lid, src) in self.dense_iter_in(node_id) { - if match_label(lid) { - result.push(( - self.id_to_label[lid as usize].clone(), - self.id_to_node[src as usize].clone(), - )); - } - } - } - - result - } - /// Add a node without any edges. Idempotent — returns the existing /// tagged id if the name is already present. /// @@ -193,8 +148,8 @@ impl CsrIndex { /// /// # Errors /// - /// Returns [`GraphError::MemoryBudget`] if the reservation for the - /// three output arrays exceeds the `Graph` engine budget. + /// Returns [`GraphError::MemoryBudget`] if the reservation for the three + /// output arrays exceeds the `Graph` engine budget. pub(crate) fn build_dense( edges: &[Vec<(u32, u32)>], collections: &[Vec], @@ -247,43 +202,35 @@ impl CsrIndex { /// Iterate dense outbound edges for a node as `(label, dst, collection)` /// (raw u32, no tag check, no deletion filter). pub(crate) fn dense_out_edges(&self, node: u32) -> impl Iterator + '_ { - let idx = node as usize; - if idx + 1 >= self.out_offsets.len() { - return Vec::new().into_iter(); - } - let start = self.out_offsets[idx] as usize; - let end = self.out_offsets[idx + 1] as usize; - (start..end) - .map(move |i| { - ( - self.out_labels[i], - self.out_targets[i], - self.out_collections.get(i).copied().unwrap_or(0), - ) - }) - .collect::>() - .into_iter() + let range = self + .out_offsets + .get(node as usize..) + .and_then(|offsets| offsets.first().zip(offsets.get(1))) + .map_or(0..0, |(&start, &end)| start as usize..end as usize); + range.map(move |i| { + ( + self.out_labels[i], + self.out_targets[i], + self.out_collections.get(i).copied().unwrap_or(0), + ) + }) } /// Iterate dense inbound edges for a node as `(label, src, collection)` /// (raw u32, no tag check, no deletion filter). pub(crate) fn dense_in_edges(&self, node: u32) -> impl Iterator + '_ { - let idx = node as usize; - if idx + 1 >= self.in_offsets.len() { - return Vec::new().into_iter(); - } - let start = self.in_offsets[idx] as usize; - let end = self.in_offsets[idx + 1] as usize; - (start..end) - .map(move |i| { - ( - self.in_labels[i], - self.in_targets[i], - self.in_collections.get(i).copied().unwrap_or(0), - ) - }) - .collect::>() - .into_iter() + let range = self + .in_offsets + .get(node as usize..) + .and_then(|offsets| offsets.first().zip(offsets.get(1))) + .map_or(0..0, |(&start, &end)| start as usize..end as usize); + range.map(move |i| { + ( + self.in_labels[i], + self.in_targets[i], + self.in_collections.get(i).copied().unwrap_or(0), + ) + }) } /// Raw u32 iteration over outbound edges (dense + buffer - deleted), @@ -355,23 +302,21 @@ impl CsrIndex { } /// Buffer-only iteration over outbound edges for a node. - pub(crate) fn buffer_out_iter(&self, node: u32) -> std::vec::IntoIter<(u32, u32)> { - let idx = node as usize; - if idx < self.buffer_out.len() { - self.buffer_out[idx].clone().into_iter() - } else { - Vec::new().into_iter() - } + pub(crate) fn buffer_out_iter(&self, node: u32) -> impl Iterator + '_ { + self.buffer_out + .get(node as usize) + .map_or(&[][..], Vec::as_slice) + .iter() + .copied() } /// Buffer-only iteration over inbound edges for a node. - pub(crate) fn buffer_in_iter(&self, node: u32) -> std::vec::IntoIter<(u32, u32)> { - let idx = node as usize; - if idx < self.buffer_in.len() { - self.buffer_in[idx].clone().into_iter() - } else { - Vec::new().into_iter() - } + pub(crate) fn buffer_in_iter(&self, node: u32) -> impl Iterator + '_ { + self.buffer_in + .get(node as usize) + .map_or(&[][..], Vec::as_slice) + .iter() + .copied() } /// Iterate all outbound edges for a tagged node. Yields @@ -453,3 +398,64 @@ impl CsrIndex { self.node_to_id.get(name).copied() } } + +#[cfg(test)] +mod tests { + // Live adjacency iteration across dense and buffered edges. + + use super::CsrIndex; + use crate::test_support::test_memory; + + fn names(csr: &CsrIndex, node: &str) -> Vec<(String, String)> { + csr.iter_out_edges(csr.node_id(node).unwrap()) + .map(|(label, target)| (csr.label_name(label).into(), csr.node_name(target).into())) + .collect() + } + + #[test] + fn live_iterators_preserve_dense_then_buffer_order() { + let mut csr = CsrIndex::new(test_memory()); + csr.add_edge("a", "FIRST", "b").unwrap(); + assert_eq!(names(&csr, "a"), vec![("FIRST".into(), "b".into())]); + csr.compact().unwrap(); + assert_eq!(names(&csr, "a"), vec![("FIRST".into(), "b".into())]); + csr.add_edge("a", "SECOND", "c").unwrap(); + assert_eq!( + names(&csr, "a"), + vec![("FIRST".into(), "b".into()), ("SECOND".into(), "c".into())] + ); + let a = csr.node_id_raw("a").unwrap(); + assert_eq!(csr.iter_out_edges_raw(a).count(), 2); + for target in ["b", "c"] { + let tagged = csr.node_id(target).unwrap(); + let inbound: Vec<_> = csr + .iter_in_edges(tagged) + .map(|(_, source)| csr.node_name(source)) + .collect(); + assert_eq!(inbound, vec!["a"]); + assert_eq!( + csr.iter_in_edges_raw(csr.node_id_raw(target).unwrap()) + .count(), + 1 + ); + } + assert_eq!(csr.iter_out_edges_raw(u32::MAX).count(), 0); + assert_eq!(csr.iter_in_edges_raw(u32::MAX).count(), 0); + } + + #[test] + fn collection_tombstones_remove_only_the_matching_copy() { + let mut csr = CsrIndex::new(test_memory()); + csr.add_edge_in_collection("a", "LINK", "b", "first") + .unwrap(); + csr.add_edge_in_collection("a", "LINK", "b", "second") + .unwrap(); + csr.compact().unwrap(); + csr.remove_edge_in_collection("a", "LINK", "b", "first"); + csr.add_edge("a", "STAGED", "c").unwrap(); + csr.remove_edge("a", "STAGED", "c"); + assert_eq!(names(&csr, "a"), vec![("LINK".into(), "b".into())]); + assert_eq!(csr.iter_in_edges(csr.node_id("b").unwrap()).count(), 1); + assert_eq!(csr.iter_in_edges(csr.node_id("c").unwrap()).count(), 0); + } +} diff --git a/nodedb-graph/src/csr/index/restore.rs b/nodedb-graph/src/csr/index/restore.rs index 7b5b0d4b6..5d9e03878 100644 --- a/nodedb-graph/src/csr/index/restore.rs +++ b/nodedb-graph/src/csr/index/restore.rs @@ -315,7 +315,7 @@ mod tests { Some(2.5) ); assert_eq!(csr.edge_weight_in_collection("a", "L", "b", "c"), Some(9.0)); - assert_eq!(csr.neighbors("a", None, Direction::Out).len(), 1); + assert_eq!(csr.neighbors("a", &[], Direction::Out).len(), 1); } #[test] @@ -330,10 +330,10 @@ mod tests { Some(2.5) ); assert_eq!(csr.edge_weight_in_collection("a", "L", "b", "c"), Some(9.0)); - assert_eq!(csr.neighbors("a", None, Direction::Out).len(), 1); + assert_eq!(csr.neighbors("a", &[], Direction::Out).len(), 1); csr.compact().expect("compact again"); assert_eq!(csr.edge_weight_in_collection("a", "L", "b", "c"), Some(9.0)); - assert_eq!(csr.neighbors("a", None, Direction::Out).len(), 1); + assert_eq!(csr.neighbors("a", &[], Direction::Out).len(), 1); } #[test] @@ -348,7 +348,7 @@ mod tests { csr.restore_edge_in_collection("a", "L", "b", "c", Some(1.0)) .expect("restore"); assert_eq!(csr.edge_weight_in_collection("a", "L", "b", "c"), Some(1.0)); - assert_eq!(csr.neighbors("a", None, Direction::Out).len(), 1); + assert_eq!(csr.neighbors("a", &[], Direction::Out).len(), 1); } #[test] @@ -360,7 +360,7 @@ mod tests { csr.remove_edge_in_collection("a", "L", "b", "c"); csr.add_edge_in_collection("a", "L", "b", "c") .expect("re-add"); - assert_eq!(csr.neighbors("a", None, Direction::Out).len(), 1); + assert_eq!(csr.neighbors("a", &[], Direction::Out).len(), 1); } #[test] @@ -383,7 +383,7 @@ mod tests { csr.add_edge("a", "L", "y") .expect("a later node takes the freed id"); csr.compact().expect("compact"); - assert_eq!(csr.neighbors("a", None, Direction::Out).len(), 2); + assert_eq!(csr.neighbors("a", &[], Direction::Out).len(), 2); } #[test] diff --git a/nodedb-graph/src/csr/index/scoped.rs b/nodedb-graph/src/csr/index/scoped.rs index 8e19fa604..8eb780d71 100644 --- a/nodedb-graph/src/csr/index/scoped.rs +++ b/nodedb-graph/src/csr/index/scoped.rs @@ -13,7 +13,8 @@ use super::types::{CsrIndex, Direction}; impl CsrIndex { /// Collection-scoped neighbor lookup by string name. /// - /// Like [`Self::neighbors`] but only traverses edges tagged with + /// Like [`Self::neighbors`], with the same `label_filter` set semantics, + /// but only traverses edges tagged with /// `collection`. A collection with no edges in this partition yields an /// empty result (never falls back to the merged, cross-collection view). /// Used by collection-scoped GraphRAG BFS so expansion never crosses a @@ -21,7 +22,7 @@ impl CsrIndex { pub fn neighbors_in_collection( &self, node: &str, - label_filter: Option<&str>, + label_filter: &[&str], direction: Direction, collection: &str, ) -> Vec<(String, String)> { diff --git a/nodedb-graph/src/csr/index/types.rs b/nodedb-graph/src/csr/index/types.rs index fb6303f11..fe7fe63b1 100644 --- a/nodedb-graph/src/csr/index/types.rs +++ b/nodedb-graph/src/csr/index/types.rs @@ -223,7 +223,7 @@ mod tests { #[test] fn neighbors_out() { let csr = make_csr(); - let n = csr.neighbors("a", None, Direction::Out); + let n = csr.neighbors("a", &[], Direction::Out); assert_eq!(n.len(), 2); let dsts: Vec<&str> = n.iter().map(|(_, d)| d.as_str()).collect(); assert!(dsts.contains(&"b")); @@ -233,7 +233,7 @@ mod tests { #[test] fn neighbors_filtered() { let csr = make_csr(); - let n = csr.neighbors("a", Some("KNOWS"), Direction::Out); + let n = csr.neighbors("a", &["KNOWS"], Direction::Out); assert_eq!(n.len(), 1); assert_eq!(n[0].1, "b"); } @@ -241,7 +241,7 @@ mod tests { #[test] fn neighbors_in() { let csr = make_csr(); - let n = csr.neighbors("b", None, Direction::In); + let n = csr.neighbors("b", &[], Direction::In); assert_eq!(n.len(), 1); assert_eq!(n[0].1, "a"); } @@ -249,9 +249,9 @@ mod tests { #[test] fn incremental_remove() { let mut csr = make_csr(); - assert_eq!(csr.neighbors("a", Some("KNOWS"), Direction::Out).len(), 1); + assert_eq!(csr.neighbors("a", &["KNOWS"], Direction::Out).len(), 1); csr.remove_edge("a", "KNOWS", "b"); - assert_eq!(csr.neighbors("a", Some("KNOWS"), Direction::Out).len(), 0); + assert_eq!(csr.neighbors("a", &["KNOWS"], Direction::Out).len(), 0); } #[test] @@ -259,7 +259,7 @@ mod tests { let mut csr = CsrIndex::new(test_memory()); csr.add_edge("a", "L", "b").unwrap(); csr.add_edge("a", "L", "b").unwrap(); - assert_eq!(csr.neighbors("a", None, Direction::Out).len(), 1); + assert_eq!(csr.neighbors("a", &[], Direction::Out).len(), 1); } #[test] @@ -267,13 +267,13 @@ mod tests { let mut csr = CsrIndex::new(test_memory()); csr.add_edge("a", "L", "b").unwrap(); csr.add_edge("b", "L", "c").unwrap(); - assert_eq!(csr.neighbors("a", None, Direction::Out).len(), 1); + assert_eq!(csr.neighbors("a", &[], Direction::Out).len(), 1); csr.compact() .expect("test governor ceiling covers this reservation"); assert!(csr.buffer_out.iter().all(|b| b.is_empty())); - assert_eq!(csr.neighbors("a", None, Direction::Out).len(), 1); - assert_eq!(csr.neighbors("b", None, Direction::Out).len(), 1); + assert_eq!(csr.neighbors("a", &[], Direction::Out).len(), 1); + assert_eq!(csr.neighbors("b", &[], Direction::Out).len(), 1); } #[test] @@ -285,12 +285,12 @@ mod tests { .expect("test governor ceiling covers this reservation"); csr.remove_edge("a", "L", "b"); - assert_eq!(csr.neighbors("a", None, Direction::Out).len(), 1); + assert_eq!(csr.neighbors("a", &[], Direction::Out).len(), 1); csr.compact() .expect("test governor ceiling covers this reservation"); - assert_eq!(csr.neighbors("a", None, Direction::Out).len(), 1); - assert_eq!(csr.neighbors("a", None, Direction::Out)[0].1, "c"); + assert_eq!(csr.neighbors("a", &[], Direction::Out).len(), 1); + assert_eq!(csr.neighbors("a", &[], Direction::Out)[0].1, "c"); } #[test] @@ -327,7 +327,7 @@ mod tests { assert_eq!(restored.node_count(), csr.node_count()); assert_eq!(restored.edge_count(), csr.edge_count()); - let n = restored.neighbors("a", Some("KNOWS"), Direction::Out); + let n = restored.neighbors("a", &["KNOWS"], Direction::Out); assert_eq!(n.len(), 1); assert_eq!(n[0].1, "b"); } @@ -362,8 +362,8 @@ mod tests { let removed = csr.remove_node_edges("a"); assert_eq!(removed, 3); - assert_eq!(csr.neighbors("a", None, Direction::Out).len(), 0); - assert_eq!(csr.neighbors("a", None, Direction::In).len(), 0); + assert_eq!(csr.neighbors("a", &[], Direction::Out).len(), 0); + assert_eq!(csr.neighbors("a", &[], Direction::In).len(), 0); } #[test] diff --git a/nodedb-graph/src/csr/persist.rs b/nodedb-graph/src/csr/persist.rs index 7c39b345a..c3192dae6 100644 --- a/nodedb-graph/src/csr/persist.rs +++ b/nodedb-graph/src/csr/persist.rs @@ -424,7 +424,7 @@ mod tests { assert_eq!(restored.edge_count(), 2); assert!(!restored.has_weights()); - let n = restored.neighbors("a", Some("KNOWS"), Direction::Out); + let n = restored.neighbors("a", &["KNOWS"], Direction::Out); assert_eq!(n.len(), 1); assert_eq!(n[0].1, "b"); } diff --git a/nodedb-graph/src/csr/rebuild/install.rs b/nodedb-graph/src/csr/rebuild/install.rs index 740034bc0..52f004a97 100644 --- a/nodedb-graph/src/csr/rebuild/install.rs +++ b/nodedb-graph/src/csr/rebuild/install.rs @@ -103,7 +103,7 @@ mod tests { fn out_of(csr: &CsrIndex, node: &str) -> Vec { let mut dsts: Vec = csr - .neighbors(node, None, Direction::Out) + .neighbors(node, &[], Direction::Out) .into_iter() .map(|(_, d)| d) .collect(); diff --git a/nodedb-graph/src/overlay_delta.rs b/nodedb-graph/src/overlay_delta.rs index 33afcf9f5..74107c7e9 100644 --- a/nodedb-graph/src/overlay_delta.rs +++ b/nodedb-graph/src/overlay_delta.rs @@ -15,6 +15,8 @@ use std::collections::{HashMap, HashSet}; +use crate::csr::index::LabelFilter; + /// Staged, not-yet-durable graph edges and deletes for one traversal scope. /// /// `out_edges`/`in_edges` mirror each other: [`Self::stage_edge`] inserts into @@ -60,37 +62,35 @@ impl GraphOverlayDelta { self.out_edges.is_empty() && self.in_edges.is_empty() && self.tombstones.is_empty() } - /// Staged out-neighbours of `src` as `(label, dst)`, honouring - /// `label_filter` (when `Some`, only edges with that exact label). + /// Staged out-neighbours of `src` as `(label, dst)`. `label_filter` empty + /// keeps every staged edge. Otherwise an edge whose label is any listed + /// label passes. pub fn out_neighbors<'a>( &'a self, src: &str, - label_filter: Option<&'a str>, + label_filter: &'a [&'a str], ) -> impl Iterator { self.out_edges.get(src).into_iter().flat_map(move |edges| { edges .iter() - .filter_map(move |(label, dst)| match label_filter { - Some(f) if f != label => None, - _ => Some((label.as_str(), dst.as_str())), - }) + .filter(move |(label, _)| LabelFilter::keeps_name(label_filter, label)) + .map(|(label, dst)| (label.as_str(), dst.as_str())) }) } - /// Staged in-neighbours of `dst` as `(label, src)`, honouring - /// `label_filter` (when `Some`, only edges with that exact label). + /// Staged in-neighbours of `dst` as `(label, src)`. `label_filter` empty + /// keeps every staged edge. Otherwise an edge whose label is any listed + /// label passes. pub fn in_neighbors<'a>( &'a self, dst: &str, - label_filter: Option<&'a str>, + label_filter: &'a [&'a str], ) -> impl Iterator { self.in_edges.get(dst).into_iter().flat_map(move |edges| { edges .iter() - .filter_map(move |(label, src)| match label_filter { - Some(f) if f != label => None, - _ => Some((label.as_str(), src.as_str())), - }) + .filter(move |(label, _)| LabelFilter::keeps_name(label_filter, label)) + .map(|(label, src)| (label.as_str(), src.as_str())) }) } @@ -109,6 +109,11 @@ impl GraphOverlayDelta { .map(String::as_str) } + /// True when a staged edge has `node` as its source or destination. + pub fn names_node(&self, node: &str) -> bool { + self.out_edges.contains_key(node) || self.in_edges.contains_key(node) + } + /// True when `src -[label]-> dst` has been staged-deleted. pub fn is_tombstoned(&self, src: &str, label: &str, dst: &str) -> bool { self.tombstones @@ -130,9 +135,9 @@ mod tests { let mut d = GraphOverlayDelta::new(); d.stage_edge("a", "knows", "b"); assert!(!d.is_empty()); - let out: Vec<_> = d.out_neighbors("a", None).collect(); + let out: Vec<_> = d.out_neighbors("a", &[]).collect(); assert_eq!(out, vec![("knows", "b")]); - let inn: Vec<_> = d.in_neighbors("b", None).collect(); + let inn: Vec<_> = d.in_neighbors("b", &[]).collect(); assert_eq!(inn, vec![("knows", "a")]); } @@ -141,10 +146,24 @@ mod tests { let mut d = GraphOverlayDelta::new(); d.stage_edge("a", "knows", "b"); d.stage_edge("a", "likes", "c"); - let out: Vec<_> = d.out_neighbors("a", Some("likes")).collect(); + let out: Vec<_> = d.out_neighbors("a", &["likes"]).collect(); assert_eq!(out, vec![("likes", "c")]); } + #[test] + fn a_label_set_keeps_every_listed_label() { + let mut d = GraphOverlayDelta::new(); + d.stage_edge("a", "knows", "b"); + d.stage_edge("a", "likes", "c"); + d.stage_edge("a", "hates", "d"); + let out: Vec<_> = d.out_neighbors("a", &["likes", "knows", "nope"]).collect(); + assert_eq!(out, vec![("knows", "b"), ("likes", "c")]); + let inn: Vec<_> = d.in_neighbors("d", &["knows", "likes"]).collect(); + assert!(inn.is_empty()); + let none: Vec<_> = d.out_neighbors("a", &["nope"]).collect(); + assert!(none.is_empty()); + } + #[test] fn tombstone_only_is_not_empty() { let mut d = GraphOverlayDelta::new(); @@ -153,4 +172,15 @@ mod tests { assert!(d.is_tombstoned("a", "knows", "b")); assert!(!d.is_tombstoned("a", "knows", "z")); } + + #[test] + fn names_node_covers_staged_endpoints_only() { + let mut d = GraphOverlayDelta::new(); + d.stage_edge("a", "knows", "b"); + d.stage_tombstone("x", "knows", "y"); + assert!(d.names_node("a")); + assert!(d.names_node("b")); + assert!(!d.names_node("x")); + assert!(!d.names_node("z")); + } } diff --git a/nodedb-graph/src/path_overlay.rs b/nodedb-graph/src/path_overlay.rs index fd9731ce3..b9aefe242 100644 --- a/nodedb-graph/src/path_overlay.rs +++ b/nodedb-graph/src/path_overlay.rs @@ -18,179 +18,212 @@ //! targets, unioned with `overlay.out_neighbors`). Each backward frontier //! node expands its IN edges symmetrically. Staged edges bypass the frontier //! bitmap: the transaction's own writes have no durable surrogate to gate on. +//! +//! `partition` is `None` for a tenant with no CSR partition. The search then +//! follows only staged edges. use std::collections::HashMap; use std::collections::hash_map::Entry; use crate::csr::CsrIndex; +use crate::csr::index::LabelFilter; use crate::overlay_delta::GraphOverlayDelta; use crate::path_params::ShortestPathParams; +use crate::traversal_overlay::node_present; + +/// String-keyed bidirectional BFS that merges the transaction's staged +/// edges/tombstones. Mirrors the durable `shortest_path_dense` loop +/// structure (alternating one forward then one backward level per depth +/// step) so that, when the overlay contributes nothing, the result and +/// path shape match the durable search. +/// +/// A path from a node to itself exists only when the node exists: the +/// partition holds it, or a staged edge names it. The dense search answers +/// the same. +pub(crate) fn shortest_path_overlay( + partition: Option<&CsrIndex>, + params: ShortestPathParams<'_>, + overlay: &GraphOverlayDelta, +) -> Option> { + let ShortestPathParams { + src, + dst, + label_filter, + max_depth, + max_visited, + frontier_bitmap, + } = params; + if src == dst { + return node_present(partition, overlay, src).then(|| vec![src.to_string()]); + } -impl CsrIndex { - /// String-keyed bidirectional BFS that merges the transaction's staged - /// edges/tombstones. Mirrors the durable `shortest_path_dense` loop - /// structure (alternating one forward then one backward level per depth - /// step) so that, when the overlay contributes nothing, the result and - /// path shape match the durable search. - pub(crate) fn shortest_path_overlay( - &self, - params: ShortestPathParams<'_>, - overlay: &GraphOverlayDelta, - ) -> Option> { - let ShortestPathParams { - src, - dst, - label_filter, - max_depth, - max_visited, - frontier_bitmap, - } = params; - if src == dst { - return Some(vec![src.to_string()]); - } - - // parent maps: node -> the neighbour it was reached from. The endpoint - // maps to itself, marking the reconstruction terminus. - let mut fwd_parent: HashMap = HashMap::new(); - let mut bwd_parent: HashMap = HashMap::new(); - fwd_parent.insert(src.to_string(), src.to_string()); - bwd_parent.insert(dst.to_string(), dst.to_string()); + // parent maps: node -> the neighbour it was reached from. The endpoint + // maps to itself, marking the reconstruction terminus. + let mut fwd_parent: HashMap = HashMap::new(); + let mut bwd_parent: HashMap = HashMap::new(); + fwd_parent.insert(src.to_string(), src.to_string()); + bwd_parent.insert(dst.to_string(), dst.to_string()); - let mut fwd_frontier: Vec = vec![src.to_string()]; - let mut bwd_frontier: Vec = vec![dst.to_string()]; + let mut fwd_frontier: Vec = vec![src.to_string()]; + let mut bwd_frontier: Vec = vec![dst.to_string()]; - let labels = self.label_filter(label_filter); + let gate = PathGate::new(partition, label_filter, frontier_bitmap, [src, dst]); + let within_depth = + |path: Vec| (path.len().saturating_sub(1) <= max_depth).then_some(path); - for _depth in 0..max_depth { - if fwd_parent.len() + bwd_parent.len() >= max_visited { - break; - } + // Round `k` meets on a path of `2k - 1` or `2k` edges, so + // `max_depth.div_ceil(2)` rounds reach every path within `max_depth`. + for _round in 0..max_depth.div_ceil(2) { + if fwd_parent.len() + bwd_parent.len() >= max_visited { + break; + } - // Each level's edges are relaxed in (neighbour, frontier node) - // name order, as the durable search relaxes them. - let mut candidates: Vec<(String, String)> = Vec::new(); - for node in std::mem::take(&mut fwd_frontier) { - for neighbor in - self.forward_neighbors(&node, labels, label_filter, frontier_bitmap, overlay) - { - candidates.push((neighbor, node.clone())); - } + // Each level's edges are relaxed in (neighbour, frontier node) + // name order, as the durable search relaxes them. + let mut candidates: Vec<(String, String)> = Vec::new(); + for node in std::mem::take(&mut fwd_frontier) { + for neighbor in forward_neighbors(&node, &gate, overlay) { + candidates.push((neighbor, node.clone())); } - candidates.sort(); - let mut next_fwd = Vec::new(); - for (neighbor, node) in candidates { - if let Some(meeting) = relax( - &neighbor, - &node, - &mut fwd_parent, - &bwd_parent, - &mut next_fwd, - ) { - return Some(reconstruct(&meeting, &fwd_parent, &bwd_parent)); - } + } + candidates.sort(); + let mut next_fwd = Vec::new(); + for (neighbor, node) in candidates { + if let Some(meeting) = relax( + &neighbor, + &node, + &mut fwd_parent, + &bwd_parent, + &mut next_fwd, + ) { + return within_depth(reconstruct(&meeting, &fwd_parent, &bwd_parent)); } - fwd_frontier = next_fwd; - - let mut candidates: Vec<(String, String)> = Vec::new(); - for node in std::mem::take(&mut bwd_frontier) { - for neighbor in - self.backward_neighbors(&node, labels, label_filter, frontier_bitmap, overlay) - { - candidates.push((neighbor, node.clone())); - } + } + fwd_frontier = next_fwd; + + let mut candidates: Vec<(String, String)> = Vec::new(); + for node in std::mem::take(&mut bwd_frontier) { + for neighbor in backward_neighbors(&node, &gate, overlay) { + candidates.push((neighbor, node.clone())); } - candidates.sort(); - let mut next_bwd = Vec::new(); - for (neighbor, node) in candidates { - if let Some(meeting) = relax( - &neighbor, - &node, - &mut bwd_parent, - &fwd_parent, - &mut next_bwd, - ) { - return Some(reconstruct(&meeting, &fwd_parent, &bwd_parent)); - } + } + candidates.sort(); + let mut next_bwd = Vec::new(); + for (neighbor, node) in candidates { + if let Some(meeting) = relax( + &neighbor, + &node, + &mut bwd_parent, + &fwd_parent, + &mut next_bwd, + ) { + return within_depth(reconstruct(&meeting, &fwd_parent, &bwd_parent)); } - bwd_frontier = next_bwd; + } + bwd_frontier = next_bwd; - if fwd_frontier.is_empty() && bwd_frontier.is_empty() { - break; - } + if fwd_frontier.is_empty() && bwd_frontier.is_empty() { + break; } - None } + None +} - /// OUT neighbours of `node`: durable CSR out edges (skipping staged - /// tombstones and bitmap-excluded durable targets) unioned with staged - /// out edges. - fn forward_neighbors( - &self, - node: &str, - labels: crate::csr::index::LabelFilter, - label_filter: Option<&str>, - frontier_bitmap: Option<&nodedb_types::SurrogateBitmap>, - overlay: &GraphOverlayDelta, - ) -> Vec { - let mut out = Vec::new(); - if let Some(&node_id) = self.node_to_id.get(node) { - self.record_access(node_id); - for (lid, dst) in self.dense_iter_out(node_id) { - if !labels.keeps(lid) { - continue; - } - let dst_name = &self.id_to_node[dst as usize]; - if overlay.is_tombstoned(node, self.label_name(lid), dst_name) { - continue; - } - if !frontier_bitmap.is_none_or(|bm| { - bm.contains(nodedb_types::Surrogate::new(self.node_surrogate_raw(dst))) - }) { - continue; - } - out.push(dst_name.clone()); +/// OUT neighbours of `node`: durable CSR out edges (skipping staged +/// tombstones and gated durable targets) unioned with staged out edges. +fn forward_neighbors(node: &str, gate: &PathGate<'_>, overlay: &GraphOverlayDelta) -> Vec { + let mut out = Vec::new(); + if let Some((csr, durable_labels)) = gate.durable.as_ref() + && let Some(&node_id) = csr.node_to_id.get(node) + { + csr.record_access(node_id); + for (lid, dst) in csr.dense_iter_out(node_id) { + if !durable_labels.keeps(lid) { + continue; } + let dst_name = &csr.id_to_node[dst as usize]; + if overlay.is_tombstoned(node, csr.label_name(lid), dst_name) { + continue; + } + if !gate.admits_durable(csr, dst, dst_name) { + continue; + } + out.push(dst_name.clone()); } - for (_, dst) in overlay.out_neighbors(node, label_filter) { - out.push(dst.to_string()); - } - out } + out.extend( + overlay + .out_neighbors(node, gate.labels) + .map(|(_, dst)| dst.to_string()), + ); + out +} - /// IN neighbours of `node`: durable CSR in edges (skipping staged - /// tombstones and bitmap-excluded durable sources) unioned with staged - /// in edges. - fn backward_neighbors( - &self, - node: &str, - labels: crate::csr::index::LabelFilter, - label_filter: Option<&str>, - frontier_bitmap: Option<&nodedb_types::SurrogateBitmap>, - overlay: &GraphOverlayDelta, - ) -> Vec { - let mut out = Vec::new(); - if let Some(&node_id) = self.node_to_id.get(node) { - self.record_access(node_id); - for (lid, src) in self.dense_iter_in(node_id) { - if !labels.keeps(lid) { - continue; - } - let src_name = &self.id_to_node[src as usize]; - if overlay.is_tombstoned(src_name, self.label_name(lid), node) { - continue; - } - if !frontier_bitmap.is_none_or(|bm| { - bm.contains(nodedb_types::Surrogate::new(self.node_surrogate_raw(src))) - }) { - continue; - } - out.push(src_name.clone()); +/// IN neighbours of `node`: durable CSR in edges (skipping staged +/// tombstones and gated durable sources) unioned with staged in edges. +fn backward_neighbors(node: &str, gate: &PathGate<'_>, overlay: &GraphOverlayDelta) -> Vec { + let mut out = Vec::new(); + if let Some((csr, durable_labels)) = gate.durable.as_ref() + && let Some(&node_id) = csr.node_to_id.get(node) + { + csr.record_access(node_id); + for (lid, src) in csr.dense_iter_in(node_id) { + if !durable_labels.keeps(lid) { + continue; + } + let src_name = &csr.id_to_node[src as usize]; + if overlay.is_tombstoned(src_name, csr.label_name(lid), node) { + continue; } + if !gate.admits_durable(csr, src, src_name) { + continue; + } + out.push(src_name.clone()); } - for (_, src) in overlay.in_neighbors(node, label_filter) { - out.push(src.to_string()); + } + out.extend( + overlay + .in_neighbors(node, gate.labels) + .map(|(_, src)| src.to_string()), + ); + out +} + +/// Which edges and nodes one overlay path search may cross. +struct PathGate<'a> { + /// Staged-edge label names. Empty keeps every label. + labels: &'a [&'a str], + /// The durable partition, and `labels` resolved against it. A label the + /// partition has never seen keeps no durable edge, so it never widens + /// the filter. `None` when the tenant has no partition. + durable: Option<(&'a CsrIndex, LabelFilter)>, + frontier_bitmap: Option<&'a nodedb_types::SurrogateBitmap>, + /// The source and destination, which the bitmap never gates. + endpoints: [&'a str; 2], +} + +impl<'a> PathGate<'a> { + fn new( + partition: Option<&'a CsrIndex>, + label_filter: &'a [&'a str], + frontier_bitmap: Option<&'a nodedb_types::SurrogateBitmap>, + endpoints: [&'a str; 2], + ) -> Self { + Self { + labels: label_filter, + durable: partition.map(|csr| (csr, csr.label_filter(label_filter))), + frontier_bitmap, + endpoints, } - out + } + + /// Staged edges bypass the bitmap: the transaction's own writes have no + /// durable surrogate to gate on. + fn admits_durable(&self, csr: &CsrIndex, id: u32, name: &str) -> bool { + self.endpoints.contains(&name) + || self.frontier_bitmap.is_none_or(|bm| { + bm.contains(nodedb_types::Surrogate::new(csr.node_surrogate_raw(id))) + }) } } @@ -270,7 +303,7 @@ mod tests { fn params<'a>( src: &'a str, dst: &'a str, - label_filter: Option<&'a str>, + label_filter: &'a [&'a str], max_depth: usize, ) -> ShortestPathParams<'a> { ShortestPathParams { @@ -283,6 +316,29 @@ mod tests { } } + /// A bitmap that holds neither endpoint still admits the durable edge + /// between them while a staged edge forces the overlay search. + #[test] + fn the_frontier_bitmap_never_gates_the_endpoints() { + use nodedb_types::{Surrogate, SurrogateBitmap}; + + let mut csr = CsrIndex::new(test_memory()); + csr.add_edge("a", "KNOWS", "c").unwrap(); + csr.set_node_surrogate("a", Surrogate::new(1)); + csr.set_node_surrogate("c", Surrogate::new(3)); + let mut bm = SurrogateBitmap::new(); + bm.insert(Surrogate::new(99)); + let mut ov = GraphOverlayDelta::new(); + ov.stage_edge("x", "KNOWS", "y"); + + let mut p = params("a", "c", &["KNOWS"], 1); + p.frontier_bitmap = Some(&bm); + assert_eq!( + csr.shortest_path(p, Some(&ov)), + Some(vec!["a".to_string(), "c".to_string()]) + ); + } + /// A staged edge completes a path the durable graph lacks: durable A->B, /// staged B->C, so A->B->C must be found only with the overlay. #[test] @@ -293,13 +349,13 @@ mod tests { ov.stage_edge("b", "KNOWS", "c"); let path = csr - .shortest_path(params("a", "c", Some("KNOWS"), 5), Some(&ov)) + .shortest_path(params("a", "c", &["KNOWS"], 5), Some(&ov)) .expect("staged edge should complete the path"); assert_eq!(path, vec!["a", "b", "c"]); // Without the overlay the path does not exist. assert!( - csr.shortest_path(params("a", "c", Some("KNOWS"), 5), None) + csr.shortest_path(params("a", "c", &["KNOWS"], 5), None) .is_none() ); } @@ -315,7 +371,7 @@ mod tests { ov.stage_edge("x", "KNOWS", "d"); let path = csr - .shortest_path(params("a", "d", Some("KNOWS"), 5), Some(&ov)) + .shortest_path(params("a", "d", &["KNOWS"], 5), Some(&ov)) .expect("staged-only path should be found"); assert_eq!(path, vec!["a", "x", "d"]); } @@ -333,7 +389,7 @@ mod tests { ov.stage_tombstone("a", "KNOWS", "d"); let path = csr - .shortest_path(params("a", "d", Some("KNOWS"), 5), Some(&ov)) + .shortest_path(params("a", "d", &["KNOWS"], 5), Some(&ov)) .expect("detour path should be found"); assert_eq!(path, vec!["a", "b", "d"]); } @@ -347,7 +403,7 @@ mod tests { ov.stage_tombstone("a", "KNOWS", "d"); assert!( - csr.shortest_path(params("a", "d", Some("KNOWS"), 5), Some(&ov)) + csr.shortest_path(params("a", "d", &["KNOWS"], 5), Some(&ov)) .is_none() ); } @@ -365,12 +421,33 @@ mod tests { // Empty overlay dispatches to dense; force the overlay code path by // calling it directly, then compare with dense. let dense = csr - .shortest_path(params("a", "d", Some("KNOWS"), 10), None) + .shortest_path(params("a", "d", &["KNOWS"], 10), None) .unwrap(); - let overlaid = csr.shortest_path_overlay(params("a", "d", Some("KNOWS"), 10), &ov); + let overlaid = + super::shortest_path_overlay(Some(&csr), params("a", "d", &["KNOWS"], 10), &ov); assert_eq!(overlaid, Some(dense)); } + /// Durable `a -FIRST-> b`, staged `b -SECOND-> d` and `b -OTHER-> d`. The + /// set follows the staged edge under its second label only. + #[test] + fn a_staged_edge_under_the_second_label_of_a_set_is_followed() { + let mut csr = CsrIndex::new(test_memory()); + csr.add_edge("a", "FIRST", "b").unwrap(); + let mut ov = GraphOverlayDelta::new(); + ov.stage_edge("b", "SECOND", "d"); + ov.stage_edge("b", "OTHER", "x"); + + assert_eq!( + csr.shortest_path(params("a", "d", &["FIRST", "SECOND"], 5), Some(&ov)), + Some(vec!["a".to_string(), "b".to_string(), "d".to_string()]) + ); + assert!( + csr.shortest_path(params("a", "x", &["FIRST", "SECOND"], 5), Some(&ov)) + .is_none() + ); + } + #[test] fn src_equals_dst() { let mut csr = CsrIndex::new(test_memory()); @@ -378,11 +455,50 @@ mod tests { let mut ov = GraphOverlayDelta::new(); ov.stage_edge("a", "KNOWS", "x"); let path = csr - .shortest_path(params("a", "a", None, 5), Some(&ov)) + .shortest_path(params("a", "a", &[], 5), Some(&ov)) .unwrap(); assert_eq!(path, vec!["a"]); } + /// A node the partition lacks and no staged edge names has no path to + /// itself, as in the dense search. A staged edge naming it makes it + /// exist. + #[test] + fn src_equals_dst_on_an_absent_node_is_none() { + let mut csr = CsrIndex::new(test_memory()); + csr.add_edge("a", "KNOWS", "b").unwrap(); + let mut ov = GraphOverlayDelta::new(); + ov.stage_edge("x", "KNOWS", "y"); + assert!( + csr.shortest_path(params("ghost", "ghost", &[], 5), Some(&ov)) + .is_none() + ); + assert!( + csr.shortest_path(params("ghost", "ghost", &[], 5), None) + .is_none() + ); + assert_eq!( + csr.shortest_path(params("y", "y", &[], 5), Some(&ov)), + Some(vec!["y".to_string()]) + ); + } + + /// With no partition, the search follows staged edges only. + #[test] + fn a_missing_partition_walks_staged_edges() { + let mut ov = GraphOverlayDelta::new(); + ov.stage_edge("a", "KNOWS", "x"); + ov.stage_edge("x", "KNOWS", "d"); + assert_eq!( + CsrIndex::shortest_path_on(None, params("a", "d", &["KNOWS"], 5), Some(&ov)), + Some(vec!["a".to_string(), "x".to_string(), "d".to_string()]) + ); + assert!(CsrIndex::shortest_path_on(None, params("a", "d", &[], 5), None).is_none()); + assert!( + CsrIndex::shortest_path_on(None, params("ghost", "ghost", &[], 5), Some(&ov)).is_none() + ); + } + #[test] fn unreachable_is_none() { let mut csr = CsrIndex::new(test_memory()); @@ -390,7 +506,7 @@ mod tests { let mut ov = GraphOverlayDelta::new(); ov.stage_edge("m", "KNOWS", "n"); assert!( - csr.shortest_path(params("a", "n", Some("KNOWS"), 5), Some(&ov)) + csr.shortest_path(params("a", "n", &["KNOWS"], 5), Some(&ov)) .is_none() ); } diff --git a/nodedb-graph/src/path_params.rs b/nodedb-graph/src/path_params.rs index 5e896365d..9f18c78a6 100644 --- a/nodedb-graph/src/path_params.rs +++ b/nodedb-graph/src/path_params.rs @@ -10,7 +10,8 @@ pub struct ShortestPathParams<'a> { pub src: &'a str, pub dst: &'a str, - pub label_filter: Option<&'a str>, + /// Empty permits every label. Otherwise, any listed label matches. + pub label_filter: &'a [&'a str], pub max_depth: usize, pub max_visited: usize, pub frontier_bitmap: Option<&'a nodedb_types::SurrogateBitmap>, diff --git a/nodedb-graph/src/test_support.rs b/nodedb-graph/src/test_support.rs index ce06bb24e..5ae557d8c 100644 --- a/nodedb-graph/src/test_support.rs +++ b/nodedb-graph/src/test_support.rs @@ -27,3 +27,29 @@ pub(crate) fn test_memory() -> ScopedMemory { EngineId::Graph, ) } + +/// `a -KNOWS-> b -KNOWS-> c -KNOWS-> d` plus `a -WORKS-> e`. +pub(crate) fn chain_csr() -> crate::csr::CsrIndex { + let mut csr = crate::csr::CsrIndex::new(test_memory()); + for (src, label, dst) in [ + ("a", "KNOWS", "b"), + ("b", "KNOWS", "c"), + ("c", "KNOWS", "d"), + ("a", "WORKS", "e"), + ] { + csr.add_edge(src, label, dst).expect("test edge"); + } + csr +} + +/// `n0 -NEXT-> n1 -> ... -> n999`, compacted. +pub(crate) fn long_chain_csr() -> crate::csr::CsrIndex { + let mut csr = crate::csr::CsrIndex::new(test_memory()); + for i in 0..999 { + csr.add_edge(&format!("n{i}"), "NEXT", &format!("n{}", i + 1)) + .expect("test edge"); + } + csr.compact() + .expect("test governor ceiling covers this reservation"); + csr +} diff --git a/nodedb-graph/src/traversal.rs b/nodedb-graph/src/traversal.rs deleted file mode 100644 index 5e5393738..000000000 --- a/nodedb-graph/src/traversal.rs +++ /dev/null @@ -1,764 +0,0 @@ -// SPDX-License-Identifier: Apache-2.0 - -//! Graph traversal algorithms on the CSR index. -//! -//! BFS, bidirectional shortest path, and subgraph materialization. -//! All algorithms respect a max-visited cap to prevent supernode fan-out -//! explosion from consuming unbounded memory. -//! -//! Access tracking and prefetch hints are integrated: each traversal records -//! node access for hot/cold partition decisions, and prefetches frontier -//! neighbors for cache efficiency. - -use std::collections::{HashMap, HashSet, VecDeque, hash_map::Entry}; - -pub use nodedb_types::config::tuning::DEFAULT_MAX_VISITED; - -use crate::bfs_params::BfsParams; -use crate::csr::{CsrIndex, Direction}; -use crate::overlay_delta::GraphOverlayDelta; -use crate::path_params::ShortestPathParams; - -impl CsrIndex { - /// BFS traversal. Returns all reachable node IDs within max_depth hops. - /// - /// `max_visited` caps the number of nodes visited to prevent supernode fan-out - /// explosion. Pass [`DEFAULT_MAX_VISITED`] for the standard limit. - /// - /// `frontier_bitmap`: when `Some`, only nodes whose surrogate is present in the - /// bitmap are eligible as traversal targets. Start nodes are not gated — only - /// newly discovered frontier nodes are checked. - /// - /// `overlay`: when `Some` and non-empty, the traversal observes the - /// transaction's staged edge writes/deletes (read-your-own-writes), - /// including through nodes reachable only via a staged edge. When `None` - /// or empty, the durable-only dense fast path runs unchanged. - pub fn traverse_bfs( - &self, - params: BfsParams<'_>, - overlay: Option<&GraphOverlayDelta>, - ) -> Vec { - match overlay { - Some(ov) if !ov.is_empty() => self.traverse_bfs_overlay(params, ov), - _ => self.traverse_bfs_dense(params), - } - } - - /// Durable-only BFS over the dense u32 CSR ids. - /// - /// The walk runs level by level. Each level's new nodes are admitted in - /// node-name order until `max_visited` nodes are visited, so a capped walk - /// admits the same nodes however the edges are stored. A cluster - /// coordinator walking the same edges across partitions admits the same - /// nodes. - fn traverse_bfs_dense(&self, params: BfsParams<'_>) -> Vec { - let BfsParams { - start_nodes, - label_filter, - direction, - max_depth, - max_visited, - frontier_bitmap, - } = params; - let labels = self.label_filter(label_filter); - let in_bitmap = |id: u32| { - frontier_bitmap.is_none_or(|bm| { - bm.contains(nodedb_types::Surrogate::new(self.node_surrogate_raw(id))) - }) - }; - let mut visited: HashSet = HashSet::new(); - let mut frontier: Vec = Vec::new(); - for &node in start_nodes { - if let Some(&id) = self.node_to_id.get(node) - && visited.insert(id) - { - frontier.push(id); - } - } - - for _depth in 0..max_depth { - if frontier.is_empty() || visited.len() >= max_visited { - break; - } - let mut candidates: Vec = Vec::new(); - for &node_id in &frontier { - // Track access for hot/cold partition decisions. - self.record_access(node_id); - if matches!(direction, Direction::Out | Direction::Both) { - for (lid, dst) in self.dense_iter_out(node_id) { - if labels.keeps(lid) && !visited.contains(&dst) && in_bitmap(dst) { - candidates.push(dst); - } - } - } - if matches!(direction, Direction::In | Direction::Both) { - for (lid, src) in self.dense_iter_in(node_id) { - if labels.keeps(lid) && !visited.contains(&src) && in_bitmap(src) { - candidates.push(src); - } - } - } - } - frontier = self.admit_by_name(candidates, &mut visited, max_visited); - } - - visited - .into_iter() - .map(|id| self.id_to_node[id as usize].clone()) - .collect() - } - - /// Admit `candidates` into `visited` in node-name order, until `visited` - /// holds `max_visited` nodes. Returns the nodes admitted, in name order. - pub(crate) fn admit_by_name( - &self, - mut candidates: Vec, - visited: &mut HashSet, - max_visited: usize, - ) -> Vec { - self.sort_by_name(&mut candidates); - candidates.dedup(); - let mut admitted = Vec::with_capacity(candidates.len()); - for id in candidates { - if visited.len() >= max_visited { - break; - } - if visited.insert(id) { - self.prefetch_node(id); - admitted.push(id); - } - } - admitted - } - - /// BFS traversal returning nodes with depth information. - /// - /// `max_visited` caps the number of nodes visited to prevent supernode fan-out - /// explosion. Pass [`DEFAULT_MAX_VISITED`] for the standard limit. - pub fn traverse_bfs_with_depth( - &self, - start_nodes: &[&str], - label_filter: Option<&str>, - direction: Direction, - max_depth: usize, - max_visited: usize, - ) -> Vec<(String, u8)> { - let filters: Vec<&str> = label_filter.into_iter().collect(); - self.traverse_bfs_with_depth_multi(start_nodes, &filters, direction, max_depth, max_visited) - } - - /// BFS traversal with multi-label filter. Empty labels = all edges. - /// - /// `max_visited` caps the number of nodes visited to prevent supernode fan-out - /// explosion. Pass [`DEFAULT_MAX_VISITED`] for the standard limit. - pub fn traverse_bfs_with_depth_multi( - &self, - start_nodes: &[&str], - label_filters: &[&str], - direction: Direction, - max_depth: usize, - max_visited: usize, - ) -> Vec<(String, u8)> { - let label_ids: Vec = label_filters - .iter() - .filter_map(|l| self.label_id(l)) - .collect(); - // Filters this partition has never seen match no edge here; they must - // not widen the filter to every edge. - let match_label = |lid: u32| label_filters.is_empty() || label_ids.contains(&lid); - let mut visited: HashMap = HashMap::new(); - let mut queue: VecDeque<(u32, u8)> = VecDeque::new(); - - for &node in start_nodes { - if let Some(&id) = self.node_to_id.get(node) { - visited.insert(id, 0); - queue.push_back((id, 0)); - } - } - - while let Some((node_id, depth)) = queue.pop_front() { - if depth as usize >= max_depth || visited.len() >= max_visited { - continue; - } - - let next_depth = depth + 1; - - if matches!(direction, Direction::Out | Direction::Both) { - for (lid, dst) in self.dense_iter_out(node_id) { - if match_label(lid) - && visited.len() < max_visited - && !visited.contains_key(&dst) - { - visited.insert(dst, next_depth); - queue.push_back((dst, next_depth)); - } - } - } - if matches!(direction, Direction::In | Direction::Both) { - for (lid, src) in self.dense_iter_in(node_id) { - if match_label(lid) - && visited.len() < max_visited - && !visited.contains_key(&src) - { - visited.insert(src, next_depth); - queue.push_back((src, next_depth)); - } - } - } - } - - visited - .into_iter() - .map(|(id, depth)| (self.id_to_node[id as usize].clone(), depth)) - .collect() - } - - /// Shortest path via bidirectional BFS. - /// - /// `max_visited` caps the combined forward+backward visited set to prevent - /// supernode fan-out explosion. Pass [`DEFAULT_MAX_VISITED`] for the standard limit. - /// - /// `frontier_bitmap`: when `Some`, only nodes whose surrogate is present in the - /// bitmap are eligible for expansion. Start and end nodes are not gated. - /// - /// `overlay`: when `Some` and non-empty, the search observes the - /// transaction's staged edge writes/deletes (read-your-own-writes), - /// including a path that must pass through a node reachable only via a - /// staged edge. When `None` or empty, the durable-only dense bidirectional - /// fast path runs unchanged. - pub fn shortest_path( - &self, - params: ShortestPathParams<'_>, - overlay: Option<&GraphOverlayDelta>, - ) -> Option> { - match overlay { - Some(ov) if !ov.is_empty() => self.shortest_path_overlay(params, ov), - _ => self.shortest_path_dense(params), - } - } - - /// Durable-only bidirectional BFS over the dense u32 CSR ids. - /// - /// Each step expands one forward level, then one backward level. A level - /// relaxes its edges in `(neighbour, frontier node)` name order, and the - /// search stops at the first node both sides reached. The cap is checked - /// before each step. - fn shortest_path_dense(&self, params: ShortestPathParams<'_>) -> Option> { - let ShortestPathParams { - src, - dst, - label_filter, - max_depth, - max_visited, - frontier_bitmap, - } = params; - let src_id = *self.node_to_id.get(src)?; - let dst_id = *self.node_to_id.get(dst)?; - if src_id == dst_id { - return Some(vec![src.to_string()]); - } - - let labels = self.label_filter(label_filter); - let in_bitmap = |id: u32| { - frontier_bitmap.is_none_or(|bm| { - bm.contains(nodedb_types::Surrogate::new(self.node_surrogate_raw(id))) - }) - }; - let mut fwd_parent: HashMap = HashMap::new(); - let mut bwd_parent: HashMap = HashMap::new(); - fwd_parent.insert(src_id, src_id); - bwd_parent.insert(dst_id, dst_id); - - let mut fwd_frontier: Vec = vec![src_id]; - let mut bwd_frontier: Vec = vec![dst_id]; - - for _depth in 0..max_depth { - if fwd_parent.len() + bwd_parent.len() >= max_visited { - break; - } - - // Each level's edges are relaxed in (neighbour, frontier node) - // name order, so the parent a node gets, and the meeting point, - // do not depend on how the edges are stored. A cluster coordinator - // relaxes cross-shard hops in the same order. - let mut candidates: Vec<(u32, u32)> = Vec::new(); - for &node in &fwd_frontier { - self.record_access(node); - for (lid, neighbor) in self.dense_iter_out(node) { - if labels.keeps(lid) && in_bitmap(neighbor) { - candidates.push((neighbor, node)); - } - } - } - self.sort_edges_by_name(&mut candidates); - let mut next_fwd = Vec::new(); - for (neighbor, node) in candidates { - if let Entry::Vacant(e) = fwd_parent.entry(neighbor) { - e.insert(node); - next_fwd.push(neighbor); - } - if bwd_parent.contains_key(&neighbor) { - return Some(self.reconstruct_path(neighbor, &fwd_parent, &bwd_parent)); - } - } - fwd_frontier = next_fwd; - - let mut candidates: Vec<(u32, u32)> = Vec::new(); - for &node in &bwd_frontier { - self.record_access(node); - for (lid, neighbor) in self.dense_iter_in(node) { - if labels.keeps(lid) && in_bitmap(neighbor) { - candidates.push((neighbor, node)); - } - } - } - self.sort_edges_by_name(&mut candidates); - let mut next_bwd = Vec::new(); - for (neighbor, node) in candidates { - if let Entry::Vacant(e) = bwd_parent.entry(neighbor) { - e.insert(node); - next_bwd.push(neighbor); - } - if fwd_parent.contains_key(&neighbor) { - return Some(self.reconstruct_path(neighbor, &fwd_parent, &bwd_parent)); - } - } - bwd_frontier = next_bwd; - - if fwd_frontier.is_empty() && bwd_frontier.is_empty() { - break; - } - } - None - } - - /// Order `(neighbour, frontier node)` edges by the two node names. - fn sort_edges_by_name(&self, edges: &mut [(u32, u32)]) { - edges.sort_by(|a, b| { - let name = |id: u32| self.node_name_checked(id); - (name(a.0), name(a.1)).cmp(&(name(b.0), name(b.1))) - }); - } - - fn reconstruct_path( - &self, - meeting: u32, - fwd_parent: &HashMap, - bwd_parent: &HashMap, - ) -> Vec { - let mut fwd_path = Vec::new(); - let mut current = meeting; - loop { - fwd_path.push(current); - let parent = fwd_parent[¤t]; - if parent == current { - break; - } - current = parent; - } - fwd_path.reverse(); - - current = bwd_parent[&meeting]; - if current != meeting { - loop { - fwd_path.push(current); - let parent = bwd_parent[¤t]; - if parent == current { - break; - } - current = parent; - } - } - - fwd_path - .into_iter() - .map(|id| self.id_to_node[id as usize].clone()) - .collect() - } - - /// Materialize a subgraph as `(src, label, dst)` edge tuples within - /// max_depth, expanding in `direction`. - /// - /// `max_visited` caps the number of nodes visited to prevent supernode fan-out - /// explosion. Pass [`DEFAULT_MAX_VISITED`] for the standard limit. - /// - /// `overlay`: when `Some` and non-empty, staged edges are included and - /// staged tombstones subtract durable edges (read-your-own-writes), - /// including through staged-only intermediate nodes. When `None` or - /// empty, the durable-only dense path runs unchanged. - pub fn subgraph( - &self, - start_nodes: &[&str], - label_filter: Option<&str>, - direction: Direction, - max_depth: usize, - max_visited: usize, - overlay: Option<&GraphOverlayDelta>, - ) -> Vec<(String, String, String)> { - match overlay { - Some(ov) if !ov.is_empty() => self.subgraph_overlay( - start_nodes, - label_filter, - direction, - max_depth, - max_visited, - ov, - ), - _ => self.subgraph_dense(start_nodes, label_filter, direction, max_depth, max_visited), - } - } - - /// Durable-only subgraph materialization over the dense u32 CSR ids. - fn subgraph_dense( - &self, - start_nodes: &[&str], - label_filter: Option<&str>, - direction: Direction, - max_depth: usize, - max_visited: usize, - ) -> Vec<(String, String, String)> { - let labels = self.label_filter(label_filter); - let mut visited: HashSet = HashSet::new(); - let mut frontier: Vec = Vec::new(); - let mut edges = Vec::new(); - - for &node in start_nodes { - if let Some(&id) = self.node_to_id.get(node) - && visited.insert(id) - { - frontier.push(id); - } - } - - // Level by level, as `traverse_bfs_dense`: every frontier node's edges - // are recorded, then the level's new nodes are admitted in name order. - for _depth in 0..max_depth { - if frontier.is_empty() || visited.len() >= max_visited { - break; - } - let mut candidates: Vec = Vec::new(); - for &node_id in &frontier { - self.record_access(node_id); - if matches!(direction, Direction::Out | Direction::Both) { - for (lid, dst) in self.dense_iter_out(node_id) { - if labels.keeps(lid) { - edges.push(( - self.id_to_node[node_id as usize].clone(), - self.label_name(lid).to_string(), - self.id_to_node[dst as usize].clone(), - )); - if !visited.contains(&dst) { - candidates.push(dst); - } - } - } - } - if matches!(direction, Direction::In | Direction::Both) { - for (lid, src) in self.dense_iter_in(node_id) { - if labels.keeps(lid) { - edges.push(( - self.id_to_node[src as usize].clone(), - self.label_name(lid).to_string(), - self.id_to_node[node_id as usize].clone(), - )); - if !visited.contains(&src) { - candidates.push(src); - } - } - } - } - } - frontier = self.admit_by_name(candidates, &mut visited, max_visited); - } - - edges - } -} - -#[cfg(test)] -mod tests { - use super::*; - use crate::test_support::test_memory; - - fn make_csr() -> CsrIndex { - let mut csr = CsrIndex::new(test_memory()); - csr.add_edge("a", "KNOWS", "b").unwrap(); - csr.add_edge("b", "KNOWS", "c").unwrap(); - csr.add_edge("c", "KNOWS", "d").unwrap(); - csr.add_edge("a", "WORKS", "e").unwrap(); - csr - } - - #[test] - fn bfs_traversal() { - let csr = make_csr(); - let mut result = csr.traverse_bfs( - BfsParams { - start_nodes: &["a"], - label_filter: Some("KNOWS"), - direction: Direction::Out, - max_depth: 2, - max_visited: DEFAULT_MAX_VISITED, - frontier_bitmap: None, - }, - None, - ); - result.sort(); - assert_eq!(result, vec!["a", "b", "c"]); - } - - #[test] - fn bfs_all_labels() { - let csr = make_csr(); - let mut result = csr.traverse_bfs( - BfsParams { - start_nodes: &["a"], - label_filter: None, - direction: Direction::Out, - max_depth: 1, - max_visited: DEFAULT_MAX_VISITED, - frontier_bitmap: None, - }, - None, - ); - result.sort(); - assert_eq!(result, vec!["a", "b", "e"]); - } - - /// `a` points at `z`, `m` and `b`, stored in that order. A cap of 3 leaves - /// room for two of them, and name order picks `b` and `m`. - fn fan_out_csr() -> CsrIndex { - let mut csr = CsrIndex::new(test_memory()); - csr.add_edge("a", "L", "z").unwrap(); - csr.add_edge("a", "L", "m").unwrap(); - csr.add_edge("a", "L", "b").unwrap(); - csr - } - - #[test] - fn a_capped_bfs_admits_each_level_in_name_order() { - let csr = fan_out_csr(); - let mut result = csr.traverse_bfs( - BfsParams { - start_nodes: &["a"], - label_filter: None, - direction: Direction::Out, - max_depth: 2, - max_visited: 3, - frontier_bitmap: None, - }, - None, - ); - result.sort(); - assert_eq!(result, vec!["a", "b", "m"]); - } - - #[test] - fn a_capped_subgraph_expands_only_admitted_levels() { - let mut csr = fan_out_csr(); - csr.add_edge("z", "L", "y").unwrap(); - csr.add_edge("b", "L", "c").unwrap(); - let mut edges = csr.subgraph(&["a"], None, Direction::Out, 3, 3, None); - edges.sort(); - // Level 1 fills the cap, so no level-1 node expands. - let expected: Vec<(String, String, String)> = ["b", "m", "z"] - .iter() - .map(|dst| ("a".to_string(), "L".to_string(), dst.to_string())) - .collect(); - assert_eq!(edges, expected); - } - - #[test] - fn bfs_cycle() { - let mut csr = CsrIndex::new(test_memory()); - csr.add_edge("a", "L", "b").unwrap(); - csr.add_edge("b", "L", "c").unwrap(); - csr.add_edge("c", "L", "a").unwrap(); - let mut result = csr.traverse_bfs( - BfsParams { - start_nodes: &["a"], - label_filter: None, - direction: Direction::Out, - max_depth: 10, - max_visited: DEFAULT_MAX_VISITED, - frontier_bitmap: None, - }, - None, - ); - result.sort(); - assert_eq!(result, vec!["a", "b", "c"]); - } - - #[test] - fn bfs_with_depth() { - let csr = make_csr(); - let result = csr.traverse_bfs_with_depth( - &["a"], - Some("KNOWS"), - Direction::Out, - 3, - DEFAULT_MAX_VISITED, - ); - let map: HashMap = result.into_iter().collect(); - assert_eq!(map["a"], 0); - assert_eq!(map["b"], 1); - assert_eq!(map["c"], 2); - assert_eq!(map["d"], 3); - } - - fn path_params<'a>( - src: &'a str, - dst: &'a str, - label_filter: Option<&'a str>, - max_depth: usize, - frontier_bitmap: Option<&'a nodedb_types::SurrogateBitmap>, - ) -> ShortestPathParams<'a> { - ShortestPathParams { - src, - dst, - label_filter, - max_depth, - max_visited: DEFAULT_MAX_VISITED, - frontier_bitmap, - } - } - - #[test] - fn shortest_path_direct() { - let csr = make_csr(); - let path = csr - .shortest_path(path_params("a", "c", Some("KNOWS"), 5, None), None) - .unwrap(); - assert_eq!(path, vec!["a", "b", "c"]); - } - - #[test] - fn shortest_path_takes_the_smallest_named_tie() { - // Two paths of equal length, the `z` one stored first. - let mut csr = CsrIndex::new(test_memory()); - csr.add_edge("a", "L", "z").unwrap(); - csr.add_edge("z", "L", "d").unwrap(); - csr.add_edge("a", "L", "b").unwrap(); - csr.add_edge("b", "L", "d").unwrap(); - let path = csr - .shortest_path(path_params("a", "d", None, 5, None), None) - .unwrap(); - assert_eq!(path, vec!["a", "b", "d"]); - } - - #[test] - fn shortest_path_same_node() { - let csr = make_csr(); - let path = csr - .shortest_path(path_params("a", "a", None, 5, None), None) - .unwrap(); - assert_eq!(path, vec!["a"]); - } - - #[test] - fn shortest_path_unreachable() { - let csr = make_csr(); - let path = csr.shortest_path(path_params("d", "a", Some("KNOWS"), 5, None), None); - assert!(path.is_none()); - } - - #[test] - fn shortest_path_depth_limit() { - let csr = make_csr(); - let path = csr.shortest_path(path_params("a", "d", Some("KNOWS"), 1, None), None); - assert!(path.is_none()); - } - - #[test] - fn subgraph_materialization() { - let csr = make_csr(); - let edges = csr.subgraph(&["a"], None, Direction::Out, 2, DEFAULT_MAX_VISITED, None); - assert_eq!(edges.len(), 3); - assert!(edges.contains(&("a".into(), "KNOWS".into(), "b".into()))); - assert!(edges.contains(&("a".into(), "WORKS".into(), "e".into()))); - assert!(edges.contains(&("b".into(), "KNOWS".into(), "c".into()))); - } - - #[test] - fn large_graph_bfs() { - let mut csr = CsrIndex::new(test_memory()); - for i in 0..999 { - csr.add_edge(&format!("n{i}"), "NEXT", &format!("n{}", i + 1)) - .unwrap(); - } - csr.compact() - .expect("test governor ceiling covers this reservation"); - - let result = csr.traverse_bfs( - BfsParams { - start_nodes: &["n0"], - label_filter: Some("NEXT"), - direction: Direction::Out, - max_depth: 100, - max_visited: DEFAULT_MAX_VISITED, - frontier_bitmap: None, - }, - None, - ); - assert_eq!(result.len(), 101); - - let path = csr - .shortest_path(path_params("n0", "n50", Some("NEXT"), 100, None), None) - .unwrap(); - assert_eq!(path.len(), 51); - } - - /// BFS with a frontier bitmap that includes only "b". Starting from "a", - /// "b" is reachable but "c" is blocked (its surrogate is not in the bitmap). - #[test] - fn bfs_frontier_bitmap_excludes_non_members() { - use nodedb_types::{Surrogate, SurrogateBitmap}; - - let mut csr = make_csr(); - // Assign surrogates: b=10, c=20, d=30. "a" and "e" get no surrogate. - csr.set_node_surrogate("b", Surrogate::new(10)); - csr.set_node_surrogate("c", Surrogate::new(20)); - csr.set_node_surrogate("d", Surrogate::new(30)); - - // Bitmap contains only "b" (surrogate 10). - let mut bm = SurrogateBitmap::new(); - bm.insert(Surrogate::new(10)); - - let mut result = csr.traverse_bfs( - BfsParams { - start_nodes: &["a"], - label_filter: Some("KNOWS"), - direction: Direction::Out, - max_depth: 10, - max_visited: DEFAULT_MAX_VISITED, - frontier_bitmap: Some(&bm), - }, - None, - ); - result.sort(); - // "a" is the start node (not gated). "b" passes the bitmap. "c" is - // excluded (surrogate 20 not in bitmap) so traversal stops there. - assert_eq!(result, vec!["a", "b"]); - } - - /// shortest_path with a bitmap that excludes the only intermediate node. - /// "b" is the only path from "a" to "c" via KNOWS edges; if "b" is blocked - /// then no path exists. - #[test] - fn shortest_path_frontier_bitmap_blocks_intermediate() { - use nodedb_types::{Surrogate, SurrogateBitmap}; - - let mut csr = make_csr(); - csr.set_node_surrogate("b", Surrogate::new(10)); - csr.set_node_surrogate("c", Surrogate::new(20)); - - // Bitmap that does NOT contain "b". - let mut bm = SurrogateBitmap::new(); - bm.insert(Surrogate::new(20)); // only "c" is in the bitmap - - let path = csr.shortest_path(path_params("a", "c", Some("KNOWS"), 5, Some(&bm)), None); - // "b" (surrogate 10) is not in the bitmap so expansion through it is - // blocked, making the path from "a" to "c" unreachable. - assert!(path.is_none()); - } -} diff --git a/nodedb-graph/src/traversal/bfs.rs b/nodedb-graph/src/traversal/bfs.rs new file mode 100644 index 000000000..e119b0ab8 --- /dev/null +++ b/nodedb-graph/src/traversal/bfs.rs @@ -0,0 +1,452 @@ +// SPDX-License-Identifier: Apache-2.0 + +//! Breadth-first traversal over CSR adjacency. + +use std::collections::HashSet; + +#[cfg(test)] +use super::DEFAULT_MAX_VISITED; +use crate::bfs_params::BfsParams; +use crate::csr::index::LabelFilter; +use crate::csr::{CsrIndex, Direction}; +use crate::overlay_delta::GraphOverlayDelta; +use crate::traversal_overlay::traverse_bfs_overlay; + +impl CsrIndex { + /// BFS traversal. Returns all reachable node IDs within max_depth hops. + /// + /// `max_visited` caps the number of nodes visited to prevent supernode fan-out + /// explosion. Pass [`crate::traversal::DEFAULT_MAX_VISITED`] for the standard limit. + /// + /// `frontier_bitmap`: when `Some`, only nodes whose surrogate is present in the + /// bitmap are eligible as traversal targets. Start nodes are not gated — only + /// newly discovered frontier nodes are checked. + /// + /// `overlay`: when `Some` and non-empty, the traversal observes the + /// transaction's staged edge writes/deletes (read-your-own-writes), + /// including through nodes reachable only via a staged edge. When `None` + /// or empty, the durable-only dense fast path runs unchanged. + pub fn traverse_bfs( + &self, + params: BfsParams<'_>, + overlay: Option<&GraphOverlayDelta>, + ) -> Vec { + Self::traverse_bfs_on(Some(self), params, overlay) + } + + /// [`Self::traverse_bfs`] over an optional partition. `None` is a tenant + /// with no CSR partition: the walk then follows only the staged edges of + /// `overlay`, and returns no node without them. + pub fn traverse_bfs_on( + partition: Option<&Self>, + params: BfsParams<'_>, + overlay: Option<&GraphOverlayDelta>, + ) -> Vec { + match (partition, overlay) { + (_, Some(ov)) if !ov.is_empty() => traverse_bfs_overlay(partition, params, ov), + (Some(csr), _) => csr.traverse_bfs_dense(params), + (None, _) => Vec::new(), + } + } + + /// Durable-only BFS over the dense u32 CSR ids. + /// + /// The walk runs level by level. Each level's new nodes are admitted in + /// node-name order until `max_visited` nodes are visited, so a capped walk + /// admits the same nodes however the edges are stored. A cluster + /// coordinator walking the same edges across partitions admits the same + /// nodes. + fn traverse_bfs_dense(&self, params: BfsParams<'_>) -> Vec { + let BfsParams { + start_nodes, + label_filter, + direction, + max_depth, + max_visited, + frontier_bitmap, + } = params; + let labels = self.label_filter(label_filter); + let in_bitmap = |id: u32| { + frontier_bitmap.is_none_or(|bm| { + bm.contains(nodedb_types::Surrogate::new(self.node_surrogate_raw(id))) + }) + }; + let mut visited: HashSet = HashSet::new(); + let mut frontier: Vec = Vec::new(); + for &node in start_nodes { + if let Some(&id) = self.node_to_id.get(node) + && visited.insert(id) + { + frontier.push(id); + } + } + + for _depth in 0..max_depth { + if frontier.is_empty() || visited.len() >= max_visited { + break; + } + let candidates = + self.level_candidates(&frontier, &labels, direction, &visited, in_bitmap); + frontier = self.admit_by_name(candidates, &mut visited, max_visited); + } + + visited + .into_iter() + .map(|id| self.id_to_node[id as usize].clone()) + .collect() + } + + /// Collect the unvisited neighbors of `frontier` that pass `labels` and + /// `eligible`. Records an access on every frontier node. + fn level_candidates( + &self, + frontier: &[u32], + labels: &LabelFilter, + direction: Direction, + visited: &HashSet, + eligible: impl Fn(u32) -> bool, + ) -> Vec { + let mut candidates: Vec = Vec::new(); + for &node_id in frontier { + // Track access for hot/cold partition decisions. + self.record_access(node_id); + if matches!(direction, Direction::Out | Direction::Both) { + for (lid, dst) in self.dense_iter_out(node_id) { + if labels.keeps(lid) && !visited.contains(&dst) && eligible(dst) { + candidates.push(dst); + } + } + } + if matches!(direction, Direction::In | Direction::Both) { + for (lid, src) in self.dense_iter_in(node_id) { + if labels.keeps(lid) && !visited.contains(&src) && eligible(src) { + candidates.push(src); + } + } + } + } + candidates + } + + /// Admit `candidates` into `visited` in node-name order, until `visited` + /// holds `max_visited` nodes. Returns the nodes admitted, in name order. + pub(crate) fn admit_by_name( + &self, + mut candidates: Vec, + visited: &mut HashSet, + max_visited: usize, + ) -> Vec { + self.sort_by_name(&mut candidates); + candidates.dedup(); + let mut admitted = Vec::with_capacity(candidates.len()); + for id in candidates { + if visited.len() >= max_visited { + break; + } + if visited.insert(id) { + self.prefetch_node(id); + admitted.push(id); + } + } + admitted + } + + /// BFS traversal returning nodes with their hop depth. + /// + /// An empty `label_filter` keeps every edge. Otherwise an edge whose label + /// is any listed label passes. The walk runs level by level and admits + /// each level's new nodes in node-name order until `max_visited` nodes are + /// visited, like [`Self::traverse_bfs`]. The depth tag saturates at + /// `u8::MAX`. + /// + /// `max_visited` caps the number of nodes visited to prevent supernode fan-out + /// explosion. Pass [`crate::traversal::DEFAULT_MAX_VISITED`] for the standard limit. + pub fn traverse_bfs_with_depth( + &self, + start_nodes: &[&str], + label_filter: &[&str], + direction: Direction, + max_depth: usize, + max_visited: usize, + ) -> Vec<(String, u8)> { + let labels = self.label_filter(label_filter); + let mut visited: HashSet = HashSet::new(); + let mut depths: Vec<(u32, u8)> = Vec::new(); + let mut frontier: Vec = Vec::new(); + for &node in start_nodes { + if let Some(&id) = self.node_to_id.get(node) + && visited.insert(id) + { + frontier.push(id); + depths.push((id, 0)); + } + } + + for depth in 0..max_depth { + if frontier.is_empty() || visited.len() >= max_visited { + break; + } + let tag = u8::try_from(depth.saturating_add(1)).unwrap_or(u8::MAX); + let candidates = + self.level_candidates(&frontier, &labels, direction, &visited, |_| true); + frontier = self.admit_by_name(candidates, &mut visited, max_visited); + depths.extend(frontier.iter().map(|&id| (id, tag))); + } + + depths + .into_iter() + .map(|(id, depth)| (self.id_to_node[id as usize].clone(), depth)) + .collect() + } +} + +#[cfg(test)] +mod tests { + use std::collections::HashMap; + + use super::*; + use crate::test_support::{chain_csr, long_chain_csr, test_memory}; + + fn bfs(csr: &CsrIndex, labels: &[&str], max_depth: usize, max_visited: usize) -> Vec { + let mut result = csr.traverse_bfs( + BfsParams { + start_nodes: &["a"], + label_filter: labels, + direction: Direction::Out, + max_depth, + max_visited, + frontier_bitmap: None, + }, + None, + ); + result.sort(); + result + } + + #[test] + fn a_label_set_reaches_both_labels() { + let csr = chain_csr(); + assert_eq!( + bfs(&csr, &["KNOWS", "WORKS"], 1, DEFAULT_MAX_VISITED), + vec!["a", "b", "e"] + ); + } + + #[test] + fn an_unknown_label_in_a_set_adds_nothing() { + let csr = chain_csr(); + assert_eq!( + bfs(&csr, &["KNOWS", "NOPE"], 3, DEFAULT_MAX_VISITED), + bfs(&csr, &["KNOWS"], 3, DEFAULT_MAX_VISITED) + ); + } + + #[test] + fn a_set_of_unknown_labels_returns_the_start_only() { + let csr = chain_csr(); + assert_eq!(bfs(&csr, &["NOPE"], 3, DEFAULT_MAX_VISITED), vec!["a"]); + } + + /// The same edges, inserted in two orders, give the same capped walk. + #[test] + fn a_capped_walk_over_a_set_ignores_insertion_order() { + let edges = [ + ("a", "K", "z"), + ("a", "W", "m"), + ("a", "K", "b"), + ("a", "W", "c"), + ("a", "X", "a0"), + ]; + let mut forward = CsrIndex::new(test_memory()); + for (src, label, dst) in edges { + forward.add_edge(src, label, dst).unwrap(); + } + let mut reverse = CsrIndex::new(test_memory()); + for (src, label, dst) in edges.iter().rev() { + reverse.add_edge(src, label, dst).unwrap(); + } + let forward_walk = bfs(&forward, &["K", "W"], 2, 3); + assert_eq!(forward_walk, vec!["a", "b", "c"]); + assert_eq!(bfs(&reverse, &["K", "W"], 2, 3), forward_walk); + } + + /// `a` points at `z`, `m` and `b`, stored in that order. A cap of 3 leaves + /// room for two of them, and name order picks `b` and `m`. + #[test] + fn a_capped_depth_walk_admits_in_name_order() { + let mut csr = CsrIndex::new(test_memory()); + csr.add_edge("a", "L", "z").unwrap(); + csr.add_edge("a", "L", "m").unwrap(); + csr.add_edge("a", "L", "b").unwrap(); + let map: HashMap = csr + .traverse_bfs_with_depth(&["a"], &[], Direction::Out, 2, 3) + .into_iter() + .collect(); + assert_eq!(map.len(), 3); + assert_eq!(map["a"], 0); + assert_eq!(map["b"], 1); + assert_eq!(map["m"], 1); + } + + /// A walk deeper than 255 hops tags every node past 255 with `u8::MAX`. + #[test] + fn the_depth_tag_saturates_past_255() { + let mut csr = CsrIndex::new(test_memory()); + for i in 0..270 { + csr.add_edge(&format!("n{i}"), "NEXT", &format!("n{}", i + 1)) + .unwrap(); + } + let map: HashMap = csr + .traverse_bfs_with_depth(&["n0"], &["NEXT"], Direction::Out, 300, DEFAULT_MAX_VISITED) + .into_iter() + .collect(); + assert_eq!(map.len(), 271); + assert_eq!(map["n0"], 0); + assert_eq!(map["n200"], 200); + assert_eq!(map["n255"], 255); + assert_eq!(map["n256"], u8::MAX); + assert_eq!(map["n270"], u8::MAX); + } + + #[test] + fn bfs_traversal() { + let csr = chain_csr(); + let mut result = csr.traverse_bfs( + BfsParams { + start_nodes: &["a"], + label_filter: &["KNOWS"], + direction: Direction::Out, + max_depth: 2, + max_visited: DEFAULT_MAX_VISITED, + frontier_bitmap: None, + }, + None, + ); + result.sort(); + assert_eq!(result, vec!["a", "b", "c"]); + } + + #[test] + fn bfs_all_labels() { + let csr = chain_csr(); + let mut result = csr.traverse_bfs( + BfsParams { + start_nodes: &["a"], + label_filter: &[], + direction: Direction::Out, + max_depth: 1, + max_visited: DEFAULT_MAX_VISITED, + frontier_bitmap: None, + }, + None, + ); + result.sort(); + assert_eq!(result, vec!["a", "b", "e"]); + } + + /// `a` points at `z`, `m` and `b`, stored in that order. A cap of 3 leaves + /// room for two of them, and name order picks `b` and `m`. + #[test] + fn a_capped_bfs_admits_each_level_in_name_order() { + let mut csr = CsrIndex::new(test_memory()); + csr.add_edge("a", "L", "z").unwrap(); + csr.add_edge("a", "L", "m").unwrap(); + csr.add_edge("a", "L", "b").unwrap(); + let mut result = csr.traverse_bfs( + BfsParams { + start_nodes: &["a"], + label_filter: &[], + direction: Direction::Out, + max_depth: 2, + max_visited: 3, + frontier_bitmap: None, + }, + None, + ); + result.sort(); + assert_eq!(result, vec!["a", "b", "m"]); + } + + #[test] + fn bfs_cycle() { + let mut csr = CsrIndex::new(test_memory()); + csr.add_edge("a", "L", "b").unwrap(); + csr.add_edge("b", "L", "c").unwrap(); + csr.add_edge("c", "L", "a").unwrap(); + let mut result = csr.traverse_bfs( + BfsParams { + start_nodes: &["a"], + label_filter: &[], + direction: Direction::Out, + max_depth: 10, + max_visited: DEFAULT_MAX_VISITED, + frontier_bitmap: None, + }, + None, + ); + result.sort(); + assert_eq!(result, vec!["a", "b", "c"]); + } + + #[test] + fn bfs_with_depth() { + let csr = chain_csr(); + let result = + csr.traverse_bfs_with_depth(&["a"], &["KNOWS"], Direction::Out, 3, DEFAULT_MAX_VISITED); + let map: HashMap = result.into_iter().collect(); + assert_eq!(map["a"], 0); + assert_eq!(map["b"], 1); + assert_eq!(map["c"], 2); + assert_eq!(map["d"], 3); + } + + #[test] + fn large_graph_bfs() { + let csr = long_chain_csr(); + let result = csr.traverse_bfs( + BfsParams { + start_nodes: &["n0"], + label_filter: &["NEXT"], + direction: Direction::Out, + max_depth: 100, + max_visited: DEFAULT_MAX_VISITED, + frontier_bitmap: None, + }, + None, + ); + assert_eq!(result.len(), 101); + } + + /// BFS with a frontier bitmap that includes only "b". Starting from "a", + /// "b" is reachable but "c" is blocked (its surrogate is not in the bitmap). + #[test] + fn bfs_frontier_bitmap_excludes_non_members() { + use nodedb_types::{Surrogate, SurrogateBitmap}; + + let mut csr = chain_csr(); + // Assign surrogates: b=10, c=20, d=30. "a" and "e" get no surrogate. + csr.set_node_surrogate("b", Surrogate::new(10)); + csr.set_node_surrogate("c", Surrogate::new(20)); + csr.set_node_surrogate("d", Surrogate::new(30)); + + // Bitmap contains only "b" (surrogate 10). + let mut bm = SurrogateBitmap::new(); + bm.insert(Surrogate::new(10)); + + let mut result = csr.traverse_bfs( + BfsParams { + start_nodes: &["a"], + label_filter: &["KNOWS"], + direction: Direction::Out, + max_depth: 10, + max_visited: DEFAULT_MAX_VISITED, + frontier_bitmap: Some(&bm), + }, + None, + ); + result.sort(); + // "a" is the start node (not gated). "b" passes the bitmap. "c" is + // excluded (surrogate 20 not in bitmap) so traversal stops there. + assert_eq!(result, vec!["a", "b"]); + } +} diff --git a/nodedb-graph/src/traversal/mod.rs b/nodedb-graph/src/traversal/mod.rs new file mode 100644 index 000000000..856bdcb48 --- /dev/null +++ b/nodedb-graph/src/traversal/mod.rs @@ -0,0 +1,14 @@ +// SPDX-License-Identifier: Apache-2.0 + +//! Graph traversal algorithms over CSR adjacency: BFS, bidirectional +//! shortest path, and subgraph materialization. +//! +//! Every algorithm respects a max-visited cap, so supernode fan-out cannot +//! consume unbounded memory. Each traversal records node access for hot/cold +//! partition decisions and prefetches frontier neighbors. + +pub mod bfs; +pub mod shortest_path; +pub mod subgraph; + +pub use nodedb_types::config::tuning::DEFAULT_MAX_VISITED; diff --git a/nodedb-graph/src/traversal/shortest_path.rs b/nodedb-graph/src/traversal/shortest_path.rs new file mode 100644 index 000000000..34d7c4be8 --- /dev/null +++ b/nodedb-graph/src/traversal/shortest_path.rs @@ -0,0 +1,460 @@ +// SPDX-License-Identifier: Apache-2.0 + +//! Bidirectional shortest paths over CSR adjacency. + +use std::collections::{HashMap, hash_map::Entry}; + +#[cfg(test)] +use super::DEFAULT_MAX_VISITED; +use crate::csr::CsrIndex; +use crate::overlay_delta::GraphOverlayDelta; +use crate::path_overlay::shortest_path_overlay; +use crate::path_params::ShortestPathParams; + +impl CsrIndex { + /// Shortest path via bidirectional BFS. + /// + /// `max_visited` caps the combined forward+backward visited set to prevent + /// supernode fan-out explosion. Pass [`crate::traversal::DEFAULT_MAX_VISITED`] + /// for the standard limit. + /// + /// `frontier_bitmap`: when `Some`, only nodes whose surrogate is present in + /// the bitmap are eligible for expansion. The start and end nodes are not + /// gated. + /// + /// `overlay`: when `Some` and non-empty, the search observes the + /// transaction's staged edge writes/deletes (read-your-own-writes), + /// including a path that must pass through a node reachable only via a + /// staged edge. When `None` or empty, the durable-only dense bidirectional + /// fast path runs unchanged. + pub fn shortest_path( + &self, + params: ShortestPathParams<'_>, + overlay: Option<&GraphOverlayDelta>, + ) -> Option> { + Self::shortest_path_on(Some(self), params, overlay) + } + + /// [`Self::shortest_path`] over an optional partition. `None` is a + /// tenant with no CSR partition: the search then follows only the staged + /// edges of `overlay`, and finds no path without them. + pub fn shortest_path_on( + partition: Option<&Self>, + params: ShortestPathParams<'_>, + overlay: Option<&GraphOverlayDelta>, + ) -> Option> { + match (partition, overlay) { + (_, Some(ov)) if !ov.is_empty() => shortest_path_overlay(partition, params, ov), + (Some(csr), _) => csr.shortest_path_dense(params), + (None, _) => None, + } + } + + /// Durable-only bidirectional BFS over the dense u32 CSR ids. + /// + /// Each step expands one forward level, then one backward level. A level + /// relaxes its edges in `(neighbour, frontier node)` name order, and the + /// search stops at the first node both sides reached. The cap is checked + /// before each step. + /// + /// Round `k` meets on a path of `2k - 1` edges (forward level) or `2k` + /// edges (backward level), so the first meeting is a shortest path, and + /// `max_depth.div_ceil(2)` rounds reach every path of at most `max_depth` + /// edges. A meeting on a longer path returns `None`. + fn shortest_path_dense(&self, params: ShortestPathParams<'_>) -> Option> { + let ShortestPathParams { + src, + dst, + label_filter, + max_depth, + max_visited, + frontier_bitmap, + } = params; + let src_id = *self.node_to_id.get(src)?; + let dst_id = *self.node_to_id.get(dst)?; + if src_id == dst_id { + return Some(vec![src.to_string()]); + } + + let labels = self.label_filter(label_filter); + // The endpoints are never gated: a bitmap that leaves them out still + // admits the edges that reach them. + let in_bitmap = |id: u32| { + id == src_id + || id == dst_id + || frontier_bitmap.is_none_or(|bm| { + bm.contains(nodedb_types::Surrogate::new(self.node_surrogate_raw(id))) + }) + }; + let within_depth = + |path: Vec| (path.len().saturating_sub(1) <= max_depth).then_some(path); + let mut fwd_parent: HashMap = HashMap::new(); + let mut bwd_parent: HashMap = HashMap::new(); + fwd_parent.insert(src_id, src_id); + bwd_parent.insert(dst_id, dst_id); + + let mut fwd_frontier: Vec = vec![src_id]; + let mut bwd_frontier: Vec = vec![dst_id]; + + for _round in 0..max_depth.div_ceil(2) { + if fwd_parent.len() + bwd_parent.len() >= max_visited { + break; + } + + // Each level's edges are relaxed in (neighbour, frontier node) + // name order, so the parent a node gets, and the meeting point, + // do not depend on how the edges are stored. A cluster coordinator + // relaxes cross-shard hops in the same order. + let mut candidates: Vec<(u32, u32)> = Vec::new(); + for &node in &fwd_frontier { + self.record_access(node); + for (lid, neighbor) in self.dense_iter_out(node) { + if labels.keeps(lid) && in_bitmap(neighbor) { + candidates.push((neighbor, node)); + } + } + } + self.sort_edges_by_name(&mut candidates); + let mut next_fwd = Vec::new(); + for (neighbor, node) in candidates { + if let Entry::Vacant(e) = fwd_parent.entry(neighbor) { + e.insert(node); + next_fwd.push(neighbor); + } + if bwd_parent.contains_key(&neighbor) { + return within_depth(self.reconstruct_path(neighbor, &fwd_parent, &bwd_parent)); + } + } + fwd_frontier = next_fwd; + + let mut candidates: Vec<(u32, u32)> = Vec::new(); + for &node in &bwd_frontier { + self.record_access(node); + for (lid, neighbor) in self.dense_iter_in(node) { + if labels.keeps(lid) && in_bitmap(neighbor) { + candidates.push((neighbor, node)); + } + } + } + self.sort_edges_by_name(&mut candidates); + let mut next_bwd = Vec::new(); + for (neighbor, node) in candidates { + if let Entry::Vacant(e) = bwd_parent.entry(neighbor) { + e.insert(node); + next_bwd.push(neighbor); + } + if fwd_parent.contains_key(&neighbor) { + return within_depth(self.reconstruct_path(neighbor, &fwd_parent, &bwd_parent)); + } + } + bwd_frontier = next_bwd; + + if fwd_frontier.is_empty() && bwd_frontier.is_empty() { + break; + } + } + None + } + + /// Order `(neighbour, frontier node)` edges by the two node names. + fn sort_edges_by_name(&self, edges: &mut [(u32, u32)]) { + edges.sort_by(|a, b| { + let name = |id: u32| self.node_name_checked(id); + (name(a.0), name(a.1)).cmp(&(name(b.0), name(b.1))) + }); + } + + fn reconstruct_path( + &self, + meeting: u32, + fwd_parent: &HashMap, + bwd_parent: &HashMap, + ) -> Vec { + let mut fwd_path = Vec::new(); + let mut current = meeting; + loop { + fwd_path.push(current); + let parent = fwd_parent[¤t]; + if parent == current { + break; + } + current = parent; + } + fwd_path.reverse(); + + current = bwd_parent[&meeting]; + if current != meeting { + loop { + fwd_path.push(current); + let parent = bwd_parent[¤t]; + if parent == current { + break; + } + current = parent; + } + } + + fwd_path + .into_iter() + .map(|id| self.id_to_node[id as usize].clone()) + .collect() + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::test_support::{chain_csr, long_chain_csr, test_memory}; + + fn path_params<'a>( + src: &'a str, + dst: &'a str, + label_filter: &'a [&'a str], + max_depth: usize, + frontier_bitmap: Option<&'a nodedb_types::SurrogateBitmap>, + ) -> ShortestPathParams<'a> { + ShortestPathParams { + src, + dst, + label_filter, + max_depth, + max_visited: DEFAULT_MAX_VISITED, + frontier_bitmap, + } + } + + #[test] + fn shortest_path_direct() { + let csr = chain_csr(); + let path = csr + .shortest_path(path_params("a", "c", &["KNOWS"], 5, None), None) + .unwrap(); + assert_eq!(path, vec!["a", "b", "c"]); + } + + #[test] + fn shortest_path_takes_the_smallest_named_tie() { + // Two paths of equal length, the `z` one stored first. + let mut csr = CsrIndex::new(test_memory()); + csr.add_edge("a", "L", "z").unwrap(); + csr.add_edge("z", "L", "d").unwrap(); + csr.add_edge("a", "L", "b").unwrap(); + csr.add_edge("b", "L", "d").unwrap(); + let path = csr + .shortest_path(path_params("a", "d", &[], 5, None), None) + .unwrap(); + assert_eq!(path, vec!["a", "b", "d"]); + } + + #[test] + fn shortest_path_over_a_long_chain() { + let csr = long_chain_csr(); + let path = csr + .shortest_path(path_params("n0", "n50", &["NEXT"], 100, None), None) + .unwrap(); + assert_eq!(path.len(), 51); + // 50 edges: a depth of 49 finds none, 50 finds it. + assert!( + csr.shortest_path(path_params("n0", "n50", &["NEXT"], 49, None), None) + .is_none() + ); + assert!( + csr.shortest_path(path_params("n0", "n50", &["NEXT"], 50, None), None) + .is_some() + ); + } + + #[test] + fn shortest_path_same_node() { + let csr = chain_csr(); + let path = csr + .shortest_path(path_params("a", "a", &[], 5, None), None) + .unwrap(); + assert_eq!(path, vec!["a"]); + } + + #[test] + fn shortest_path_unreachable() { + let csr = chain_csr(); + let path = csr.shortest_path(path_params("d", "a", &["KNOWS"], 5, None), None); + assert!(path.is_none()); + } + + #[test] + fn shortest_path_depth_limit() { + let csr = chain_csr(); + let path = csr.shortest_path(path_params("a", "d", &["KNOWS"], 1, None), None); + assert!(path.is_none()); + } + + /// shortest_path with a bitmap that excludes the only intermediate node. + /// "b" is the only path from "a" to "c" via KNOWS edges; if "b" is blocked + /// then no path exists. + #[test] + fn shortest_path_frontier_bitmap_blocks_intermediate() { + use nodedb_types::{Surrogate, SurrogateBitmap}; + + let mut csr = chain_csr(); + csr.set_node_surrogate("b", Surrogate::new(10)); + csr.set_node_surrogate("c", Surrogate::new(20)); + + // Bitmap that does NOT contain "b". + let mut bm = SurrogateBitmap::new(); + bm.insert(Surrogate::new(20)); // only "c" is in the bitmap + + let path = csr.shortest_path(path_params("a", "c", &["KNOWS"], 5, Some(&bm)), None); + // "b" (surrogate 10) is not in the bitmap so expansion through it is + // blocked, making the path from "a" to "c" unreachable. + assert!(path.is_none()); + } + + /// A bitmap that holds neither endpoint still admits the edge between + /// them: the endpoints are never gated. + #[test] + fn the_frontier_bitmap_never_gates_the_endpoints() { + use nodedb_types::{Surrogate, SurrogateBitmap}; + + let mut csr = CsrIndex::new(test_memory()); + csr.add_edge("a", "L", "c").unwrap(); + csr.set_node_surrogate("a", Surrogate::new(1)); + csr.set_node_surrogate("c", Surrogate::new(3)); + let mut bm = SurrogateBitmap::new(); + bm.insert(Surrogate::new(99)); + + let path = csr.shortest_path(path_params("a", "c", &[], 1, Some(&bm)), None); + assert_eq!(path, Some(vec!["a".to_string(), "c".to_string()])); + } + + fn params<'a>(labels: &'a [&'a str], depth: usize) -> ShortestPathParams<'a> { + ShortestPathParams { + src: "a", + dst: "d", + label_filter: labels, + max_depth: depth, + max_visited: DEFAULT_MAX_VISITED, + frontier_bitmap: None, + } + } + + #[test] + fn dense_path_matches_any_listed_label() { + let mut csr = CsrIndex::new(test_memory()); + csr.add_edge("a", "FIRST", "b").unwrap(); + csr.add_edge("b", "SECOND", "d").unwrap(); + assert_eq!( + csr.shortest_path(params(&["FIRST", "SECOND"], 2), None), + Some(vec!["a".into(), "b".into(), "d".into()]) + ); + assert!( + csr.shortest_path(params(&["FIRST", "SECOND"], 1), None) + .is_none() + ); + assert!(csr.shortest_path(params(&["FIRST"], 2), None).is_none()); + assert!(csr.shortest_path(params(&["unknown"], 2), None).is_none()); + assert!(csr.shortest_path(params(&[], 2), None).is_some()); + assert!( + csr.shortest_path(params(&["unknown", "FIRST", "SECOND"], 2), None) + .is_some() + ); + assert!( + csr.shortest_path(params(&["FIRST", "SECOND"], 0), None) + .is_none() + ); + } + + /// `a -FIRST-> b -SECOND-> d` plus a shortcut `a -OTHER-> d`. An unknown + /// label in the set must not widen it to the shortcut. + #[test] + fn an_unknown_label_in_the_set_does_not_widen_it() { + let mut csr = CsrIndex::new(test_memory()); + csr.add_edge("a", "FIRST", "b").unwrap(); + csr.add_edge("b", "SECOND", "d").unwrap(); + csr.add_edge("a", "OTHER", "d").unwrap(); + assert_eq!( + csr.shortest_path(params(&["FIRST", "SECOND", "unknown"], 2), None), + Some(vec!["a".into(), "b".into(), "d".into()]) + ); + assert!( + csr.shortest_path(params(&["FIRST", "unknown"], 2), None) + .is_none() + ); + assert!( + csr.shortest_path(params(&["unknown", "never"], 2), None) + .is_none() + ); + } + + #[test] + fn overlay_path_matches_labels_without_durable_ids() { + let csr = CsrIndex::new(test_memory()); + let mut overlay = GraphOverlayDelta::new(); + overlay.stage_edge("a", "FIRST", "x"); + overlay.stage_edge("x", "SECOND", "d"); + assert_eq!( + csr.shortest_path(params(&["FIRST", "SECOND"], 2), Some(&overlay)), + Some(vec!["a".into(), "x".into(), "d".into()]) + ); + assert!( + csr.shortest_path(params(&["FIRST", "SECOND"], 1), Some(&overlay)) + .is_none() + ); + assert!( + csr.shortest_path(params(&["unknown"], 2), Some(&overlay)) + .is_none() + ); + assert!( + csr.shortest_path(params(&["FIRST"], 2), Some(&overlay)) + .is_none() + ); + assert!(csr.shortest_path(params(&[], 2), Some(&overlay)).is_some()); + assert!( + csr.shortest_path(params(&["FIRST", "SECOND"], 0), Some(&overlay)) + .is_none() + ); + } + + #[test] + fn mixed_labels_preserve_forward_and_backward_tombstones() { + let mut csr = CsrIndex::new(test_memory()); + csr.add_edge("a", "FIRST", "b").unwrap(); + csr.add_edge("b", "SECOND", "d").unwrap(); + let labels = ["FIRST", "SECOND"]; + for (src, label, dst) in [("a", "FIRST", "b"), ("b", "SECOND", "d")] { + let mut overlay = GraphOverlayDelta::new(); + overlay.stage_tombstone(src, label, dst); + assert!( + csr.shortest_path(params(&labels, 2), Some(&overlay)) + .is_none() + ); + } + } + + #[test] + fn odd_length_paths_count_every_edge() { + let mut csr = CsrIndex::new(test_memory()); + csr.add_edge("a", "FIRST", "b").unwrap(); + csr.add_edge("b", "SECOND", "c").unwrap(); + csr.add_edge("c", "FIRST", "d").unwrap(); + let labels = ["FIRST", "SECOND"]; + assert!(csr.shortest_path(params(&labels, 2), None).is_none()); + assert_eq!( + csr.shortest_path(params(&labels, 3), None), + Some(vec!["a".into(), "b".into(), "c".into(), "d".into()]) + ); + + let staged_csr = CsrIndex::new(test_memory()); + let mut overlay = GraphOverlayDelta::new(); + overlay.stage_edge("a", "FIRST", "b"); + overlay.stage_edge("b", "SECOND", "c"); + overlay.stage_edge("c", "FIRST", "d"); + assert!( + staged_csr + .shortest_path(params(&labels, 2), Some(&overlay)) + .is_none() + ); + assert_eq!( + staged_csr.shortest_path(params(&labels, 3), Some(&overlay)), + Some(vec!["a".into(), "b".into(), "c".into(), "d".into()]) + ); + } +} diff --git a/nodedb-graph/src/traversal/subgraph.rs b/nodedb-graph/src/traversal/subgraph.rs new file mode 100644 index 000000000..4864ce14e --- /dev/null +++ b/nodedb-graph/src/traversal/subgraph.rs @@ -0,0 +1,337 @@ +// SPDX-License-Identifier: Apache-2.0 + +//! Subgraph materialization over CSR adjacency. + +use std::collections::HashSet; + +#[cfg(test)] +use super::DEFAULT_MAX_VISITED; +use crate::csr::{CsrIndex, Direction}; +use crate::overlay_delta::GraphOverlayDelta; +use crate::traversal_overlay::{OverlaySubgraphParams, subgraph_overlay}; + +impl CsrIndex { + /// Materialize a subgraph as `(src, label, dst)` edge tuples within + /// max_depth, expanding in `direction`. + /// + /// Every expanded node records each of its edges in `direction`. The last + /// admitted level is not expanded: it records only its edges in + /// `direction` whose other endpoint is an admitted node, and admits no + /// node. + /// + /// An empty `label_filter` keeps every edge. Otherwise an edge whose label + /// is any listed label passes. + /// + /// `max_visited` caps the number of nodes visited to prevent supernode fan-out + /// explosion. Pass [`crate::traversal::DEFAULT_MAX_VISITED`] for the standard limit. + /// + /// `overlay`: when `Some` and non-empty, staged edges are included and + /// staged tombstones subtract durable edges (read-your-own-writes), + /// including through staged-only intermediate nodes. When `None` or + /// empty, the durable-only dense path runs unchanged. + pub fn subgraph( + &self, + start_nodes: &[&str], + label_filter: &[&str], + direction: Direction, + max_depth: usize, + max_visited: usize, + overlay: Option<&GraphOverlayDelta>, + ) -> Vec<(String, String, String)> { + Self::subgraph_on( + Some(self), + start_nodes, + label_filter, + direction, + max_depth, + max_visited, + overlay, + ) + } + + /// [`Self::subgraph`] over an optional partition. `None` is a tenant + /// with no CSR partition: the walk then follows only the staged edges of + /// `overlay`, and returns no edge without them. + pub fn subgraph_on( + partition: Option<&Self>, + start_nodes: &[&str], + label_filter: &[&str], + direction: Direction, + max_depth: usize, + max_visited: usize, + overlay: Option<&GraphOverlayDelta>, + ) -> Vec<(String, String, String)> { + match (partition, overlay) { + (_, Some(ov)) if !ov.is_empty() => subgraph_overlay( + partition, + OverlaySubgraphParams { + start_nodes, + label_filter, + direction, + max_depth, + max_visited, + }, + ov, + ), + (Some(csr), _) => { + csr.subgraph_dense(start_nodes, label_filter, direction, max_depth, max_visited) + } + (None, _) => Vec::new(), + } + } + + /// Durable-only subgraph materialization over the dense u32 CSR ids. + fn subgraph_dense( + &self, + start_nodes: &[&str], + label_filter: &[&str], + direction: Direction, + max_depth: usize, + max_visited: usize, + ) -> Vec<(String, String, String)> { + let labels = self.label_filter(label_filter); + let mut visited: HashSet = HashSet::new(); + let mut frontier: Vec = Vec::new(); + let mut edges = Vec::new(); + // Each physical edge once: `Both` reaches an edge from both ends, and + // one triple can be stored under several collections. + let mut seen: HashSet<(u32, u32, u32)> = HashSet::new(); + + for &node in start_nodes { + if let Some(&id) = self.node_to_id.get(node) + && visited.insert(id) + { + frontier.push(id); + } + } + + // Level by level, as `traverse_bfs_dense`: every frontier node's edges + // are recorded, then the level's new nodes are admitted in name order. + for _depth in 0..max_depth { + if frontier.is_empty() || visited.len() >= max_visited { + break; + } + let mut candidates: Vec = Vec::new(); + for &node_id in &frontier { + self.record_access(node_id); + if matches!(direction, Direction::Out | Direction::Both) { + for (lid, dst) in self.dense_iter_out(node_id) { + if labels.keeps(lid) { + if seen.insert((node_id, lid, dst)) { + edges.push(( + self.id_to_node[node_id as usize].clone(), + self.label_name(lid).to_string(), + self.id_to_node[dst as usize].clone(), + )); + } + if !visited.contains(&dst) { + candidates.push(dst); + } + } + } + } + if matches!(direction, Direction::In | Direction::Both) { + for (lid, src) in self.dense_iter_in(node_id) { + if labels.keeps(lid) { + if seen.insert((src, lid, node_id)) { + edges.push(( + self.id_to_node[src as usize].clone(), + self.label_name(lid).to_string(), + self.id_to_node[node_id as usize].clone(), + )); + } + if !visited.contains(&src) { + candidates.push(src); + } + } + } + } + } + frontier = self.admit_by_name(candidates, &mut visited, max_visited); + } + + // The last admitted level is never expanded. Its edges to admitted + // nodes, itself included, are still part of the subgraph. + for &node_id in &frontier { + if matches!(direction, Direction::Out | Direction::Both) { + for (lid, dst) in self.dense_iter_out(node_id) { + if labels.keeps(lid) + && visited.contains(&dst) + && seen.insert((node_id, lid, dst)) + { + edges.push(( + self.id_to_node[node_id as usize].clone(), + self.label_name(lid).to_string(), + self.id_to_node[dst as usize].clone(), + )); + } + } + } + if matches!(direction, Direction::In | Direction::Both) { + for (lid, src) in self.dense_iter_in(node_id) { + if labels.keeps(lid) + && visited.contains(&src) + && seen.insert((src, lid, node_id)) + { + edges.push(( + self.id_to_node[src as usize].clone(), + self.label_name(lid).to_string(), + self.id_to_node[node_id as usize].clone(), + )); + } + } + } + } + + edges + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::test_support::{chain_csr, test_memory}; + + #[test] + fn subgraph_materialization() { + let csr = chain_csr(); + let edges = csr.subgraph(&["a"], &[], Direction::Out, 2, DEFAULT_MAX_VISITED, None); + assert_eq!(edges.len(), 3); + assert!(edges.contains(&("a".into(), "KNOWS".into(), "b".into()))); + assert!(edges.contains(&("a".into(), "WORKS".into(), "e".into()))); + assert!(edges.contains(&("b".into(), "KNOWS".into(), "c".into()))); + } + #[test] + fn a_capped_subgraph_expands_only_admitted_levels() { + // `a` points at `z`, `m` and `b`, stored in that order. + let mut csr = CsrIndex::new(test_memory()); + csr.add_edge("a", "L", "z").unwrap(); + csr.add_edge("a", "L", "m").unwrap(); + csr.add_edge("a", "L", "b").unwrap(); + csr.add_edge("z", "L", "y").unwrap(); + csr.add_edge("b", "L", "c").unwrap(); + let mut edges = csr.subgraph(&["a"], &[], Direction::Out, 3, 3, None); + edges.sort(); + // Level 1 fills the cap, so no level-1 node expands. + let expected: Vec<(String, String, String)> = ["b", "m", "z"] + .iter() + .map(|dst| ("a".to_string(), "L".to_string(), dst.to_string())) + .collect(); + assert_eq!(edges, expected); + } + + #[test] + fn both_directions_return_each_physical_edge_once() { + let mut csr = CsrIndex::new(test_memory()); + csr.add_edge("a", "L", "b").unwrap(); + csr.add_edge("b", "L", "a").unwrap(); + csr.add_edge("a", "SELF", "a").unwrap(); + csr.add_edge_in_collection("a", "L", "c", "first").unwrap(); + csr.add_edge_in_collection("a", "L", "c", "second").unwrap(); + let mut edges = csr.subgraph(&["a"], &[], Direction::Both, 3, DEFAULT_MAX_VISITED, None); + edges.sort(); + let expected: Vec<(String, String, String)> = [ + ("a", "L", "b"), + ("a", "L", "c"), + ("a", "SELF", "a"), + ("b", "L", "a"), + ] + .iter() + .map(|(s, l, d)| (s.to_string(), l.to_string(), d.to_string())) + .collect(); + assert_eq!(edges, expected); + } + + fn triples(edges: &[(&str, &str, &str)]) -> Vec<(String, String, String)> { + let mut out: Vec<(String, String, String)> = edges + .iter() + .map(|(s, l, d)| (s.to_string(), l.to_string(), d.to_string())) + .collect(); + out.sort(); + out + } + + fn boundary_csr() -> CsrIndex { + let mut csr = CsrIndex::new(test_memory()); + for (src, dst) in [ + ("a", "b"), + ("a", "c"), + ("b", "c"), + ("c", "b"), + ("c", "a"), + ("c", "x"), + ] { + csr.add_edge(src, "L", dst).unwrap(); + } + csr + } + + #[test] + fn boundary_nodes_record_edges_among_admitted_nodes() { + let csr = boundary_csr(); + // A non-empty overlay runs the string-keyed path. + let mut staged = GraphOverlayDelta::new(); + staged.stage_tombstone("q", "L", "r"); + for overlay in [None, Some(&staged)] { + let mut edges = + csr.subgraph(&["a"], &[], Direction::Out, 1, DEFAULT_MAX_VISITED, overlay); + edges.sort(); + assert_eq!( + edges, + triples(&[ + ("a", "L", "b"), + ("a", "L", "c"), + ("b", "L", "c"), + ("c", "L", "a"), + ("c", "L", "b"), + ]), + "c -> x leaves the admitted set" + ); + } + } + + #[test] + fn boundary_edges_follow_the_walk_direction() { + let mut csr = CsrIndex::new(test_memory()); + for (src, dst) in [("b", "a"), ("c", "a"), ("b", "c"), ("a", "z")] { + csr.add_edge(src, "L", dst).unwrap(); + } + let mut edges = csr.subgraph(&["a"], &[], Direction::In, 1, DEFAULT_MAX_VISITED, None); + edges.sort(); + assert_eq!( + edges, + triples(&[("b", "L", "a"), ("b", "L", "c"), ("c", "L", "a")]) + ); + } + + #[test] + fn depth_zero_records_only_the_start_nodes_edges_among_themselves() { + let mut csr = CsrIndex::new(test_memory()); + csr.add_edge("a", "L", "a").unwrap(); + csr.add_edge("a", "L", "b").unwrap(); + let edges = csr.subgraph(&["a"], &[], Direction::Out, 0, DEFAULT_MAX_VISITED, None); + assert_eq!(edges, triples(&[("a", "L", "a")])); + } + + #[test] + fn a_label_set_returns_parallel_edges_of_both_labels() { + let mut csr = CsrIndex::new(test_memory()); + csr.add_edge("a", "K", "b").unwrap(); + csr.add_edge("a", "W", "b").unwrap(); + csr.add_edge("a", "X", "c").unwrap(); + let mut edges = csr.subgraph( + &["a"], + &["K", "W"], + Direction::Out, + 2, + DEFAULT_MAX_VISITED, + None, + ); + edges.sort(); + let expected: Vec<(String, String, String)> = [("a", "K", "b"), ("a", "W", "b")] + .iter() + .map(|(s, l, d)| (s.to_string(), l.to_string(), d.to_string())) + .collect(); + assert_eq!(edges, expected); + } +} diff --git a/nodedb-graph/src/traversal_overlay.rs b/nodedb-graph/src/traversal_overlay.rs index 640f3fdea..7190a4f23 100644 --- a/nodedb-graph/src/traversal_overlay.rs +++ b/nodedb-graph/src/traversal_overlay.rs @@ -12,194 +12,283 @@ //! //! Both walks run level by level, as the durable paths do: each level's new //! nodes are admitted in node-name order under `max_visited`. +//! +//! `partition` is `None` for a tenant with no CSR partition. The walks then +//! follow only staged edges, and allocate no CSR. +//! +//! A start node joins the walk only when it exists: the partition holds it, +//! or a staged edge names it. The dense walks drop an absent start the same +//! way. use std::collections::HashSet; use crate::bfs_params::BfsParams; +use crate::csr::index::LabelFilter; use crate::csr::{CsrIndex, Direction}; use crate::overlay_delta::GraphOverlayDelta; -impl CsrIndex { - /// String-keyed BFS that merges the transaction's staged edges/tombstones. - pub(crate) fn traverse_bfs_overlay( - &self, - params: BfsParams<'_>, - overlay: &GraphOverlayDelta, - ) -> Vec { - let BfsParams { - start_nodes, - label_filter, - direction, - max_depth, - max_visited, - frontier_bitmap, - } = params; - let labels = self.label_filter(label_filter); - let in_bitmap = |id: u32| { - frontier_bitmap.is_none_or(|bm| { - bm.contains(nodedb_types::Surrogate::new(self.node_surrogate_raw(id))) - }) - }; - let mut visited: HashSet = HashSet::new(); - let mut frontier: Vec = Vec::new(); - for &node in start_nodes { - if visited.insert(node.to_string()) { - frontier.push(node.to_string()); - } +/// True when `node` exists for an overlay walk: `partition` holds it, or a +/// staged edge names it. +pub(crate) fn node_present( + partition: Option<&CsrIndex>, + overlay: &GraphOverlayDelta, + node: &str, +) -> bool { + partition.is_some_and(|csr| csr.node_to_id.contains_key(node)) || overlay.names_node(node) +} + +/// Admit each present start node once into `visited`. Returns the admitted +/// nodes in input order. +fn admit_starts( + partition: Option<&CsrIndex>, + overlay: &GraphOverlayDelta, + start_nodes: &[&str], + visited: &mut HashSet, +) -> Vec { + let mut frontier = Vec::new(); + for &node in start_nodes { + if node_present(partition, overlay, node) && visited.insert(node.to_string()) { + frontier.push(node.to_string()); } + } + frontier +} + +/// String-keyed BFS that merges the transaction's staged edges/tombstones. +pub(crate) fn traverse_bfs_overlay( + partition: Option<&CsrIndex>, + params: BfsParams<'_>, + overlay: &GraphOverlayDelta, +) -> Vec { + let BfsParams { + start_nodes, + label_filter, + direction, + max_depth, + max_visited, + frontier_bitmap, + } = params; + let durable_labels = partition.map(|csr| csr.label_filter(label_filter)); + let mut visited: HashSet = HashSet::new(); + let mut frontier = admit_starts(partition, overlay, start_nodes, &mut visited); - let want_out = matches!(direction, Direction::Out | Direction::Both); - let want_in = matches!(direction, Direction::In | Direction::Both); + let want_out = matches!(direction, Direction::Out | Direction::Both); + let want_in = matches!(direction, Direction::In | Direction::Both); - for _depth in 0..max_depth { - if frontier.is_empty() || visited.len() >= max_visited { - break; - } - let mut candidates: Vec = Vec::new(); - for node in &frontier { - // Durable CSR expansion for nodes that carry a surrogate. - if let Some(&node_id) = self.node_to_id.get(node.as_str()) { - self.record_access(node_id); - if want_out { - for (lid, dst) in self.dense_iter_out(node_id) { - let dst_name = &self.id_to_node[dst as usize]; - if labels.keeps(lid) - && !overlay.is_tombstoned(node, self.label_name(lid), dst_name) - && in_bitmap(dst) - && !visited.contains(dst_name) - { - candidates.push(dst_name.clone()); - } + for _depth in 0..max_depth { + if frontier.is_empty() || visited.len() >= max_visited { + break; + } + let mut candidates: Vec = Vec::new(); + for node in &frontier { + // Durable CSR expansion for nodes that carry a surrogate. + if let (Some(csr), Some(labels)) = (partition, durable_labels.as_ref()) + && let Some(&node_id) = csr.node_to_id.get(node.as_str()) + { + let in_bitmap = |id: u32| { + frontier_bitmap.is_none_or(|bm| { + bm.contains(nodedb_types::Surrogate::new(csr.node_surrogate_raw(id))) + }) + }; + csr.record_access(node_id); + if want_out { + for (lid, dst) in csr.dense_iter_out(node_id) { + let dst_name = &csr.id_to_node[dst as usize]; + if labels.keeps(lid) + && !overlay.is_tombstoned(node, csr.label_name(lid), dst_name) + && in_bitmap(dst) + && !visited.contains(dst_name) + { + candidates.push(dst_name.clone()); } } - if want_in { - for (lid, src) in self.dense_iter_in(node_id) { - let src_name = &self.id_to_node[src as usize]; - if labels.keeps(lid) - && !overlay.is_tombstoned(src_name, self.label_name(lid), node) - && in_bitmap(src) - && !visited.contains(src_name) - { - candidates.push(src_name.clone()); - } + } + if want_in { + for (lid, src) in csr.dense_iter_in(node_id) { + let src_name = &csr.id_to_node[src as usize]; + if labels.keeps(lid) + && !overlay.is_tombstoned(src_name, csr.label_name(lid), node) + && in_bitmap(src) + && !visited.contains(src_name) + { + candidates.push(src_name.clone()); } } } + } - // Staged edges — followed for durable and staged-only nodes - // alike. Staged edges are the transaction's own writes, so - // bitmap gating (which needs a durable surrogate) does not - // apply. - if want_out { - candidates.extend( - overlay - .out_neighbors(node, label_filter) - .map(|(_, dst)| dst.to_string()) - .filter(|dst| !visited.contains(dst)), - ); - } - if want_in { - candidates.extend( - overlay - .in_neighbors(node, label_filter) - .map(|(_, src)| src.to_string()) - .filter(|src| !visited.contains(src)), - ); - } + // Staged edges — followed for durable and staged-only nodes + // alike. Staged edges are the transaction's own writes, so + // bitmap gating (which needs a durable surrogate) does not + // apply. + if want_out { + candidates.extend( + overlay + .out_neighbors(node, label_filter) + .map(|(_, dst)| dst.to_string()) + .filter(|dst| !visited.contains(dst)), + ); + } + if want_in { + candidates.extend( + overlay + .in_neighbors(node, label_filter) + .map(|(_, src)| src.to_string()) + .filter(|src| !visited.contains(src)), + ); } - frontier = admit_names(candidates, &mut visited, max_visited); } + frontier = admit_names(candidates, &mut visited, max_visited); + } - visited.into_iter().collect() + visited.into_iter().collect() +} + +/// The arguments of one overlay subgraph walk. +pub(crate) struct OverlaySubgraphParams<'a> { + pub start_nodes: &'a [&'a str], + pub label_filter: &'a [&'a str], + pub direction: Direction, + pub max_depth: usize, + pub max_visited: usize, +} + +/// String-keyed subgraph materialization merging staged edges/tombstones. +pub(crate) fn subgraph_overlay( + partition: Option<&CsrIndex>, + params: OverlaySubgraphParams<'_>, + overlay: &GraphOverlayDelta, +) -> Vec<(String, String, String)> { + let OverlaySubgraphParams { + start_nodes, + label_filter, + direction, + max_depth, + max_visited, + } = params; + let mut visited: HashSet = HashSet::new(); + let mut edges: Vec<(String, String, String)> = Vec::new(); + // Each physical edge once: `Both` reaches an edge from both ends, and + // one triple can be stored under several collections. + let mut seen: HashSet<(String, String, String)> = HashSet::new(); + let mut frontier = admit_starts(partition, overlay, start_nodes, &mut visited); + + let scope = OverlayEdgeScope { + partition: partition.map(|csr| (csr, csr.label_filter(label_filter))), + label_filter, + want_out: matches!(direction, Direction::Out | Direction::Both), + want_in: matches!(direction, Direction::In | Direction::Both), + overlay, + }; + + for _depth in 0..max_depth { + if frontier.is_empty() || visited.len() >= max_visited { + break; + } + let mut candidates: Vec = Vec::new(); + for node in &frontier { + for (edge, neighbor) in overlay_node_edges(node, &scope) { + push_once(&mut edges, &mut seen, edge); + if !visited.contains(&neighbor) { + candidates.push(neighbor); + } + } + } + frontier = admit_names(candidates, &mut visited, max_visited); } - /// String-keyed subgraph materialization merging staged edges/tombstones. - pub(crate) fn subgraph_overlay( - &self, - start_nodes: &[&str], - label_filter: Option<&str>, - direction: Direction, - max_depth: usize, - max_visited: usize, - overlay: &GraphOverlayDelta, - ) -> Vec<(String, String, String)> { - let labels = self.label_filter(label_filter); - let mut visited: HashSet = HashSet::new(); - let mut frontier: Vec = Vec::new(); - let mut edges: Vec<(String, String, String)> = Vec::new(); - - for &node in start_nodes { - if visited.insert(node.to_string()) { - frontier.push(node.to_string()); + // The last admitted level is never expanded. Its edges to admitted + // nodes, itself included, are still part of the subgraph. + for node in &frontier { + for (edge, neighbor) in overlay_node_edges(node, &scope) { + if visited.contains(&neighbor) { + push_once(&mut edges, &mut seen, edge); } } + } - let want_out = matches!(direction, Direction::Out | Direction::Both); - let want_in = matches!(direction, Direction::In | Direction::Both); + edges +} - for _depth in 0..max_depth { - if frontier.is_empty() || visited.len() >= max_visited { - break; - } - let mut candidates: Vec = Vec::new(); - for node in &frontier { - if let Some(&node_id) = self.node_to_id.get(node.as_str()) { - self.record_access(node_id); - if want_out { - for (lid, dst) in self.dense_iter_out(node_id) { - if !labels.keeps(lid) { - continue; - } - let label = self.label_name(lid); - let dst_name = &self.id_to_node[dst as usize]; - if overlay.is_tombstoned(node, label, dst_name) { - continue; - } - edges.push((node.clone(), label.to_string(), dst_name.clone())); - if !visited.contains(dst_name) { - candidates.push(dst_name.clone()); - } - } - } - if want_in { - for (lid, src) in self.dense_iter_in(node_id) { - if !labels.keeps(lid) { - continue; - } - let label = self.label_name(lid); - let src_name = &self.id_to_node[src as usize]; - if overlay.is_tombstoned(src_name, label, node) { - continue; - } - edges.push((src_name.clone(), label.to_string(), node.clone())); - if !visited.contains(src_name) { - candidates.push(src_name.clone()); - } - } - } +/// Each edge of `node` in the walk's directions, durable then staged, as +/// `(physical edge, neighbour)`. A tombstoned durable edge is skipped. +fn overlay_node_edges( + node: &str, + scope: &OverlayEdgeScope<'_>, +) -> Vec<((String, String, String), String)> { + let mut out = Vec::new(); + if let Some((csr, labels)) = scope.partition.as_ref() + && let Some(&node_id) = csr.node_to_id.get(node) + { + csr.record_access(node_id); + if scope.want_out { + for (lid, dst) in csr.dense_iter_out(node_id) { + if !labels.keeps(lid) { + continue; } - - if want_out { - for (label, dst) in overlay.out_neighbors(node, label_filter) { - edges.push((node.clone(), label.to_string(), dst.to_string())); - if !visited.contains(dst) { - candidates.push(dst.to_string()); - } - } + let label = csr.label_name(lid); + let dst_name = &csr.id_to_node[dst as usize]; + if !scope.overlay.is_tombstoned(node, label, dst_name) { + out.push(( + (node.to_string(), label.to_string(), dst_name.clone()), + dst_name.clone(), + )); } - if want_in { - for (label, src) in overlay.in_neighbors(node, label_filter) { - edges.push((src.to_string(), label.to_string(), node.clone())); - if !visited.contains(src) { - candidates.push(src.to_string()); - } - } + } + } + if scope.want_in { + for (lid, src) in csr.dense_iter_in(node_id) { + if !labels.keeps(lid) { + continue; + } + let label = csr.label_name(lid); + let src_name = &csr.id_to_node[src as usize]; + if !scope.overlay.is_tombstoned(src_name, label, node) { + out.push(( + (src_name.clone(), label.to_string(), node.to_string()), + src_name.clone(), + )); } } - frontier = admit_names(candidates, &mut visited, max_visited); } + } + if scope.want_out { + for (label, dst) in scope.overlay.out_neighbors(node, scope.label_filter) { + out.push(( + (node.to_string(), label.to_string(), dst.to_string()), + dst.to_string(), + )); + } + } + if scope.want_in { + for (label, src) in scope.overlay.in_neighbors(node, scope.label_filter) { + out.push(( + (src.to_string(), label.to_string(), node.to_string()), + src.to_string(), + )); + } + } + out +} + +/// What one overlay subgraph walk keeps of each node's edges. +struct OverlayEdgeScope<'a> { + /// The durable partition and the walk's labels resolved against it. + partition: Option<(&'a CsrIndex, LabelFilter)>, + label_filter: &'a [&'a str], + want_out: bool, + want_in: bool, + overlay: &'a GraphOverlayDelta, +} - edges +/// Append `edge` unless `seen` already holds it. +fn push_once( + edges: &mut Vec<(String, String, String)>, + seen: &mut HashSet<(String, String, String)>, + edge: (String, String, String), +) { + if seen.insert(edge.clone()) { + edges.push(edge); } } @@ -250,7 +339,7 @@ mod tests { let mut r = csr.traverse_bfs( BfsParams { start_nodes: &["a"], - label_filter: Some("KNOWS"), + label_filter: &["KNOWS"], direction: Direction::Out, max_depth: 2, max_visited: DEFAULT_MAX_VISITED, @@ -273,7 +362,7 @@ mod tests { let mut r = csr.traverse_bfs( BfsParams { start_nodes: &["a"], - label_filter: Some("KNOWS"), + label_filter: &["KNOWS"], direction: Direction::Out, max_depth: 2, max_visited: 3, @@ -294,7 +383,7 @@ mod tests { let mut r = csr.traverse_bfs( BfsParams { start_nodes: &["a"], - label_filter: Some("KNOWS"), + label_filter: &["KNOWS"], direction: Direction::Out, max_depth: 2, max_visited: DEFAULT_MAX_VISITED, @@ -318,7 +407,7 @@ mod tests { let edges = csr.subgraph( &["a"], - Some("KNOWS"), + &["KNOWS"], Direction::Out, 1, DEFAULT_MAX_VISITED, @@ -329,6 +418,33 @@ mod tests { assert!(!edges.contains(&("a".into(), "KNOWS".into(), "b".into()))); } + #[test] + fn subgraph_both_returns_each_physical_edge_once() { + // Durable a->b. Staged b->a and a self-loop on a. + let csr = base(); + let mut ov = GraphOverlayDelta::new(); + ov.stage_edge("b", "KNOWS", "a"); + ov.stage_edge("a", "KNOWS", "a"); + let mut edges = csr.subgraph( + &["a"], + &[], + Direction::Both, + 3, + DEFAULT_MAX_VISITED, + Some(&ov), + ); + edges.sort(); + let expected: Vec<(String, String, String)> = [ + ("a", "KNOWS", "a"), + ("a", "KNOWS", "b"), + ("b", "KNOWS", "a"), + ] + .iter() + .map(|(s, l, d)| (s.to_string(), l.to_string(), d.to_string())) + .collect(); + assert_eq!(edges, expected); + } + #[test] fn subgraph_in_direction_surfaces_staged_in_edge() { // Staged in-edge z->a; querying subgraph In from "a" surfaces it. @@ -338,7 +454,7 @@ mod tests { let edges = csr.subgraph( &["a"], - Some("KNOWS"), + &["KNOWS"], Direction::In, 1, DEFAULT_MAX_VISITED, @@ -347,6 +463,126 @@ mod tests { assert!(edges.contains(&("z".into(), "KNOWS".into(), "a".into()))); } + /// Durable `a -KNOWS-> b`. Staged `a -LIKES-> x` and `a -HATES-> y`. The + /// set `["KNOWS", "LIKES"]` follows the staged edge under its second label. + #[test] + fn a_staged_edge_under_the_second_label_of_a_set_is_followed() { + let csr = base(); + let mut ov = GraphOverlayDelta::new(); + ov.stage_edge("a", "LIKES", "x"); + ov.stage_edge("a", "HATES", "y"); + + let mut r = csr.traverse_bfs( + BfsParams { + start_nodes: &["a"], + label_filter: &["KNOWS", "LIKES"], + direction: Direction::Out, + max_depth: 1, + max_visited: DEFAULT_MAX_VISITED, + frontier_bitmap: None, + }, + Some(&ov), + ); + r.sort(); + assert_eq!(r, vec!["a", "b", "x"]); + + let mut edges = csr.subgraph( + &["a"], + &["KNOWS", "LIKES"], + Direction::Out, + 1, + DEFAULT_MAX_VISITED, + Some(&ov), + ); + edges.sort(); + let expected: Vec<(String, String, String)> = [("a", "KNOWS", "b"), ("a", "LIKES", "x")] + .iter() + .map(|(s, l, d)| (s.to_string(), l.to_string(), d.to_string())) + .collect(); + assert_eq!(edges, expected); + } + + fn out_params<'a>(starts: &'a [&'a str], depth: usize) -> BfsParams<'a> { + BfsParams { + start_nodes: starts, + label_filter: &[], + direction: Direction::Out, + max_depth: depth, + max_visited: DEFAULT_MAX_VISITED, + frontier_bitmap: None, + } + } + + /// A start the partition lacks and no staged edge names is dropped, as + /// in the dense walk. A start a staged edge names is kept. + #[test] + fn an_absent_start_is_dropped() { + let csr = base(); + let mut ov = GraphOverlayDelta::new(); + ov.stage_edge("x", "KNOWS", "y"); + + let mut r = csr.traverse_bfs(out_params(&["a", "ghost", "x"], 1), Some(&ov)); + r.sort(); + assert_eq!(r, vec!["a", "b", "x", "y"]); + let dense = csr.traverse_bfs(out_params(&["ghost"], 1), None); + assert!(dense.is_empty()); + + let edges = csr.subgraph( + &["ghost"], + &[], + Direction::Out, + 2, + DEFAULT_MAX_VISITED, + Some(&ov), + ); + assert!(edges.is_empty()); + } + + /// With no partition, the walks follow staged edges only, and keep only + /// starts a staged edge names. + #[test] + fn a_missing_partition_walks_staged_edges() { + let mut ov = GraphOverlayDelta::new(); + ov.stage_edge("a", "KNOWS", "x"); + ov.stage_edge("x", "KNOWS", "y"); + + let mut r = CsrIndex::traverse_bfs_on(None, out_params(&["a", "ghost"], 1), Some(&ov)); + r.sort(); + assert_eq!(r, vec!["a", "x"]); + let mut r = CsrIndex::traverse_bfs_on(None, out_params(&["a"], 2), Some(&ov)); + r.sort(); + assert_eq!(r, vec!["a", "x", "y"]); + assert!(CsrIndex::traverse_bfs_on(None, out_params(&["a"], 2), None).is_empty()); + + let mut edges = CsrIndex::subgraph_on( + None, + &["a"], + &[], + Direction::Out, + 2, + DEFAULT_MAX_VISITED, + Some(&ov), + ); + edges.sort(); + let expected: Vec<(String, String, String)> = [("a", "KNOWS", "x"), ("x", "KNOWS", "y")] + .iter() + .map(|(s, l, d)| (s.to_string(), l.to_string(), d.to_string())) + .collect(); + assert_eq!(edges, expected); + assert!( + CsrIndex::subgraph_on( + None, + &["a"], + &[], + Direction::Out, + 2, + DEFAULT_MAX_VISITED, + None + ) + .is_empty() + ); + } + #[test] fn empty_overlay_matches_durable() { let csr = base(); @@ -354,7 +590,7 @@ mod tests { let mut with = csr.traverse_bfs( BfsParams { start_nodes: &["a"], - label_filter: None, + label_filter: &[], direction: Direction::Out, max_depth: 2, max_visited: DEFAULT_MAX_VISITED, @@ -365,7 +601,7 @@ mod tests { let mut without = csr.traverse_bfs( BfsParams { start_nodes: &["a"], - label_filter: None, + label_filter: &[], direction: Direction::Out, max_depth: 2, max_visited: DEFAULT_MAX_VISITED, diff --git a/nodedb-graph/src/traversal_surrogate.rs b/nodedb-graph/src/traversal_surrogate.rs index 82c835039..07aea3391 100644 --- a/nodedb-graph/src/traversal_surrogate.rs +++ b/nodedb-graph/src/traversal_surrogate.rs @@ -30,8 +30,9 @@ pub struct SurrogateBfsParams<'a> { /// name-seeded walk working on an index whose surrogate bindings have not /// been populated. pub seeds: &'a [u32], - /// Restrict expansion to one edge label. - pub label_filter: Option<&'a str>, + /// Edge labels to expand over. Empty keeps every edge. Otherwise an edge + /// whose label is any listed label passes. + pub label_filter: &'a [&'a str], pub direction: Direction, /// Hops from the seeds. `0` returns just the addressable seeds. pub max_depth: usize, @@ -127,14 +128,14 @@ impl CsrIndex { let Some(collection_id) = self.collection_id(collection) else { return hops; }; - // An unknown label matches no edge. Distinguishing that from "no filter" - // matters: falling through to unfiltered expansion would return another - // label's neighbourhood under the caller's label. - if label_filter.is_some_and(|l| self.label_id(l).is_none()) { + // A set of unknown labels matches no edge. Falling through to + // unfiltered expansion would return other labels' neighbourhoods under + // the caller's labels. + let labels = self.label_filter(label_filter); + if labels.keeps_none() { self.seed_only(seeds, &mut hops); return hops; } - let label_id = label_filter.and_then(|l| self.label_id(l)); let mut visited: HashSet = HashSet::with_capacity(max_visited.min(1024)); let mut frontier: Vec = Vec::new(); @@ -170,10 +171,7 @@ impl CsrIndex { neighbors.extend(self.iter_in_edges_raw_in(node, collection_id)); } for (lid, other) in neighbors { - if label_id.is_some_and(|f| f != lid) - || visited.contains(&other) - || !offered.insert(other) - { + if !labels.keeps(lid) || visited.contains(&other) || !offered.insert(other) { continue; } candidates.push(other); @@ -202,8 +200,8 @@ impl CsrIndex { nodes.sort_by(|a, b| self.node_name_checked(*a).cmp(&self.node_name_checked(*b))); } - /// Record the addressable seeds and nothing else. Used when the requested - /// edge label does not exist in this partition, so no expansion is possible + /// Record the addressable seeds and nothing else. Used when no requested + /// edge label exists in this partition, so no expansion is possible /// but the seeds themselves are still legitimately reachable at depth 0. fn seed_only(&self, seeds: &[u32], hops: &mut SurrogateHops) { let mut seen: HashSet = HashSet::new(); @@ -250,7 +248,7 @@ mod tests { fn params<'a>(seeds: &'a [u32], collection: &'a str) -> SurrogateBfsParams<'a> { SurrogateBfsParams { seeds, - label_filter: None, + label_filter: &[], direction: Direction::Out, max_depth: 5, max_visited: 1000, @@ -369,12 +367,56 @@ mod tests { let csr = seeded_csr(); let seeds = [local(&csr, "a")]; let mut p = params(&seeds, "people"); - p.label_filter = Some("never_inserted"); + p.label_filter = &["never_inserted"]; let hops = csr.traverse_surrogates_in_collection(p); assert!(hops.reached.contains(Surrogate::new(10))); assert!(!hops.reached.contains(Surrogate::new(20))); } + /// `a -knows-> b -likes-> c -hates-> d` in `people`. + fn three_label_csr() -> CsrIndex { + let mut csr = CsrIndex::new(test_memory()); + for (src, label, dst) in [ + ("a", "knows", "b"), + ("b", "likes", "c"), + ("c", "hates", "d"), + ] { + csr.add_edge_in_collection(src, label, dst, "people") + .unwrap_or_else(|e| panic!("seed edge {src}->{dst}: {e}")); + } + csr + } + + fn reached_names<'a>(csr: &'a CsrIndex, hops: &SurrogateHops) -> Vec<&'a str> { + let mut names: Vec<&str> = hops + .distances + .iter() + .filter_map(|&(l, _)| csr.node_name_checked(l)) + .collect(); + names.sort_unstable(); + names + } + + #[test] + fn a_label_set_expands_over_every_listed_label() { + let csr = three_label_csr(); + let seeds = [local(&csr, "a")]; + let mut p = params(&seeds, "people"); + p.label_filter = &["knows", "likes"]; + let hops = csr.traverse_surrogates_in_collection(p); + assert_eq!(reached_names(&csr, &hops), vec!["a", "b", "c"]); + } + + #[test] + fn an_unknown_label_in_a_set_adds_nothing() { + let csr = three_label_csr(); + let seeds = [local(&csr, "a")]; + let mut p = params(&seeds, "people"); + p.label_filter = &["knows", "never_inserted"]; + let hops = csr.traverse_surrogates_in_collection(p); + assert_eq!(reached_names(&csr, &hops), vec!["a", "b"]); + } + #[test] fn max_depth_zero_returns_only_the_seeds() { let csr = seeded_csr(); diff --git a/nodedb-physical/src/physical_plan/columnar.rs b/nodedb-physical/src/physical_plan/columnar.rs index 2cce68d98..db7ef7e81 100644 --- a/nodedb-physical/src/physical_plan/columnar.rs +++ b/nodedb-physical/src/physical_plan/columnar.rs @@ -173,13 +173,13 @@ pub enum ColumnarOp { /// Update rows matching filter predicates. /// /// Uses `MutationEngine` for plain/spatial profiles. - /// `updates` is a list of (field_name, json_value_bytes) pairs. Update { collection: QualifiedCollection, /// Serialized `Vec` (MessagePack). filters: Vec, - /// Field assignments: `(column_name, json_value_bytes)`. - updates: Vec<(String, Vec)>, + /// Field assignments: `(column_name, value)`. An `Expr` value + /// evaluates against each matched row's pre-image. + updates: Vec<(String, super::document::UpdateValue)>, /// Compiled row-level-security WRITE predicate, evaluated against each /// row's post-image once the assignments have been applied, or the /// reason no predicate is attached. @@ -240,7 +240,7 @@ pub enum ColumnarOp { /// Serialized `Vec` — the statement's WHERE clause. filters: Vec, /// Field assignments for an UPDATE. Empty for a DELETE. - updates: Vec<(String, Vec)>, + updates: Vec<(String, super::document::UpdateValue)>, /// True for UPDATE, false for DELETE. is_update: bool, rls_write_check: RlsWriteCheck, diff --git a/nodedb-physical/src/physical_plan/document/declared_column.rs b/nodedb-physical/src/physical_plan/document/declared_column.rs new file mode 100644 index 000000000..35708355f --- /dev/null +++ b/nodedb-physical/src/physical_plan/document/declared_column.rs @@ -0,0 +1,130 @@ +// SPDX-License-Identifier: Apache-2.0 + +//! A declared numeric column of a schemaless document or KV collection. +//! +//! These engines store a row as a MessagePack map. The Data Plane re-types +//! every value written to one of these columns by the rule the strict encoder +//! applies, declared integer or float width included. The Control Plane +//! derives the list from the catalog and ships it on `DocumentOp::Register`: +//! from the declared text for a schemaless collection, from the typed schema +//! for a KV collection. + +use nodedb_types::columnar::{ColumnDef, ColumnType, FloatWidth, IntWidth}; + +/// One declared column whose stored value the write path re-types. +#[derive( + Debug, + Clone, + PartialEq, + serde::Serialize, + serde::Deserialize, + zerompk::ToMessagePack, + zerompk::FromMessagePack, +)] +pub struct DeclaredColumn { + /// The field name, as the catalog records it. + pub name: String, + /// `Int64`, `Float64`, or `Decimal(Some(typmod))`. + pub column_type: ColumnType, + /// The declared integer width. `SMALLINT` and `INTEGER` bound the value. + pub int_width: Option, + /// The declared float width. `REAL` refuses a finite value past `f32`. + pub float_width: Option, +} + +impl DeclaredColumn { + /// The column `name` declared as `declared`, the DDL type text with its + /// modifiers. + /// + /// `Some` only for a type that fixes the stored numeric value: an + /// integer, a float, or a `DECIMAL(p,s)`. Every other declared type is + /// stored as given and returns `None`. + pub fn from_declared(name: &str, declared: &str) -> Option { + let column_type = ColumnType::from_declared_type(declared)?; + if !fixes_stored_value(&column_type) { + return None; + } + Some(Self { + name: name.to_string(), + column_type, + int_width: IntWidth::from_declared_type(declared), + float_width: FloatWidth::from_declared_type(declared), + }) + } + + /// The typed schema column `column`, with its declared width. + /// + /// `Some` only for a column type that fixes the stored numeric value, as + /// [`Self::from_declared`] decides. + pub fn from_column_def(column: &ColumnDef) -> Option { + fixes_stored_value(&column.column_type).then(|| Self { + name: column.name.clone(), + column_type: column.column_type, + int_width: column.int_width, + float_width: column.float_width, + }) + } +} + +/// Whether a column of `column_type` fixes the numeric value it stores: an +/// integer, a float, or a `DECIMAL(p,s)`. +fn fixes_stored_value(column_type: &ColumnType) -> bool { + matches!( + column_type, + ColumnType::Int64 | ColumnType::Float64 | ColumnType::Decimal(Some(_)) + ) +} + +#[cfg(test)] +mod tests { + use nodedb_types::columnar::DecimalTypmod; + + use super::*; + + #[test] + fn numeric_declarations_resolve_with_their_width() { + let small = DeclaredColumn::from_declared("n", "SMALLINT NOT NULL").expect("smallint"); + assert_eq!(small.column_type, ColumnType::Int64); + assert_eq!(small.int_width, Some(IntWidth::I16)); + + let real = DeclaredColumn::from_declared("r", "REAL").expect("real"); + assert_eq!(real.column_type, ColumnType::Float64); + assert_eq!(real.float_width, Some(FloatWidth::F32)); + + let typmod = DecimalTypmod::new(5, 2).expect("valid typmod"); + let dec = DeclaredColumn::from_declared("d", "DECIMAL(5, 2)").expect("decimal"); + assert_eq!(dec.column_type, ColumnType::Decimal(Some(typmod))); + } + + #[test] + fn schema_column_keeps_its_declared_width() { + let small = ColumnDef::nullable("n", ColumnType::Int64).with_declared_width("SMALLINT"); + let declared = DeclaredColumn::from_column_def(&small).expect("integer column"); + assert_eq!(declared.int_width, Some(IntWidth::I16)); + assert_eq!( + declared, + DeclaredColumn::from_declared("n", "SMALLINT").expect("smallint") + ); + assert!( + DeclaredColumn::from_column_def(&ColumnDef::nullable("t", ColumnType::String)) + .is_none() + ); + } + + #[test] + fn other_declarations_are_not_retyped() { + for declared in [ + "TEXT", + "DECIMAL", + "BOOL", + "TIMESTAMP", + "VECTOR(3)", + "unknown", + ] { + assert!( + DeclaredColumn::from_declared("c", declared).is_none(), + "{declared}" + ); + } + } +} diff --git a/nodedb-physical/src/physical_plan/document/mod.rs b/nodedb-physical/src/physical_plan/document/mod.rs index 9fb7cc651..39ee94ba1 100644 --- a/nodedb-physical/src/physical_plan/document/mod.rs +++ b/nodedb-physical/src/physical_plan/document/mod.rs @@ -2,6 +2,7 @@ //! Document / sparse engine operations dispatched to the Data Plane. +pub mod declared_column; pub mod enforcement_types; pub mod merge_types; pub mod ollp_edge; @@ -12,6 +13,7 @@ pub mod timeseries_schema; pub mod types; pub mod update_value; +pub use declared_column::DeclaredColumn; pub use enforcement_types::{ RetentionDuration, RetentionUnit, StateTransitionDef, TransitionCheckDef, TransitionRule, }; diff --git a/nodedb-physical/src/physical_plan/document/op.rs b/nodedb-physical/src/physical_plan/document/op.rs index 0a127cd39..6af3676cf 100644 --- a/nodedb-physical/src/physical_plan/document/op.rs +++ b/nodedb-physical/src/physical_plan/document/op.rs @@ -268,6 +268,21 @@ pub enum DocumentOp { /// `WITH (primary='vector')`. Both plain and vector-primary rows are /// legal MessagePack maps — this is the only way to tell them apart. vector_primary: Option>, + /// Declared `VECTOR(n)` columns of a schemaless collection, as + /// `(column, n)`. A document write indexes each one into its + /// field's vector index, as a strict schema's vector columns are. + /// Empty for every other collection. + vector_fields: Vec<(String, usize)>, + /// Declared numeric columns of a schemaless or KV collection. Every + /// value a write stores under one of them is re-typed to it. Empty + /// for every other collection. + declared_columns: Vec, + /// The collection's declared key column, per + /// `nodedb_types::declared_key`. `None` when rows key by the implicit + /// `id` or `_rowid`. A row read from the sparse store renders its + /// identity under this column, else under `id`, and only when the + /// row lacks that column. + declared_key: Option, }, /// Lookup documents by secondary index value. diff --git a/nodedb-physical/src/physical_plan/graph/op.rs b/nodedb-physical/src/physical_plan/graph/op.rs index 955abf8eb..5501358c4 100644 --- a/nodedb-physical/src/physical_plan/graph/op.rs +++ b/nodedb-physical/src/physical_plan/graph/op.rs @@ -82,7 +82,8 @@ pub enum GraphOp { /// mapping back to a collection, authorized via the index DDL instead. collection: Option, start_nodes: Vec, - edge_label: Option, + /// Empty keeps every edge. Otherwise an edge with any listed label. + edge_labels: Vec, direction: Direction, depth: usize, options: GraphTraversalOptions, @@ -98,7 +99,8 @@ pub enum GraphOp { /// See `Hop::collection`. collection: Option, node_id: String, - edge_label: Option, + /// Empty keeps every edge. Otherwise an edge with any listed label. + edge_labels: Vec, direction: Direction, /// RLS filters applied to neighbor nodes before returning. rls_filters: Vec, @@ -117,11 +119,20 @@ pub enum GraphOp { /// See `Hop::collection`. collection: Option, node_ids: Vec, - edge_label: Option, + /// Empty keeps every edge. Otherwise an edge with any listed label. + edge_labels: Vec, direction: Direction, max_results: u32, /// RLS filters applied to neighbor nodes before returning. rls_filters: Vec, + /// Edge-property predicate, AND-ed. Evaluated against each crossed + /// edge's current property object in `collection` before the row + /// counts against `max_results`. Empty admits every edge. Non-empty + /// requires `collection`. + edge_predicate: Vec, + /// Each row carries the crossed edge's current property object. + /// Requires `collection`. + with_properties: bool, }, /// Shortest path between two nodes. @@ -130,7 +141,8 @@ pub enum GraphOp { collection: Option, src: String, dst: String, - edge_label: Option, + /// Empty keeps every edge. Otherwise an edge with any listed label. + edge_labels: Vec, max_depth: usize, options: GraphTraversalOptions, /// RLS filters applied to path nodes before returning. @@ -145,7 +157,8 @@ pub enum GraphOp { /// See `Hop::collection`. collection: Option, start_nodes: Vec, - edge_label: Option, + /// Empty keeps every edge. Otherwise an edge with any listed label. + edge_labels: Vec, depth: usize, options: GraphTraversalOptions, /// RLS filters applied to subgraph nodes/edges before returning. diff --git a/nodedb-physical/src/physical_plan/kv/mod.rs b/nodedb-physical/src/physical_plan/kv/mod.rs index bbbdc278a..3898a4298 100644 --- a/nodedb-physical/src/physical_plan/kv/mod.rs +++ b/nodedb-physical/src/physical_plan/kv/mod.rs @@ -7,8 +7,10 @@ pub mod counter_shape; pub mod op; pub mod resolved_mutation; pub mod sorted_read; +pub mod transfer_amount; pub use counter_shape::KvCounterShape; pub use op::KvOp; pub use resolved_mutation::{KvResolveOutcome, KvResolvedMutation}; pub use sorted_read::{SortedIndexRead, SortedIndexSpec}; +pub use transfer_amount::TransferAmount; diff --git a/nodedb-physical/src/physical_plan/kv/op.rs b/nodedb-physical/src/physical_plan/kv/op.rs index 9e1366330..70490f4c8 100644 --- a/nodedb-physical/src/physical_plan/kv/op.rs +++ b/nodedb-physical/src/physical_plan/kv/op.rs @@ -7,6 +7,7 @@ use nodedb_types::{QualifiedCollection, RlsWriteCheck, Surrogate}; use super::counter_shape::KvCounterShape; use super::resolved_mutation::KvResolvedMutation; use super::sorted_read::{SortedIndexRead, SortedIndexSpec}; +use super::transfer_amount::TransferAmount; use crate::physical_plan::document::ReturningSpec; /// KV engine physical operations. @@ -402,8 +403,9 @@ pub enum KvOp { source_key: Vec, dest_key: Vec, field: String, - /// Amount to transfer (encoded as f64 bytes). - amount: f64, + /// Amount to transfer, typed by the field it moves: exact for a + /// `DECIMAL` field, a float otherwise. + amount: TransferAmount, /// Debit (source) row's identity, content-addressed on `(collection, /// source_key)`, threaded to the write-back. debit_surrogate: Surrogate, diff --git a/nodedb-physical/src/physical_plan/kv/transfer_amount.rs b/nodedb-physical/src/physical_plan/kv/transfer_amount.rs new file mode 100644 index 000000000..658d1a104 --- /dev/null +++ b/nodedb-physical/src/physical_plan/kv/transfer_amount.rs @@ -0,0 +1,141 @@ +// SPDX-License-Identifier: Apache-2.0 + +//! The amount a fungible `Transfer` moves. +//! +//! The planner types the amount against the field it moves, the way an +//! assignment types its value. A `DECIMAL` field takes an exact amount, so +//! the balance moves by exactly the literal. Every other field takes a +//! float amount, moved by float arithmetic. + +use std::fmt; + +use rust_decimal::Decimal; + +/// The amount a fungible `Transfer` moves, typed by the field it moves. +#[derive(Debug, Clone, Copy, PartialEq, serde::Serialize, serde::Deserialize)] +pub enum TransferAmount { + /// The amount of a field that is not `DECIMAL`. + Float(f64), + /// The exact amount of a `DECIMAL` field. + Decimal(Decimal), +} + +/// Wire tag of [`TransferAmount::Float`]. +const TAG_FLOAT: u8 = 0; +/// Wire tag of [`TransferAmount::Decimal`]. +const TAG_DECIMAL: u8 = 1; + +impl TransferAmount { + /// Whether the amount is above zero. A transfer moves a positive amount. + pub fn is_positive(self) -> bool { + match self { + Self::Float(f) => f > 0.0, + Self::Decimal(d) => d > Decimal::ZERO, + } + } + + /// The amount as an `f64`. A decimal converts through its text, so the + /// result is the `f64` nearest the exact value. + pub fn to_f64(self) -> f64 { + match self { + Self::Float(f) => f, + Self::Decimal(d) => d.to_string().parse().unwrap_or(f64::NAN), + } + } + + /// The amount as an exact decimal. A float converts through its + /// shortest round-trip text, so the literal `0.1` stays `0.1`. `None` + /// when the float lies outside the `Decimal` range. + pub fn to_decimal(self) -> Option { + match self { + Self::Float(f) => f.to_string().parse().ok(), + Self::Decimal(d) => Some(d), + } + } +} + +impl fmt::Display for TransferAmount { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + match self { + Self::Float(v) => write!(f, "{v}"), + Self::Decimal(d) => write!(f, "{d}"), + } + } +} + +/// Encoded as `[tag, payload]`. A decimal payload is its 16-byte +/// `Decimal::serialize` form, the form `nodedb_types::Value` uses. +impl zerompk::ToMessagePack for TransferAmount { + fn write(&self, writer: &mut W) -> zerompk::Result<()> { + writer.write_array_len(2)?; + match self { + Self::Float(f) => { + writer.write_u8(TAG_FLOAT)?; + writer.write_f64(*f) + } + Self::Decimal(d) => { + writer.write_u8(TAG_DECIMAL)?; + writer.write_binary(&d.serialize()) + } + } + } +} + +impl<'de> zerompk::FromMessagePack<'de> for TransferAmount { + fn read>(reader: &mut R) -> zerompk::Result { + reader.check_array_len(2)?; + match reader.read_u8()? { + TAG_FLOAT => Ok(Self::Float(reader.read_f64()?)), + TAG_DECIMAL => { + let bytes = reader.read_binary()?; + let buf = + <[u8; 16]>::try_from(&bytes[..]).map_err(|_| zerompk::Error::BufferTooSmall)?; + Ok(Self::Decimal(Decimal::deserialize(buf))) + } + tag => Err(zerompk::Error::InvalidMarker(tag)), + } + } +} + +#[cfg(test)] +mod tests { + use super::*; + + fn dec(text: &str) -> Decimal { + text.parse().expect("test decimal parses") + } + + #[test] + fn both_variants_roundtrip_through_msgpack() { + for amount in [ + TransferAmount::Float(30.5), + TransferAmount::Decimal(dec("12.345")), + TransferAmount::Decimal(dec("-0.01")), + ] { + let bytes = zerompk::to_msgpack_vec(&amount).expect("encode"); + let back: TransferAmount = zerompk::from_msgpack(&bytes).expect("decode"); + assert_eq!(back, amount); + } + } + + #[test] + fn an_unknown_tag_is_refused() { + let bytes = zerompk::to_msgpack_vec(&(7u8, 1.0f64)).expect("encode"); + assert!(zerompk::from_msgpack::(&bytes).is_err()); + } + + #[test] + fn a_float_literal_converts_to_its_exact_text() { + assert_eq!(TransferAmount::Float(0.1).to_decimal(), Some(dec("0.1"))); + assert_eq!(TransferAmount::Float(1e300).to_decimal(), None); + assert_eq!(TransferAmount::Decimal(dec("0.1")).to_f64(), 0.1); + } + + #[test] + fn only_a_positive_amount_is_positive() { + assert!(TransferAmount::Decimal(dec("0.01")).is_positive()); + assert!(!TransferAmount::Decimal(Decimal::ZERO).is_positive()); + assert!(!TransferAmount::Float(-1.0).is_positive()); + assert!(!TransferAmount::Float(f64::NAN).is_positive()); + } +} diff --git a/nodedb-physical/src/physical_plan/meta.rs b/nodedb-physical/src/physical_plan/meta.rs index 727acc5e7..52b842968 100644 --- a/nodedb-physical/src/physical_plan/meta.rs +++ b/nodedb-physical/src/physical_plan/meta.rs @@ -10,13 +10,6 @@ use nodedb_types::{QualifiedCollection, TenantId, Value}; pub use super::meta_calvin::PassiveReadKeyId; -/// Byte length of the [`MetaOp::MarkSavepoint`] response payload / -/// [`MetaOp::RollbackToSavepoint`] marker set: three little-endian `u64`s -/// (value/TTL overlay, GRAPH overlay, ARRAY overlay journal lengths). -/// Encoder and decoder both key off this constant so the layout can only -/// change in one place. -pub const SAVEPOINT_MARKER_BYTES: usize = 24; - /// Meta / maintenance physical operations. #[derive( Debug, @@ -487,29 +480,29 @@ pub enum MetaOp { /// Mark a savepoint in the per-transaction staging overlays. /// - /// A single savepoint spans the value/TTL overlay and the parallel GRAPH - /// and ARRAY overlays, which keep independent undo journals. The Data - /// Plane returns a 24-byte composite marker — three little-endian `u64`s: - /// the value overlay's journal length, then the GRAPH overlay's, then the - /// ARRAY overlay's — so the Control Plane can record all three as the - /// savepoint's rollback markers. - /// In-memory only — savepoints append no WAL. Keyed by the request's - /// `txn_id`. - MarkSavepoint { txn_id: nodedb_types::id::TxnId }, + /// A core holds one value/TTL, one GRAPH and one ARRAY overlay per + /// transaction, shared by every vShard it hosts. The core records the + /// three undo-journal lengths under `savepoint`. The Control Plane sends + /// the mark through each vShard the transaction has staged to, so a core + /// can receive it more than once. The core keeps its first record: no + /// write stages between those sends. `savepoint` is unique for the + /// transaction and grows with each mark. In-memory only, no WAL. + MarkSavepoint { + txn_id: nodedb_types::id::TxnId, + savepoint: u64, + }, /// Roll the per-transaction staging overlays back to a savepoint. /// - /// Replays every overlay's undo journal from its end down to its marker - /// in reverse — restoring each recorded prior slot (or removing it when - /// absent) in the value/TTL overlay to `value_marker`, in the GRAPH - /// overlay to `graph_marker`, and in the ARRAY overlay to `array_marker` - /// — then truncates each journal to its marker. The transaction stays - /// open. In-memory only. Keyed by `txn_id`. + /// Replays each overlay's undo journal in reverse down to the length the + /// core recorded under `savepoint`, restoring each prior slot or removing + /// it when absent, then truncates the journal there. A core with no + /// record had no staged vShard at the mark, so its overlays rewind to + /// empty. The core drops its records of later savepoints. The + /// transaction stays open. In-memory only. RollbackToSavepoint { txn_id: nodedb_types::id::TxnId, - value_marker: u64, - graph_marker: u64, - array_marker: u64, + savepoint: u64, }, /// Record the per-key write versions of a committed Calvin transaction's diff --git a/nodedb-physical/src/physical_plan/mod.rs b/nodedb-physical/src/physical_plan/mod.rs index 787c86e99..5e1919506 100644 --- a/nodedb-physical/src/physical_plan/mod.rs +++ b/nodedb-physical/src/physical_plan/mod.rs @@ -42,11 +42,11 @@ pub use cluster_event::{ClusterEventOp, MAX_REMOTE_CDC_COMMITTED_OFFSETS}; pub use columnar::{ColumnarInsertIntent, ColumnarOp}; pub use crdt::{CrdtOp, CrdtWriteVerb}; pub use document::{ - BalancedDef, DocumentOp, DocumentResolveOutcome, DocumentResolvedMutation, EnforcementOptions, - GeneratedColumnSpec, MaterializedSumBinding, OllpPredictedEdge, PeriodLockConfig, - RedoSumTargets, RegisteredIndex, RegisteredIndexState, ResolvedSumTarget, ReturningColumns, - ReturningItem, ReturningSpec, StorageMode, SumTargetKey, TimeseriesSchema, UpdateValue, - resolved_sum_surrogate, + BalancedDef, DeclaredColumn, DocumentOp, DocumentResolveOutcome, DocumentResolvedMutation, + EnforcementOptions, GeneratedColumnSpec, MaterializedSumBinding, OllpPredictedEdge, + PeriodLockConfig, RedoSumTargets, RegisteredIndex, RegisteredIndexState, ResolvedSumTarget, + ReturningColumns, ReturningItem, ReturningSpec, StorageMode, SumTargetKey, TimeseriesSchema, + UpdateValue, resolved_sum_surrogate, }; pub use exchange::{ExchangeMode, ExchangeOp}; pub use graph::{ @@ -55,8 +55,9 @@ pub use graph::{ }; pub use kv::{ KvCounterShape, KvOp, KvResolveOutcome, KvResolvedMutation, SortedIndexRead, SortedIndexSpec, + TransferAmount, }; -pub use meta::{MetaOp, SAVEPOINT_MARKER_BYTES}; +pub use meta::MetaOp; pub use meta_home::{HomeAnswer, HomeVersion, HomeVersionProbe}; pub use meta_restore::{RestoredEdgeVersion, RestoredIdentity, RestoredRedo, RestoredRow}; pub use meta_snapshot::{CutCaptureRequest, SnapshotClearTarget}; @@ -67,7 +68,7 @@ pub use routing::plan_contains_cluster_partitioned_leaf; pub use set_op::SetOpKind; pub use sort_key::SortKeySpec; pub use spatial::{SpatialOp, SpatialPredicate}; -pub use text::TextOp; +pub use text::{ScoreScanBound, ScoreScanOrder, TextOp, TextScoreSpec}; pub use timeseries::{TimeseriesOp, TimeseriesResolve, UNBOUNDED_TIME_RANGE}; pub use vector::{ VectorDirectWriteIntent, VectorOp, VectorResolveOutcome, VectorResolvedMutation, diff --git a/nodedb-physical/src/physical_plan/text.rs b/nodedb-physical/src/physical_plan/text.rs index 9a8dfc487..ccfb7fe08 100644 --- a/nodedb-physical/src/physical_plan/text.rs +++ b/nodedb-physical/src/physical_plan/text.rs @@ -2,8 +2,69 @@ //! Full-text search operations dispatched to the Data Plane. +use nodedb_types::text_search::QueryMode; use nodedb_types::{QualifiedCollection, SurrogateBitmap}; +/// One per-row BM25 score column: `bm25_score(field, query)` under `alias`. +#[derive( + Debug, + Clone, + PartialEq, + serde::Serialize, + serde::Deserialize, + zerompk::ToMessagePack, + zerompk::FromMessagePack, +)] +pub struct TextScoreSpec { + /// `None` reads the whole-document index. + pub field: Option, + pub query: String, + /// Boolean combination of the query terms. + pub mode: QueryMode, + /// Fuzzy (Levenshtein) fallback for a term with no exact posting. + pub fuzzy: bool, + /// Output column the score lands in. A row the scoped index holds but + /// the query does not match carries `0.0` there. A row the index does + /// not hold carries `null`. + pub alias: String, +} + +/// The order a bounded score scan keeps its best rows in: one score column. +#[derive( + Debug, + Clone, + PartialEq, + serde::Serialize, + serde::Deserialize, + zerompk::ToMessagePack, + zerompk::FromMessagePack, +)] +pub struct ScoreScanOrder { + /// Alias of the score column the rows are ordered by. + pub alias: String, + pub ascending: bool, + /// Whether `null` scores sort before every number. + pub nulls_first: bool, +} + +/// The row bound of a score scan. +#[derive( + Debug, + Clone, + PartialEq, + serde::Serialize, + serde::Deserialize, + zerompk::ToMessagePack, + zerompk::FromMessagePack, +)] +pub struct ScoreScanBound { + /// Rows the scan returns at most. + pub rows: usize, + /// `Some`: the scan returns the first `rows` rows in this order. `None`: + /// any `rows` admitted rows. + pub order: Option, +} + /// Full-text search physical operations. #[derive( Debug, @@ -18,35 +79,49 @@ pub enum TextOp { /// BM25 full-text search on the inverted index. Search { collection: QualifiedCollection, + /// Field index the query reads. `None` reads the whole-document index. + field: Option, query: String, + /// Hits returned, best first. `usize::MAX` returns every match. top_k: usize, + /// Boolean combination of the query terms. + mode: QueryMode, /// Enable fuzzy matching (Levenshtein) for typo tolerance. fuzzy: bool, /// Pre-computed bitmap of eligible surrogates (from prefilter evaluation). /// `None` = no prefilter; all postings are eligible. prefilter: Option, - /// RLS post-score filters (serialized `Vec`). - /// Applied after BM25 scoring, before returning to client. - /// Result count may be less than requested `top_k`. + /// Residual WHERE predicates (serialized `Vec`). They + /// restrict candidates before ranking, so `top_k` counts only rows + /// that satisfy them. + filters: Vec, + /// RLS filters (serialized `Vec`). Like `filters`, they + /// restrict candidates before ranking, so `top_k` counts only rows + /// the policy admits. rls_filters: Vec, + /// Score columns injected into each hit. + scores: Vec, }, - /// Full-collection scan with per-row BM25 score injection. + /// Every row `filters` admit, each with its score columns. /// - /// Scans every document in the collection, runs FTS scoring for each - /// document against `query`, and returns all documents with a score - /// column appended under `score_alias`. Documents that do not match the - /// query receive a `null` score. This is the physical plan used when - /// `bm25_score(field, term)` appears as a SELECT projection without a - /// restricting WHERE clause — all rows must be present in the result set - /// so the query planner cannot emit the hit-only `TextOp::Search` shape. + /// The physical plan for `bm25_score(field, term)` with no `text_match`: + /// every admitted row is present. A row the score's index holds but its + /// query does not match carries `0.0`. A row the index does not hold + /// carries `null`. BM25ScoreScan { collection: QualifiedCollection, - query: String, - /// Column name under which the BM25 score is injected into each row. - score_alias: String, - /// Enable fuzzy matching for the scoring pass. - fuzzy: bool, + /// Residual WHERE predicates (serialized `Vec`). + filters: Vec, + /// RLS filters (serialized `Vec`). A row that fails one + /// is dropped. + rls_filters: Vec, + /// Score columns injected into each row. + scores: Vec, + /// The query's LIMIT pushed into the scan. `None` returns every + /// admitted row. The relational tail still applies its own ORDER BY + /// and LIMIT over the rows returned. + bound: Option, }, /// Exact phrase search: all terms must appear consecutively in the document. @@ -56,20 +131,42 @@ pub enum TextOp { /// is positional: documents with the phrase closer to the start rank higher. PhraseSearch { collection: QualifiedCollection, - /// Ordered sequence of terms to match as a phrase. + /// Field index the phrase reads. `None` reads the whole-document index. + field: Option, + /// The phrase words as written, in order. The Data Plane analyzes + /// them once with the collection's analyzer and matches the tokens + /// as a contiguous sequence. terms: Vec, + /// Hits returned, best first. `usize::MAX` returns every match. top_k: usize, /// Pre-computed bitmap of eligible surrogates (from prefilter evaluation). prefilter: Option, + /// Residual WHERE predicates (serialized `Vec`), applied + /// before ranking. + filters: Vec, + /// RLS filters (serialized `Vec`), applied before + /// ranking. + rls_filters: Vec, + /// Score columns injected into each hit. + scores: Vec, }, /// Hybrid search: vector similarity + BM25 text, fused via RRF. HybridSearch { collection: QualifiedCollection, + /// Vector column the vector leg searches. + vector_field: String, query_vector: Vec, + /// Field index the text leg reads. `None` reads the whole-document index. + text_field: Option, query_text: String, + /// Residual WHERE predicates (serialized `Vec`). They + /// restrict both legs before fusion. + filters: Vec, top_k: usize, ef_search: usize, + /// Boolean combination of the text leg's query terms. + mode: QueryMode, fuzzy: bool, /// Weight for vector results in RRF (0.0–1.0). Default: 0.5. vector_weight: f32, @@ -92,8 +189,9 @@ pub enum TextOp { collection: QualifiedCollection, /// Pre-assigned global surrogate for `(collection, doc_id)`. surrogate: nodedb_types::Surrogate, - /// Concatenated text to index. - text: String, + /// `(field, text)` per top-level string field. Empty removes the + /// document from every index. + fields: Vec<(String, String)>, /// Sync provenance: identifies the originating peer and sequence for idempotency. #[serde(default)] provenance: Option, @@ -121,8 +219,15 @@ pub enum TextOp { /// to `reciprocal_rank_fusion_weighted` with per-source k-constants. HybridSearchTriple { collection: QualifiedCollection, + /// Vector column the vector leg searches. + vector_field: String, query_vector: Vec, + /// Field index the text leg reads. `None` reads the whole-document index. + text_field: Option, query_text: String, + /// Residual WHERE predicates (serialized `Vec`). They + /// restrict every leg before fusion. + filters: Vec, /// Node id used as the BFS seed for the graph leg. graph_seed_id: String, /// Maximum BFS depth from the seed node. @@ -131,6 +236,8 @@ pub enum TextOp { graph_edge_label: Option, top_k: usize, ef_search: usize, + /// Boolean combination of the text leg's query terms. + mode: QueryMode, fuzzy: bool, /// Per-source RRF k constants: (vector_k, text_k, graph_k). rrf_k: (f64, f64, f64), diff --git a/nodedb-query/src/expr/eval.rs b/nodedb-query/src/expr/eval.rs index e73651660..fec2d9141 100644 --- a/nodedb-query/src/expr/eval.rs +++ b/nodedb-query/src/expr/eval.rs @@ -12,7 +12,7 @@ use nodedb_types::Value; use crate::value_ops::{coerced_eq, is_truthy, to_value_number, value_to_f64}; use super::binary::eval_binary_op; -use super::types::SqlExpr; +use super::types::{SqlExpr, WHOLE_ROW_COLUMN}; /// Error type for row-scope `SqlExpr` evaluation. /// @@ -25,7 +25,8 @@ use super::types::SqlExpr; /// silent `NULL`; /// - a function argument it cannot compute on: vectors of different /// dimensions, an argument of the wrong type, a malformed JSONPath. -/// SQLSTATE `22000`. A `NULL` argument stays `NULL` instead. +/// SQLSTATE `22000`. A `NULL` argument stays `NULL` instead; +/// - an exact integer SUM past the `Decimal` range, SQLSTATE `22000`. #[derive(Debug, Clone, PartialEq, Eq, thiserror::Error)] pub enum EvalError { #[error("division by zero")] @@ -62,6 +63,9 @@ pub enum EvalError { path: String, reason: String, }, + /// An exact integer aggregate total lies outside the `Decimal` range. + #[error("{function}(): integer total out of range")] + NumericOverflow { function: &'static str }, } /// Row scope for `SqlExpr::eval_scope`: how `Column(..)` and `OldColumn(..)` @@ -81,24 +85,33 @@ struct RowScope<'a> { impl<'a> RowScope<'a> { fn column(&self, name: &str) -> Value { - self.new_doc.get(name).cloned().unwrap_or(Value::Null) + row_field(self.new_doc, name) } fn old_column(&self, name: &str) -> Value { match self.old_doc { - Some(old) => old.get(name).cloned().unwrap_or(Value::Null), + Some(old) => row_field(old, name), None => Value::Null, } } fn excluded_column(&self, name: &str) -> Value { match self.excluded_doc { - Some(excluded) => excluded.get(name).cloned().unwrap_or(Value::Null), + Some(excluded) => row_field(excluded, name), None => Value::Null, } } } +/// The value `name` references in `doc`: the whole document for +/// [`WHOLE_ROW_COLUMN`], else the field, else `Null`. +fn row_field(doc: &Value, name: &str) -> Value { + if name == WHOLE_ROW_COLUMN { + return doc.clone(); + } + doc.get(name).cloned().unwrap_or(Value::Null) +} + impl SqlExpr { /// Evaluate this expression against a document. /// @@ -275,6 +288,24 @@ mod tests { assert_eq!(expr.eval(&doc()).unwrap(), Value::Null); } + #[test] + fn whole_row_column_is_the_document() { + let expr = SqlExpr::Column(WHOLE_ROW_COLUMN.into()); + assert_eq!(expr.eval(&doc()).unwrap(), doc()); + let old = SqlExpr::OldColumn(WHOLE_ROW_COLUMN.into()); + let before = Value::Object(Default::default()); + assert_eq!(old.eval_with_old(&doc(), &before).unwrap(), before); + } + + #[test] + fn to_jsonb_of_the_whole_row_keeps_every_field_type() { + let expr = SqlExpr::Function { + name: "to_jsonb".into(), + args: vec![SqlExpr::Column(WHOLE_ROW_COLUMN.into())], + }; + assert_eq!(expr.eval(&doc()).unwrap(), doc()); + } + #[test] fn literal() { let expr = SqlExpr::Literal(Value::Integer(42)); diff --git a/nodedb-query/src/expr/mod.rs b/nodedb-query/src/expr/mod.rs index b2c25fe79..5b12f7d8f 100644 --- a/nodedb-query/src/expr/mod.rs +++ b/nodedb-query/src/expr/mod.rs @@ -13,4 +13,4 @@ pub mod eval; pub mod types; pub use eval::EvalError; -pub use types::{BinaryOp, CastType, ComputedColumn, GroupKeySpec, SqlExpr}; +pub use types::{BinaryOp, CastType, ComputedColumn, GroupKeySpec, SqlExpr, WHOLE_ROW_COLUMN}; diff --git a/nodedb-query/src/expr/types.rs b/nodedb-query/src/expr/types.rs index 331a220f8..c8e3319ff 100644 --- a/nodedb-query/src/expr/types.rs +++ b/nodedb-query/src/expr/types.rs @@ -4,10 +4,17 @@ use nodedb_types::Value; +/// The column name that references the whole row. +/// +/// SQL `*` as a function argument (`to_jsonb(*)`) lowers to +/// `Column(WHOLE_ROW_COLUMN)`, which evaluates to the row document itself. +pub const WHOLE_ROW_COLUMN: &str = "*"; + /// A serializable SQL expression that can be evaluated against a document. #[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)] pub enum SqlExpr { /// Column reference: extract field value from the document. + /// [`WHOLE_ROW_COLUMN`] is the whole document. Column(String), /// Literal value. Literal(Value), diff --git a/nodedb-query/src/functions/json/dispatch.rs b/nodedb-query/src/functions/json/dispatch.rs index 2e2dd3fc1..1868f6523 100644 --- a/nodedb-query/src/functions/json/dispatch.rs +++ b/nodedb-query/src/functions/json/dispatch.rs @@ -26,6 +26,13 @@ pub(in crate::functions) fn try_eval( } fn try_eval_value(name: &str, args: &[Value]) -> Option { + // `to_jsonb(v)`: `v` as a JSON value. A `Value` already holds the JSON + // data model, so the value passes through with every type intact. + // `to_jsonb(*)` is the whole row as one JSON object. + if name == "to_jsonb" { + return Some(args.first().cloned().unwrap_or(Value::Null)); + } + // PostgreSQL JSON operator functions (lowered from AST BinaryOp). let pg_result = match name { "pg_json_get" => { diff --git a/nodedb-query/src/fusion.rs b/nodedb-query/src/fusion.rs index 02f29d8fa..d77744284 100644 --- a/nodedb-query/src/fusion.rs +++ b/nodedb-query/src/fusion.rs @@ -61,9 +61,7 @@ fn finish_fusion( .collect(); fused.sort_unstable_by(|a, b| { - b.rrf_score - .partial_cmp(&a.rrf_score) - .unwrap_or(std::cmp::Ordering::Equal) + crate::numeric_cmp::cmp_f64(b.rrf_score, a.rrf_score) .then_with(|| a.document_id.cmp(&b.document_id)) }); fused.truncate(top_k); diff --git a/nodedb-query/src/json_expr.rs b/nodedb-query/src/json_expr.rs index 012841eec..ec55cc236 100644 --- a/nodedb-query/src/json_expr.rs +++ b/nodedb-query/src/json_expr.rs @@ -9,6 +9,7 @@ //! than sorting the row under NULL. use crate::expr::{EvalError, SqlExpr}; +use crate::json_ops::compare_json_numbers; /// Evaluate `expr` against a JSON row. /// @@ -31,9 +32,10 @@ pub fn eval_expr_on_json( /// Total order over JSON values for sorting. /// -/// NULLs sort first; numbers compare numerically; strings lexicographically; -/// mixed or composite values fall back to their rendered form so the ordering -/// stays deterministic. +/// NULLs sort first; numbers compare exactly (integers past `2^53` and +/// integer/float pairs without rounding); strings lexicographically; mixed +/// or composite values fall back to their rendered form so the ordering +/// stays deterministic and total. pub fn compare_json(a: &serde_json::Value, b: &serde_json::Value) -> std::cmp::Ordering { use serde_json::Value; match (a, b) { @@ -41,11 +43,89 @@ pub fn compare_json(a: &serde_json::Value, b: &serde_json::Value) -> std::cmp::O (Value::Null, _) => std::cmp::Ordering::Less, (_, Value::Null) => std::cmp::Ordering::Greater, (Value::Bool(x), Value::Bool(y)) => x.cmp(y), - (Value::Number(x), Value::Number(y)) => x - .as_f64() - .partial_cmp(&y.as_f64()) - .unwrap_or(std::cmp::Ordering::Equal), + (Value::Number(x), Value::Number(y)) => match compare_json_numbers(x, y) { + Some(order) => order, + None => a.to_string().cmp(&b.to_string()), + }, (Value::String(x), Value::String(y)) => x.cmp(y), _ => a.to_string().cmp(&b.to_string()), } } + +#[cfg(test)] +mod tests { + use super::*; + use serde_json::json; + use std::cmp::Ordering; + + fn sorted(mut values: Vec) -> Vec { + values.sort_by(compare_json); + values + } + + /// `2^53 + 1` and `2^53` collapse to one `f64`. Integers order exactly. + #[test] + fn integers_past_two_pow_53_order_exactly() { + let above = json!(9_007_199_254_740_993_i64); + let at = json!(9_007_199_254_740_992_i64); + assert_eq!(compare_json(&above, &at), Ordering::Greater); + assert_eq!(compare_json(&at, &above), Ordering::Less); + assert_eq!(sorted(vec![above.clone(), at.clone()]), vec![at, above]); + } + + #[test] + fn nanosecond_timestamps_order_exactly() { + let t0 = json!(1_700_000_000_000_000_001_i64); + let t1 = json!(1_700_000_000_000_000_002_i64); + let t2 = json!(1_700_000_000_000_000_003_i64); + assert_eq!( + sorted(vec![t2.clone(), t0.clone(), t1.clone()]), + vec![t0, t1, t2] + ); + } + + #[test] + fn u64_against_i64_orders_exactly() { + let u_max = json!(u64::MAX); + let i_max = json!(i64::MAX); + let i_min = json!(i64::MIN); + assert_eq!(compare_json(&u_max, &i_max), Ordering::Greater); + assert_eq!(compare_json(&i_max, &u_max), Ordering::Less); + assert_eq!( + sorted(vec![u_max.clone(), i_min.clone(), i_max.clone()]), + vec![i_min, i_max, u_max] + ); + } + + #[test] + fn integer_against_float_near_two_pow_53_orders_without_rounding() { + let above = json!(9_007_199_254_740_993_i64); + let float = json!(9_007_199_254_740_992.0_f64); + assert_eq!(compare_json(&above, &float), Ordering::Greater); + assert_eq!(compare_json(&float, &above), Ordering::Less); + assert_eq!( + compare_json(&json!(9_007_199_254_740_992_i64), &float), + Ordering::Equal + ); + } + + #[test] + fn small_values_keep_their_order() { + assert_eq!( + sorted(vec![json!(2), json!(null), json!(1.5), json!(-1), json!(1)]), + vec![json!(null), json!(-1), json!(1), json!(1.5), json!(2)] + ); + assert_eq!(compare_json(&json!(3), &json!(3.0)), Ordering::Equal); + assert_eq!(compare_json(&json!("a"), &json!("b")), Ordering::Less); + assert_eq!(compare_json(&json!(false), &json!(true)), Ordering::Less); + } + + /// A NaN has no JSON number form: it reaches the sort as NULL. + #[test] + fn nan_sorts_as_null() { + let nan = json!(f64::NAN); + assert!(nan.is_null()); + assert_eq!(compare_json(&nan, &json!(1)), Ordering::Less); + assert_eq!(compare_json(&nan, &json!(null)), Ordering::Equal); + } +} diff --git a/nodedb-query/src/json_ops.rs b/nodedb-query/src/json_ops.rs index 56feb2d24..929a053bb 100644 --- a/nodedb-query/src/json_ops.rs +++ b/nodedb-query/src/json_ops.rs @@ -7,6 +7,8 @@ use std::cmp::Ordering; +use crate::numeric_cmp::{Numeric, cmp_numeric, numeric_eq, parse_numeric_str}; + /// Coerce a JSON value to f64. /// /// - Numbers: `as_f64()` directly @@ -22,14 +24,52 @@ pub fn json_to_f64(v: &serde_json::Value, coerce_bool: bool) -> Option { } } +/// A JSON number read without rounding an integer through `f64`. +fn number_reading(n: &serde_json::Number) -> Option { + if let Some(i) = n.as_i64() { + return Some(Numeric::Int(i128::from(i))); + } + if let Some(u) = n.as_u64() { + return Some(Numeric::Int(i128::from(u))); + } + n.as_f64().map(Numeric::Float) +} + +/// `v` read as a number: numbers, numeric strings, and bools (`true` = 1). +/// Integer text stays exact, and fractional decimal text reads as an exact +/// `Decimal`. +fn numeric_reading(v: &serde_json::Value) -> Option { + match v { + serde_json::Value::Number(n) => number_reading(n), + serde_json::Value::Bool(b) => Some(Numeric::Int(i128::from(*b))), + serde_json::Value::String(s) => parse_numeric_str(s), + _ => None, + } +} + +/// Exact order of two JSON numbers. `i64` and `u64` pairs compare exactly, +/// and an integer against a float compares without rounding. `None` when a +/// side has no numeric reading. +pub fn compare_json_numbers(a: &serde_json::Number, b: &serde_json::Number) -> Option { + Some(cmp_numeric(number_reading(a)?, number_reading(b)?)) +} + +/// Numeric order of `a` and `b` when both have a numeric reading (number, +/// numeric string, bool). `None` when a side is not numeric. The order is +/// total: NaN sorts above every number and equals NaN. +pub fn numeric_order(a: &serde_json::Value, b: &serde_json::Value) -> Option { + let (na, nb) = (numeric_reading(a)?, numeric_reading(b)?); + Some(cmp_numeric(na, nb)) +} + /// Compare two JSON values with type coercion. /// -/// Tries numeric comparison first (with bool coercion), then falls -/// back to string comparison. +/// Numeric comparison first (with bool coercion, exact for integers and +/// decimal text, NaN above every number), then string comparison of the +/// display forms. pub fn compare_json(a: &serde_json::Value, b: &serde_json::Value) -> Ordering { - // Try numeric comparison with bool coercion. - if let (Some(na), Some(nb)) = (json_to_f64(a, true), json_to_f64(b, true)) { - return na.partial_cmp(&nb).unwrap_or(Ordering::Equal); + if let Some(order) = numeric_order(a, b) { + return order; } // Fallback: string comparison. let sa = json_to_display_string(a); @@ -52,16 +92,17 @@ pub fn compare_json_optional( /// Check equality with type coercion. /// -/// Handles `"5" == 5` by coercing both sides to f64 when one is a -/// number and the other is a numeric string. +/// Handles `"5" == 5` by reading both sides as numbers when each has a +/// numeric reading. A pair with an integer or decimal side compares +/// exactly. A float pair is equal within `f64::EPSILON`. pub fn coerced_eq(a: &serde_json::Value, b: &serde_json::Value) -> bool { if a == b { return true; } - if let (Some(af), Some(bf)) = (json_to_f64(a, true), json_to_f64(b, true)) { - return (af - bf).abs() < f64::EPSILON; + match (numeric_reading(a), numeric_reading(b)) { + (Some(na), Some(nb)) => numeric_eq(na, nb), + _ => false, } - false } /// Check if a JSON value is truthy (for boolean contexts). @@ -141,6 +182,144 @@ mod tests { ); } + /// `2^53 + 1` and `2^53` collapse to one `f64`. Integers compare exactly. + #[test] + fn integers_past_two_pow_53_compare_exactly() { + let above = json!(9_007_199_254_740_993_i64); + let at = json!(9_007_199_254_740_992_i64); + assert!(!coerced_eq(&above, &at)); + assert_eq!(compare_json(&above, &at), Ordering::Greater); + assert_eq!(compare_json(&at, &above), Ordering::Less); + assert!(coerced_eq(&above, &json!(9_007_199_254_740_993_u64))); + } + + #[test] + fn nanosecond_timestamps_one_tick_apart() { + let t0 = json!(1_700_000_000_000_000_001_i64); + let t1 = json!(1_700_000_000_000_000_002_i64); + assert!(!coerced_eq(&t0, &t1)); + assert_eq!(compare_json(&t0, &t1), Ordering::Less); + assert_eq!(compare_json(&t1, &t0), Ordering::Greater); + } + + #[test] + fn u64_against_i64_compares_exactly() { + let u_max = json!(u64::MAX); + let i_max = json!(i64::MAX); + assert!(!coerced_eq(&u_max, &i_max)); + assert_eq!(compare_json(&u_max, &i_max), Ordering::Greater); + assert_eq!(compare_json(&i_max, &u_max), Ordering::Less); + assert_eq!(compare_json(&json!(u64::MAX - 1), &u_max), Ordering::Less); + assert_eq!(compare_json(&json!(i64::MIN), &u_max), Ordering::Less); + // `i64::MAX + 1` as `u64` against `i64::MAX`. + assert_eq!( + compare_json(&json!(9_223_372_036_854_775_808_u64), &i_max), + Ordering::Greater + ); + assert_eq!( + compare_json_numbers( + &serde_json::Number::from(u64::MAX), + &serde_json::Number::from(i64::MAX) + ), + Some(Ordering::Greater) + ); + } + + #[test] + fn integer_against_float_compares_without_rounding() { + let above = json!(9_007_199_254_740_993_i64); + let float = json!(9_007_199_254_740_992.0_f64); + assert!(!coerced_eq(&above, &float)); + assert_eq!(compare_json(&above, &float), Ordering::Greater); + assert_eq!(compare_json(&float, &above), Ordering::Less); + assert!(coerced_eq(&json!(9_007_199_254_740_992_i64), &float)); + assert_eq!(compare_json(&json!(2), &json!(2.5)), Ordering::Less); + assert_eq!(compare_json(&json!(-2), &json!(-2.5)), Ordering::Greater); + assert!(coerced_eq(&json!(3), &json!(3.0))); + // `u64::MAX` rounds up to `2^64` as an `f64`, so it is below that float. + assert_eq!( + compare_json(&json!(u64::MAX), &json!(18_446_744_073_709_551_616.0_f64)), + Ordering::Less + ); + assert!(!coerced_eq( + &json!(u64::MAX), + &json!(18_446_744_073_709_551_616.0_f64) + )); + assert_eq!( + compare_json(&json!(i64::MIN), &json!(-9_223_372_036_854_775_808.0_f64)), + Ordering::Equal + ); + } + + #[test] + fn integer_strings_compare_exactly_against_integers() { + let text = json!("9007199254740993"); + assert!(!coerced_eq(&text, &json!(9_007_199_254_740_992_i64))); + assert!(coerced_eq(&text, &json!(9_007_199_254_740_993_i64))); + assert_eq!( + compare_json(&text, &json!(9_007_199_254_740_992_i64)), + Ordering::Greater + ); + assert!(coerced_eq(&json!("18446744073709551615"), &json!(u64::MAX))); + assert!(coerced_eq(&json!("5.0"), &json!(5))); + } + + #[test] + fn small_numbers_keep_their_order() { + assert_eq!(compare_json(&json!(1), &json!(2)), Ordering::Less); + assert_eq!(compare_json(&json!(1.5), &json!(1.25)), Ordering::Greater); + assert_eq!(compare_json(&json!(-3), &json!(-3)), Ordering::Equal); + assert_eq!(compare_json(&json!(true), &json!(0)), Ordering::Greater); + assert_eq!(compare_json(&json!("abc"), &json!("abd")), Ordering::Less); + // Float equality is exact, as PostgreSQL's float8 `=` is. + assert!(!coerced_eq(&json!(0.1 + 0.2), &json!(0.3))); + assert!(coerced_eq(&json!(0.5), &json!(0.5))); + assert!(!coerced_eq(&json!("abc"), &json!(1))); + } + + /// NaN text sorts above every number and equals NaN, so a sort over it + /// is total. + #[test] + fn nan_text_sorts_above_every_number() { + assert_eq!(compare_json(&json!("NaN"), &json!(1)), Ordering::Greater); + assert_eq!(compare_json(&json!(1e308), &json!("NaN")), Ordering::Less); + assert_eq!(compare_json(&json!("NaN"), &json!("NaN")), Ordering::Equal); + assert_eq!( + compare_json(&json!("NaN"), &json!("inf")), + Ordering::Greater + ); + assert!(coerced_eq(&json!("NaN"), &json!("nan"))); + let mut values = [ + json!("NaN"), + json!(2), + json!("-inf"), + json!("NaN"), + json!(1.5), + ]; + values.sort_by(compare_json); + assert_eq!(values[0], json!("-inf")); + assert_eq!(values[1], json!(1.5)); + assert_eq!(values[2], json!(2)); + } + + /// Two decimal strings one hundredth apart past `2^53` compare exactly. + #[test] + fn fractional_decimal_text_compares_exactly() { + let low = json!("12345678901234567.01"); + let high = json!("12345678901234567.02"); + assert_eq!(compare_json(&low, &high), Ordering::Less); + assert_eq!(compare_json(&high, &low), Ordering::Greater); + assert!(!coerced_eq(&low, &high)); + assert_eq!( + compare_json( + &json!("9007199254740993.5"), + &json!(9_007_199_254_740_993_i64) + ), + Ordering::Greater + ); + assert!(coerced_eq(&json!("2.50"), &json!("2.5"))); + } + #[test] fn truthiness() { assert!(is_truthy(&json!(true))); diff --git a/nodedb-query/src/lib.rs b/nodedb-query/src/lib.rs index e3de07134..b973ec1a0 100644 --- a/nodedb-query/src/lib.rs +++ b/nodedb-query/src/lib.rs @@ -20,6 +20,7 @@ pub mod json_expr; pub mod json_ops; pub mod metadata_filter; pub mod msgpack_scan; +pub mod numeric_sum; pub mod partition_hash; pub mod scan_filter; pub mod simd_agg; @@ -37,6 +38,8 @@ pub use fusion::{ reciprocal_rank_fusion_linear, reciprocal_rank_fusion_weighted, }; pub use json_expr::{compare_json, eval_expr_on_json}; +pub use nodedb_types::numeric_cmp; +pub use numeric_sum::ExactSum; pub use partition_hash::{partition_hash, partition_hash_seeded}; pub use scan_filter::ScanFilter; pub use window::{ diff --git a/nodedb-query/src/metadata_filter.rs b/nodedb-query/src/metadata_filter.rs index 77c77f0a9..622e917fa 100644 --- a/nodedb-query/src/metadata_filter.rs +++ b/nodedb-query/src/metadata_filter.rs @@ -1,97 +1,219 @@ // SPDX-License-Identifier: Apache-2.0 -//! Bridge between `MetadataFilter` (nodedb-types) and document evaluation. +//! Evaluation of `MetadataFilter` (nodedb-types) against a document. //! -//! Converts the typed `MetadataFilter` enum into runtime evaluation against -//! JSON documents. Used by both Origin (vector search pre-filter) and Lite -//! (vector search post-filter against CRDT state). +//! One evaluator serves every document shape: a field lookup closure +//! resolves top-level fields. Adapters cover JSON documents (shape sync, +//! vector post-filter), field maps (Lite edge properties) and plain-msgpack +//! maps (Origin edge properties). +//! +//! Semantics: +//! - `Eq` on a missing field matches only a `Null` value. `Ne` negates `Eq`. +//! - `Gt`/`Gte`/`Lt`/`Lte` are false when the field is missing or null, the +//! value is null, or the pair has no order. +//! - `In` on a missing field is false. `NotIn` on a missing field is true. +//! - `And([])` is true. `Or([])` is false. +//! +//! Coercion, one rule for equality and order: +//! - Numeric: numbers, decimals, numeric strings and bools (`true` = 1) +//! compare as numbers. Integer pairs compare exactly. +//! - Instant: a typed instant compares with an instant or an ISO-8601 string +//! by epoch microseconds. +//! - Text: strings, UUIDs and ULIDs order lexicographically. +//! - Any other pair of different kinds is unequal and has no order. So +//! `{"score": "n/a"}` fails `score > 5`, and an array fails every order. +use std::borrow::Cow; +use std::cmp::Ordering; +use std::collections::HashMap; +use std::convert::Infallible; + +use nodedb_types::Value; use nodedb_types::filter::MetadataFilter; -use crate::json_ops::{coerced_eq, compare_json_optional}; +use crate::msgpack_scan::FieldIndex; +use crate::msgpack_scan::reader::read_value; +use crate::value_ops::{coerced_eq, involves_instant, numeric_order}; -/// Evaluate a `MetadataFilter` against a JSON document. -/// -/// Returns `true` if the document matches the filter. -pub fn matches_metadata_filter(doc: &serde_json::Value, filter: &MetadataFilter) -> bool { - match filter { - MetadataFilter::Eq { field, value } => { - let field_val = doc.get(field.as_str()); - let filter_val = value_to_json(value); - match field_val { - Some(fv) => coerced_eq(fv, &filter_val), - None => filter_val.is_null(), - } - } - MetadataFilter::Ne { field, value } => { - let field_val = doc.get(field.as_str()); - let filter_val = value_to_json(value); - match field_val { - Some(fv) => !coerced_eq(fv, &filter_val), - None => !filter_val.is_null(), - } - } +/// A plain-msgpack property map the evaluator cannot read. +#[derive(Debug, Clone, PartialEq, Eq, thiserror::Error)] +pub enum PropertyMapError { + /// The bytes are not one well-formed MessagePack map. + #[error("property bytes are not a well-formed MessagePack map")] + NotAMap, + /// The map holds `field`, but its value does not decode. + #[error("property field '{field}' does not decode")] + Field { field: String }, +} + +/// Evaluate `filter` against a document whose top-level fields `lookup` resolves. +pub fn matches_metadata_filter_with<'a, F>(lookup: &F, filter: &MetadataFilter) -> bool +where + F: Fn(&str) -> Option>, +{ + let infallible = |name: &str| Ok::<_, Infallible>(lookup(name)); + match try_matches_metadata_filter_with(&infallible, filter) { + Ok(admitted) => admitted, + Err(never) => match never {}, + } +} + +/// Evaluate `filter` against a document whose top-level fields `lookup` +/// resolves. A lookup error stops evaluation and returns that error. +pub fn try_matches_metadata_filter_with<'a, F, E>( + lookup: &F, + filter: &MetadataFilter, +) -> Result +where + F: Fn(&str) -> Result>, E>, +{ + Ok(match filter { + MetadataFilter::Eq { field, value } => equals(lookup(field)?.as_deref(), value), + MetadataFilter::Ne { field, value } => !equals(lookup(field)?.as_deref(), value), MetadataFilter::Gt { field, value } => { - let field_val = doc.get(field.as_str()); - let filter_val = value_to_json(value); - compare_json_optional(field_val, Some(&filter_val)) == std::cmp::Ordering::Greater + ordered(lookup(field)?.as_deref(), value, |o| o == Ordering::Greater) } MetadataFilter::Gte { field, value } => { - let field_val = doc.get(field.as_str()); - let filter_val = value_to_json(value); - let cmp = compare_json_optional(field_val, Some(&filter_val)); - cmp == std::cmp::Ordering::Greater || cmp == std::cmp::Ordering::Equal + ordered(lookup(field)?.as_deref(), value, |o| o != Ordering::Less) } MetadataFilter::Lt { field, value } => { - let field_val = doc.get(field.as_str()); - let filter_val = value_to_json(value); - compare_json_optional(field_val, Some(&filter_val)) == std::cmp::Ordering::Less + ordered(lookup(field)?.as_deref(), value, |o| o == Ordering::Less) } MetadataFilter::Lte { field, value } => { - let field_val = doc.get(field.as_str()); - let filter_val = value_to_json(value); - let cmp = compare_json_optional(field_val, Some(&filter_val)); - cmp == std::cmp::Ordering::Less || cmp == std::cmp::Ordering::Equal + ordered(lookup(field)?.as_deref(), value, |o| o != Ordering::Greater) } MetadataFilter::In { field, values } => { - let field_val = match doc.get(field.as_str()) { - Some(v) => v, - None => return false, - }; - values - .iter() - .any(|v| coerced_eq(field_val, &value_to_json(v))) + lookup(field)?.is_some_and(|found| values.iter().any(|v| filter_eq(&found, v))) } MetadataFilter::NotIn { field, values } => { - let field_val = match doc.get(field.as_str()) { - Some(v) => v, - None => return true, - }; - !values - .iter() - .any(|v| coerced_eq(field_val, &value_to_json(v))) - } - MetadataFilter::And(filters) => filters.iter().all(|f| matches_metadata_filter(doc, f)), - MetadataFilter::Or(filters) => filters.iter().any(|f| matches_metadata_filter(doc, f)), - MetadataFilter::Not(inner) => !matches_metadata_filter(doc, inner), + lookup(field)?.is_none_or(|found| !values.iter().any(|v| filter_eq(&found, v))) + } + MetadataFilter::And(children) => { + for child in children { + if !try_matches_metadata_filter_with(lookup, child)? { + return Ok(false); + } + } + true + } + MetadataFilter::Or(children) => { + for child in children { + if try_matches_metadata_filter_with(lookup, child)? { + return Ok(true); + } + } + false + } + MetadataFilter::Not(inner) => !try_matches_metadata_filter_with(lookup, inner)?, + // `MetadataFilter` is `#[non_exhaustive]`: a variant this evaluator + // does not know admits nothing. + _ => false, + }) +} + +fn equals(found: Option<&Value>, value: &Value) -> bool { + match found { + Some(found) => filter_eq(found, value), + None => value.is_null(), + } +} + +/// Equality under the module's coercion rule: [`coerced_eq`], plus text +/// kinds equal by their text. +fn filter_eq(a: &Value, b: &Value) -> bool { + coerced_eq(a, b) || matches!((text(a), text(b)), (Some(x), Some(y)) if x == y) +} + +fn ordered(found: Option<&Value>, value: &Value, accept: impl Fn(Ordering) -> bool) -> bool { + match found { + Some(found) if !found.is_null() && !value.is_null() => { + filter_order(found, value).is_some_and(accept) + } _ => false, } } -/// Convert a `nodedb_types::Value` to `serde_json::Value` for comparison. -fn value_to_json(value: &nodedb_types::value::Value) -> serde_json::Value { - match value { - nodedb_types::value::Value::Null => serde_json::Value::Null, - nodedb_types::value::Value::Bool(b) => serde_json::Value::Bool(*b), - nodedb_types::value::Value::Integer(i) => serde_json::json!(i), - nodedb_types::value::Value::Float(f) => serde_json::Number::from_f64(*f) - .map(serde_json::Value::Number) - .unwrap_or(serde_json::Value::Null), - nodedb_types::value::Value::String(s) => serde_json::Value::String(s.clone()), - _ => serde_json::to_value(value).unwrap_or(serde_json::Value::Null), +/// Order under the module's coercion rule. `None` for a pair of different +/// kinds outside the coercions. NaN sorts above every number and equals NaN. +fn filter_order(a: &Value, b: &Value) -> Option { + if involves_instant(a, b) { + return a.partial_cmp_coerced(b); + } + if let Some(order) = numeric_order(a, b) { + return Some(order); + } + match (text(a), text(b)) { + (Some(x), Some(y)) => Some(x.cmp(y)), + _ => None, } } +/// The text of a string-kind value: a string, UUID or ULID. +fn text(v: &Value) -> Option<&str> { + match v { + Value::String(s) | Value::Uuid(s) | Value::Ulid(s) => Some(s), + _ => None, + } +} + +/// `filter` against a field map (Lite edge properties). +pub fn matches_metadata_fields(fields: &HashMap, filter: &MetadataFilter) -> bool { + matches_metadata_filter_with(&|name: &str| fields.get(name).map(Cow::Borrowed), filter) +} + +/// Every filter of `filters` against a field map. An empty list admits. +pub fn matches_all_fields(fields: &HashMap, filters: &[MetadataFilter]) -> bool { + filters.iter().all(|f| matches_metadata_fields(fields, f)) +} + +/// Every filter of `filters` against a plain-msgpack map. Decodes only the +/// fields a filter names. Empty bytes evaluate as `{}`. An empty filter list +/// admits without reading `doc`. +/// +/// Bytes that are not one well-formed map, and a named field whose value +/// does not decode, are an error: a property map the predicate cannot read +/// is not a map the predicate admitted. +pub fn matches_all_msgpack( + doc: &[u8], + filters: &[MetadataFilter], +) -> Result { + if filters.is_empty() { + return Ok(true); + } + let index = if doc.is_empty() { + FieldIndex::empty() + } else { + FieldIndex::build(doc, 0).ok_or(PropertyMapError::NotAMap)? + }; + let lookup = |name: &str| -> Result>, PropertyMapError> { + let Some((start, end)) = index.get(name) else { + return Ok(None); + }; + let decoded = read_value(doc, start).or_else(|| { + doc.get(start..end) + .and_then(|bytes| nodedb_types::json_msgpack::value_from_msgpack(bytes).ok()) + }); + match decoded { + Some(value) => Ok(Some(Cow::Owned(value))), + None => Err(PropertyMapError::Field { + field: name.to_owned(), + }), + } + }; + for filter in filters { + if !try_matches_metadata_filter_with(&lookup, filter)? { + return Ok(false); + } + } + Ok(true) +} + +/// `filter` against a JSON document (shape-sync predicates, vector post-filter). +pub fn matches_metadata_filter(doc: &serde_json::Value, filter: &MetadataFilter) -> bool { + let lookup = |name: &str| doc.get(name).map(|v| Cow::Owned(Value::from(v.clone()))); + matches_metadata_filter_with(&lookup, filter) +} + #[cfg(test)] mod tests { use super::*; @@ -99,6 +221,27 @@ mod tests { use nodedb_types::value::Value; use serde_json::json; + fn gt(field: &str, value: Value) -> MetadataFilter { + MetadataFilter::Gt { + field: field.into(), + value, + } + } + + fn lt(field: &str, value: Value) -> MetadataFilter { + MetadataFilter::Lt { + field: field.into(), + value, + } + } + + fn ne(field: &str, value: Value) -> MetadataFilter { + MetadataFilter::Ne { + field: field.into(), + value, + } + } + #[test] fn eq_match() { let doc = json!({"status": "active", "age": 25}); @@ -116,11 +259,10 @@ mod tests { #[test] fn gt_numeric() { let doc = json!({"age": 30}); - let filter = MetadataFilter::Gt { - field: "age".into(), - value: Value::Integer(25), - }; - assert!(matches_metadata_filter(&doc, &filter)); + assert!(matches_metadata_filter( + &doc, + >("age", Value::Integer(25)) + )); } #[test] @@ -128,10 +270,7 @@ mod tests { let doc = json!({"status": "active", "age": 30}); let filter = MetadataFilter::and(vec![ MetadataFilter::eq("status", "active"), - MetadataFilter::Gt { - field: "age".into(), - value: Value::Integer(25), - }, + gt("age", Value::Integer(25)), ]); assert!(matches_metadata_filter(&doc, &filter)); } @@ -141,10 +280,7 @@ mod tests { let doc = json!({"status": "inactive", "age": 30}); let filter = MetadataFilter::or(vec![ MetadataFilter::eq("status", "active"), - MetadataFilter::Gt { - field: "age".into(), - value: Value::Integer(25), - }, + gt("age", Value::Integer(25)), ]); assert!(matches_metadata_filter(&doc, &filter)); } @@ -172,4 +308,327 @@ mod tests { let filter = MetadataFilter::eq("status", "active"); assert!(!matches_metadata_filter(&doc, &filter)); } + + #[test] + fn eq_and_ne_on_missing_field_and_null() { + let doc = json!({"present": null}); + assert!(matches_metadata_filter( + &doc, + &MetadataFilter::eq("absent", Value::Null) + )); + assert!(matches_metadata_filter( + &doc, + &MetadataFilter::eq("present", Value::Null) + )); + assert!(!matches_metadata_filter(&doc, &ne("absent", Value::Null))); + assert!(matches_metadata_filter( + &doc, + &ne("absent", Value::from("x")) + )); + assert!(!matches_metadata_filter( + &doc, + &MetadataFilter::eq("absent", "x") + )); + } + + #[test] + fn eq_coerces_numeric_strings_and_bools() { + let doc = json!({"n": "5", "flag": true}); + assert!(matches_metadata_filter( + &doc, + &MetadataFilter::eq("n", Value::Integer(5)) + )); + assert!(matches_metadata_filter( + &doc, + &MetadataFilter::eq("flag", Value::Integer(1)) + )); + } + + #[test] + fn ordered_comparisons_reject_missing_and_null() { + let doc = json!({"score": null}); + for filter in [ + gt("missing", Value::Integer(0)), + lt("missing", Value::Integer(100)), + MetadataFilter::Lte { + field: "missing".into(), + value: Value::Integer(100), + }, + MetadataFilter::Gte { + field: "missing".into(), + value: Value::Integer(-100), + }, + lt("score", Value::Integer(100)), + gt("score", Value::Integer(-100)), + ] { + assert!(!matches_metadata_filter(&doc, &filter), "{filter:?}"); + } + let doc = json!({"score": 3}); + assert!(!matches_metadata_filter(&doc, <("score", Value::Null))); + } + + /// NaN sorts above every number and equals NaN, as in PostgreSQL. + #[test] + fn ordered_comparisons_place_nan_above_every_number() { + let fields = HashMap::from([("score".to_owned(), Value::Float(f64::NAN))]); + let nan_gte_nan = MetadataFilter::Gte { + field: "score".into(), + value: Value::Float(f64::NAN), + }; + for filter in [gt("score", Value::Integer(0)), nan_gte_nan] { + assert!(matches_metadata_fields(&fields, &filter), "{filter:?}"); + } + let nan_lte_zero = MetadataFilter::Lte { + field: "score".into(), + value: Value::Integer(0), + }; + for filter in [lt("score", Value::Integer(0)), nan_lte_zero] { + assert!(!matches_metadata_fields(&fields, &filter), "{filter:?}"); + } + } + + #[test] + fn in_and_not_in_on_missing_field() { + let doc = json!({}); + let values = vec![Value::from("a")]; + assert!(!matches_metadata_filter( + &doc, + &MetadataFilter::In { + field: "k".into(), + values: values.clone(), + } + )); + assert!(matches_metadata_filter( + &doc, + &MetadataFilter::NotIn { + field: "k".into(), + values, + } + )); + } + + #[test] + fn empty_and_admits_empty_or_rejects() { + let doc = json!({}); + assert!(matches_metadata_filter( + &doc, + &MetadataFilter::And(Vec::new()) + )); + assert!(!matches_metadata_filter( + &doc, + &MetadataFilter::Or(Vec::new()) + )); + } + + #[test] + fn nested_not() { + let doc = json!({"score": 9}); + let filter = MetadataFilter::Not(Box::new(MetadataFilter::Not(Box::new(gt( + "score", + Value::Integer(5), + ))))); + assert!(matches_metadata_filter(&doc, &filter)); + let filter = MetadataFilter::Not(Box::new(MetadataFilter::Or(vec![ + lt("score", Value::Integer(5)), + MetadataFilter::eq("score", Value::Integer(9)), + ]))); + assert!(!matches_metadata_filter(&doc, &filter)); + } + + #[test] + fn msgpack_matches_field_map() { + let fields = HashMap::from([ + ("score".to_owned(), Value::Integer(9)), + ("kind".to_owned(), Value::from("road")), + ("closed".to_owned(), Value::Bool(false)), + ]); + let bytes = nodedb_types::json_msgpack::value_to_msgpack(&Value::Object(fields.clone())) + .expect("encode"); + let cases = [ + vec![gt("score", Value::Integer(5))], + vec![lt("score", Value::Integer(5))], + vec![ + MetadataFilter::In { + field: "kind".into(), + values: vec![Value::from("road"), Value::from("rail")], + }, + MetadataFilter::Not(Box::new(MetadataFilter::eq("closed", true))), + ], + vec![MetadataFilter::eq("missing", Value::Null)], + vec![lt("missing", Value::Integer(5))], + Vec::new(), + ]; + for filters in cases { + assert_eq!( + matches_all_msgpack(&bytes, &filters), + Ok(matches_all_fields(&fields, &filters)), + "{filters:?}" + ); + } + } + + #[test] + fn empty_msgpack_evaluates_as_empty_object() { + assert_eq!( + matches_all_msgpack(&[], &[MetadataFilter::eq("missing", Value::Null)]), + Ok(true) + ); + assert_eq!( + matches_all_msgpack(&[], &[lt("score", Value::Integer(5))]), + Ok(false) + ); + assert_eq!(matches_all_msgpack(&[], &[]), Ok(true)); + } + + /// `{"score": 9}` with the value's tag replaced by the reserved `0xc1`. + fn corrupt_score_map() -> Vec { + let mut bytes = + nodedb_types::json_msgpack::json_to_msgpack(&json!({"score": 9})).expect("encode"); + let last = bytes.len() - 1; + bytes[last] = 0xc1; + bytes + } + + #[test] + fn a_named_field_that_does_not_decode_is_an_error() { + let bytes = corrupt_score_map(); + let result = matches_all_msgpack(&bytes, &[gt("score", Value::Integer(5))]); + assert!(result.is_err(), "{result:?}"); + assert_eq!( + matches_all_msgpack(&bytes, &[]), + Ok(true), + "no filter reads nothing" + ); + } + + #[test] + fn bytes_that_are_not_a_map_are_an_error() { + let list = nodedb_types::json_msgpack::json_to_msgpack(&json!([1, 2])).expect("encode"); + assert_eq!( + matches_all_msgpack(&list, &[MetadataFilter::eq("k", Value::Null)]), + Err(PropertyMapError::NotAMap) + ); + // A map header that claims more entries than the bytes hold. + assert_eq!( + matches_all_msgpack(&[0x82, 0xa1, b'k', 0x01], &[gt("k", Value::Integer(0))]), + Err(PropertyMapError::NotAMap) + ); + } + + #[test] + fn integers_past_two_pow_53_filter_exactly() { + let fields = HashMap::from([("ts".to_owned(), Value::Integer(9_007_199_254_740_993))]); + let at = Value::Integer(9_007_199_254_740_992); + assert!(!matches_metadata_fields( + &fields, + &MetadataFilter::eq("ts", at.clone()) + )); + assert!(matches_metadata_fields(&fields, >("ts", at.clone()))); + assert!(!matches_metadata_fields(&fields, <("ts", at))); + } + + /// Each pair of different kinds outside the coercions has no order. + #[test] + fn ordered_comparisons_across_kinds_are_unordered() { + let fields = HashMap::from([ + ("text".to_owned(), Value::from("n/a")), + ("list".to_owned(), Value::Array(vec![Value::Integer(9)])), + ("num".to_owned(), Value::Integer(9)), + ( + "obj".to_owned(), + Value::Object(HashMap::from([("a".to_owned(), Value::Integer(1))])), + ), + ("flag".to_owned(), Value::Bool(true)), + ]); + let unordered = [ + ("text", Value::Integer(5)), + ("list", Value::Integer(5)), + ("num", Value::from("abc")), + ("num", Value::Array(vec![Value::Integer(1)])), + ("obj", Value::Integer(0)), + ("obj", Value::from("a")), + ("list", Value::from("a")), + ("flag", Value::from("yes")), + ]; + for (field, value) in unordered { + for filter in [ + gt(field, value.clone()), + lt(field, value.clone()), + MetadataFilter::Gte { + field: field.into(), + value: value.clone(), + }, + MetadataFilter::Lte { + field: field.into(), + value: value.clone(), + }, + ] { + assert!(!matches_metadata_fields(&fields, &filter), "{filter:?}"); + } + } + } + + /// The coercions order the same pairs that compare equal. + #[test] + fn coerced_kinds_order_and_compare_alike() { + let fields = HashMap::from([ + ("n".to_owned(), Value::from("10")), + ("flag".to_owned(), Value::Bool(true)), + ("name".to_owned(), Value::from("bob")), + ( + "id".to_owned(), + Value::Uuid("550e8400-e29b-41d4-a716-446655440000".into()), + ), + ]); + assert!(matches_metadata_fields( + &fields, + >("n", Value::Integer(9)) + )); + assert!(matches_metadata_fields( + &fields, + &MetadataFilter::eq("n", Value::Integer(10)) + )); + assert!(matches_metadata_fields( + &fields, + >("flag", Value::Integer(0)) + )); + assert!(matches_metadata_fields( + &fields, + &MetadataFilter::eq("flag", Value::Integer(1)) + )); + assert!(matches_metadata_fields( + &fields, + >("name", Value::from("alice")) + )); + assert!(matches_metadata_fields( + &fields, + <("name", Value::from("carol")) + )); + let id = Value::from("550e8400-e29b-41d4-a716-446655440000"); + assert!(matches_metadata_fields( + &fields, + &MetadataFilter::eq("id", id.clone()) + )); + assert!(matches_metadata_fields( + &fields, + &MetadataFilter::Gte { + field: "id".into(), + value: id, + } + )); + } + + #[test] + fn a_non_numeric_string_fails_a_numeric_order_over_msgpack() { + let bytes = + nodedb_types::json_msgpack::json_to_msgpack(&json!({"score": "n/a"})).expect("encode"); + assert_eq!( + matches_all_msgpack(&bytes, &[gt("score", Value::Integer(5))]), + Ok(false) + ); + assert_eq!( + matches_all_msgpack(&bytes, &[lt("score", Value::Integer(5))]), + Ok(false) + ); + } } diff --git a/nodedb-query/src/msgpack_scan/aggregate.rs b/nodedb-query/src/msgpack_scan/aggregate.rs deleted file mode 100644 index e101b3661..000000000 --- a/nodedb-query/src/msgpack_scan/aggregate.rs +++ /dev/null @@ -1,735 +0,0 @@ -// SPDX-License-Identifier: Apache-2.0 - -//! Zero-deserialization aggregate computation on raw MessagePack documents. -//! -//! Replaces `compute_aggregate(op, field, docs: &[serde_json::Value])` with -//! direct binary field extraction. Each document is `&[u8]` (MessagePack map). -//! When an expression is provided, decodes msgpack → `nodedb_types::Value` -//! directly (no JSON intermediate) and evaluates the expression once per document. - -use std::cmp::Ordering; -use std::collections::HashSet; - -use nodedb_types::Value; - -use crate::expr::EvalError; -use crate::msgpack_scan::compare::compare_field_bytes; -use crate::msgpack_scan::field::extract_field; -use crate::msgpack_scan::reader::{read_f64, read_null, read_str}; -use crate::value_ops; - -/// Compute an aggregate function over raw MessagePack documents. -/// -/// Each entry in `docs` is a complete MessagePack map (the raw bytes from storage). -/// Returns the result as `Value` — conversion to JSON happens at the -/// response boundary only. -pub fn compute_aggregate_binary( - op: &str, - field: &str, - expr: Option<&crate::expr::SqlExpr>, - docs: &[&[u8]], -) -> Result { - Ok(match op { - "count" => { - if field == "*" && expr.is_none() { - Value::Integer(docs.len() as i64) - } else { - let count = docs - .iter() - .map(|d| extract_as_value(d, field, expr)) - .collect::, _>>()? - .into_iter() - .flatten() - .filter(|v| !v.is_null()) - .count(); - Value::Integer(count as i64) - } - } - - "sum" => { - let total: f64 = docs - .iter() - .map(|d| extract_f64_val(d, field, expr)) - .collect::, _>>()? - .into_iter() - .flatten() - .sum(); - Value::Float(total) - } - - "avg" => { - let (sum, count) = docs - .iter() - .map(|d| extract_f64_val(d, field, expr)) - .collect::, _>>()? - .into_iter() - .flatten() - .fold((0.0f64, 0u64), |(s, c), v| (s + v, c + 1)); - if count == 0 { - Value::Null - } else { - Value::Float(sum / count as f64) - } - } - - "min" => find_minmax(docs, field, expr, false)?, - "max" => find_minmax(docs, field, expr, true)?, - - "count_distinct" => { - let mut seen = HashSet::new(); - for doc in docs { - if let Some(bytes) = extract_value_bytes(doc, field, expr)? - && !value_bytes_are_null(&bytes) - { - seen.insert(bytes); - } - } - Value::Integer(seen.len() as i64) - } - - "stddev" | "stddev_pop" => { - stat_aggregate(docs, field, expr, |variance, _n| variance.sqrt(), true)? - } - - "stddev_samp" => stat_aggregate(docs, field, expr, |variance, _n| variance.sqrt(), false)?, - - "variance" | "var_pop" => stat_aggregate(docs, field, expr, |variance, _n| variance, true)?, - - "var_samp" => stat_aggregate(docs, field, expr, |variance, _n| variance, false)?, - - "array_agg" => { - let values: Vec = docs - .iter() - .map(|d| extract_as_value(d, field, expr)) - .collect::, _>>()? - .into_iter() - .flatten() - .filter(|v| !v.is_null()) - .collect(); - Value::Array(values) - } - - "array_agg_distinct" => { - let mut seen_bytes = HashSet::new(); - let mut values = Vec::new(); - for doc in docs { - // When expr is present, evaluate once and derive both bytes and value - // from the result to avoid double-decoding the document. - if let Some(expr) = expr { - let Some(val) = eval_expr_on_doc(doc, expr)? else { - continue; - }; - if val.is_null() { - continue; - } - let bytes = zerompk::to_msgpack_vec(&val).unwrap_or_default(); - if seen_bytes.insert(bytes) { - values.push(val); - } - } else if let Some(bytes) = extract_value_bytes(doc, field, None)? - && !value_bytes_are_null(&bytes) - && seen_bytes.insert(bytes) - && let Some(v) = value_from_field(doc, field) - { - values.push(v); - } - } - Value::Array(values) - } - - "string_agg" | "group_concat" => { - let values: Vec = docs - .iter() - .map(|d| extract_str_val(d, field, expr)) - .collect::, _>>()? - .into_iter() - .flatten() - .collect(); - Value::String(values.join(",")) - } - - "approx_count_distinct" => { - let mut hll = nodedb_types::approx::HyperLogLog::new(); - for doc in docs { - if let Some(bytes) = extract_value_bytes(doc, field, expr)? - && !value_bytes_are_null(&bytes) - { - // Hash the raw bytes for HLL. - let hash = hash_bytes(&bytes); - hll.add(hash); - } - } - Value::Integer(hll.estimate().round() as i64) - } - - "approx_percentile" => { - // Format: field is "quantile:actual_field" (e.g. "0.95:latency"). - let (pct, actual_field) = if let Some(idx) = field.find(':') { - match field[..idx].parse::() { - Ok(p) => (p, &field[idx + 1..]), - Err(_) => return Ok(Value::Null), // invalid quantile - } - } else { - (0.5, field) - }; - let mut digest = nodedb_types::approx::TDigest::new(); - for doc in docs { - if let Some(v) = extract_f64_val(doc, actual_field, expr)? { - digest.add(v); - } - } - let result = digest.quantile(pct); - if result.is_nan() { - Value::Null - } else { - Value::Float(result) - } - } - - "approx_topk" => { - // Format: field is "k:actual_field" (e.g. "10:region"). - let (k, actual_field) = if let Some(idx) = field.find(':') { - match field[..idx].parse::() { - Ok(k) => (k, &field[idx + 1..]), - Err(_) => return Ok(Value::Null), // invalid k - } - } else { - (10, field) - }; - let mut ss = nodedb_types::approx::SpaceSaving::new(k); - for doc in docs { - if let Some(bytes) = extract_value_bytes(doc, actual_field, expr)? - && !value_bytes_are_null(&bytes) - { - ss.add(hash_bytes(&bytes)); - } - } - // Return as array of [hash, count, error] tuples. - let top = ss.top_k(); - let arr: Vec = top - .into_iter() - .map(|(item, count, error)| { - Value::Object( - [ - ("item".to_string(), Value::Integer(item as i64)), - ("count".to_string(), Value::Integer(count as i64)), - ("error".to_string(), Value::Integer(error as i64)), - ] - .into_iter() - .collect(), - ) - }) - .collect(); - Value::Array(arr) - } - - "percentile_cont" => { - let (pct, actual_field) = if let Some(idx) = field.find(':') { - match field[..idx].parse::() { - Ok(p) => (p, &field[idx + 1..]), - Err(_) => return Ok(Value::Null), // invalid quantile - } - } else { - (0.5, field) - }; - let mut values: Vec = docs - .iter() - .map(|d| extract_f64_val(d, actual_field, expr)) - .collect::, _>>()? - .into_iter() - .flatten() - .collect(); - if values.is_empty() { - return Ok(Value::Null); - } - values.sort_by(|a, b| a.partial_cmp(b).unwrap_or(Ordering::Equal)); - let idx = (pct * (values.len() - 1) as f64).clamp(0.0, (values.len() - 1) as f64); - let lower = idx.floor() as usize; - let upper = idx.ceil() as usize; - let frac = idx - lower as f64; - let result = values[lower] * (1.0 - frac) + values[upper] * frac; - Value::Float(result) - } - - _ => Value::Null, - }) -} - -// ── Internal helpers ─────────────────────────────────────────────────── - -/// Decode a msgpack document directly to `nodedb_types::Value` and evaluate -/// the expression. No JSON intermediate — msgpack → Value → eval → Value. -/// -/// `Ok(None)` means the row is skipped (document could not be decoded); -/// `Err(EvalError::DivisionByZero)` means the expression divided/modded by -/// zero and must propagate as a statement failure rather than being folded -/// to `None`. -#[inline] -fn eval_expr_on_doc(doc: &[u8], expr: &crate::expr::SqlExpr) -> Result, EvalError> { - let Ok(doc_val) = nodedb_types::json_msgpack::value_from_msgpack(doc) else { - return Ok(None); - }; - Ok(Some(expr.eval(&doc_val)?)) -} - -/// Extract a numeric value from a field or expression result. -#[inline] -fn extract_f64_val( - doc: &[u8], - field: &str, - expr: Option<&crate::expr::SqlExpr>, -) -> Result, EvalError> { - if let Some(expr) = expr { - return Ok(eval_expr_on_doc(doc, expr)?.and_then(|v| value_ops::value_to_f64(&v, false))); - } - let Some((start, _end)) = extract_field(doc, 0, field) else { - return Ok(None); - }; - Ok(read_f64(doc, start)) -} - -/// Extract a string from a field or expression result. -fn extract_str_val( - doc: &[u8], - field: &str, - expr: Option<&crate::expr::SqlExpr>, -) -> Result, EvalError> { - if let Some(expr) = expr { - return Ok(eval_expr_on_doc(doc, expr)?.map(|v| value_ops::value_to_display_string(&v))); - } - let Some((start, _end)) = extract_field(doc, 0, field) else { - return Ok(None); - }; - Ok(read_str(doc, start).map(|s| s.to_string())) -} - -/// Extract a field as `Value`. Uses direct msgpack→Value for scalars; -/// falls back to full decode only for complex types. -fn extract_as_value( - doc: &[u8], - field: &str, - expr: Option<&crate::expr::SqlExpr>, -) -> Result, EvalError> { - if let Some(expr) = expr { - return eval_expr_on_doc(doc, expr); - } - Ok(value_from_field(doc, field)) -} - -#[inline] -fn value_from_field(doc: &[u8], field: &str) -> Option { - let (start, end) = extract_field(doc, 0, field)?; - // Fast path: scalar types (null, bool, int, float, string). - if let Some(v) = crate::msgpack_scan::reader::read_value(doc, start) { - return Some(v); - } - // Slow path: complex types (array, map, bin) — decode field bytes directly. - let field_bytes = &doc[start..end]; - nodedb_types::json_msgpack::value_from_msgpack(field_bytes).ok() -} - -/// Find min or max across docs by comparing raw field bytes. -fn find_minmax( - docs: &[&[u8]], - field: &str, - expr: Option<&crate::expr::SqlExpr>, - want_max: bool, -) -> Result { - if let Some(expr) = expr { - // Evaluate expression once per doc; compare on Value - // since the result may be any type (not a raw field). - let mut best: Option = None; - for doc in docs { - let Some(value) = eval_expr_on_doc(doc, expr)? else { - continue; - }; - if value.is_null() { - continue; - } - let replace = match &best { - None => true, - Some(current) => { - let ord = value_ops::compare_values(&value, current); - if want_max { - ord == Ordering::Greater - } else { - ord == Ordering::Less - } - } - }; - if replace { - best = Some(value); - } - } - return Ok(best.unwrap_or(Value::Null)); - } - - let mut best_doc: Option<&[u8]> = None; - let mut best_range: Option<(usize, usize)> = None; - - for doc in docs { - if let Some(range) = extract_field(doc, 0, field) { - if read_null(doc, range.0) { - continue; - } - match best_range { - None => { - best_doc = Some(doc); - best_range = Some(range); - } - Some(br) => { - let Some(bd) = best_doc else { continue }; - let cmp = compare_field_bytes(doc, range, bd, br); - let replace = if want_max { - cmp == Ordering::Greater - } else { - cmp == Ordering::Less - }; - if replace { - best_doc = Some(doc); - best_range = Some(range); - } - } - } - } - } - - Ok(match (best_doc, best_range) { - (Some(doc), Some((start, end))) => { - if let Some(v) = crate::msgpack_scan::reader::read_value(doc, start) { - return Ok(v); - } - let bytes = &doc[start..end]; - nodedb_types::json_msgpack::value_from_msgpack(bytes).unwrap_or(Value::Null) - } - _ => Value::Null, - }) -} - -/// Compute stddev or variance. `population` = true for population variant. -/// `finalize` transforms the variance into the final result. -fn stat_aggregate( - docs: &[&[u8]], - field: &str, - expr: Option<&crate::expr::SqlExpr>, - finalize: fn(f64, usize) -> f64, - population: bool, -) -> Result { - let values: Vec = docs - .iter() - .map(|d| extract_f64_val(d, field, expr)) - .collect::, _>>()? - .into_iter() - .flatten() - .collect(); - if values.len() < 2 { - return Ok(Value::Null); - } - let mean = values.iter().sum::() / values.len() as f64; - let divisor = if population { - values.len() as f64 - } else { - (values.len() - 1) as f64 - }; - let variance = values.iter().map(|v| (v - mean).powi(2)).sum::() / divisor; - Ok(Value::Float(finalize(variance, values.len()))) -} - -fn extract_value_bytes( - doc: &[u8], - field: &str, - expr: Option<&crate::expr::SqlExpr>, -) -> Result>, EvalError> { - if let Some(expr) = expr { - let Some(val) = eval_expr_on_doc(doc, expr)? else { - return Ok(None); - }; - return Ok(nodedb_types::json_msgpack::value_to_msgpack(&val).ok()); - } - let Some((start, end)) = extract_field(doc, 0, field) else { - return Ok(None); - }; - Ok(Some(doc[start..end].to_vec())) -} - -/// Check if msgpack bytes represent null. Msgpack null is the single byte 0xc0. -fn value_bytes_are_null(bytes: &[u8]) -> bool { - bytes == [0xc0] -} - -/// FNV-1a hash for raw bytes (used by approx aggregates to feed HLL/SpaceSaving). -fn hash_bytes(bytes: &[u8]) -> u64 { - let mut h: u64 = 0xcbf29ce484222325; - for &b in bytes { - h ^= b as u64; - h = h.wrapping_mul(0x100000001b3); - } - h -} - -#[cfg(test)] -mod tests { - use super::*; - use serde_json::json; - - fn encode(v: &serde_json::Value) -> Vec { - nodedb_types::json_msgpack::json_to_msgpack(v).expect("encode") - } - - #[test] - fn count() { - let d1 = encode(&json!({"x": 1})); - let d2 = encode(&json!({"x": 2})); - let d3 = encode(&json!({"x": 3})); - let docs: Vec<&[u8]> = vec![&d1, &d2, &d3]; - assert_eq!( - compute_aggregate_binary("count", "x", None, &docs).unwrap(), - Value::Integer(3) - ); - } - - #[test] - fn sum() { - let d1 = encode(&json!({"v": 10})); - let d2 = encode(&json!({"v": 20})); - let d3 = encode(&json!({"v": 30})); - let docs: Vec<&[u8]> = vec![&d1, &d2, &d3]; - assert_eq!( - compute_aggregate_binary("sum", "v", None, &docs).unwrap(), - Value::Float(60.0) - ); - } - - #[test] - fn avg() { - let d1 = encode(&json!({"v": 10})); - let d2 = encode(&json!({"v": 20})); - let docs: Vec<&[u8]> = vec![&d1, &d2]; - assert_eq!( - compute_aggregate_binary("avg", "v", None, &docs).unwrap(), - Value::Float(15.0) - ); - } - - #[test] - fn avg_empty() { - let d1 = encode(&json!({"other": 1})); - let docs: Vec<&[u8]> = vec![&d1]; - assert_eq!( - compute_aggregate_binary("avg", "v", None, &docs).unwrap(), - Value::Null - ); - } - - #[test] - fn min_max() { - let d1 = encode(&json!({"v": 5})); - let d2 = encode(&json!({"v": 1})); - let d3 = encode(&json!({"v": 9})); - let docs: Vec<&[u8]> = vec![&d1, &d2, &d3]; - - let min = compute_aggregate_binary("min", "v", None, &docs).unwrap(); - let max = compute_aggregate_binary("max", "v", None, &docs).unwrap(); - assert_eq!(min, Value::Integer(1)); - assert_eq!(max, Value::Integer(9)); - } - - #[test] - fn count_distinct() { - let d1 = encode(&json!({"v": "a"})); - let d2 = encode(&json!({"v": "b"})); - let d3 = encode(&json!({"v": "a"})); - let docs: Vec<&[u8]> = vec![&d1, &d2, &d3]; - assert_eq!( - compute_aggregate_binary("count_distinct", "v", None, &docs).unwrap(), - Value::Integer(2) - ); - } - - #[test] - fn string_agg() { - let d1 = encode(&json!({"n": "alice"})); - let d2 = encode(&json!({"n": "bob"})); - let docs: Vec<&[u8]> = vec![&d1, &d2]; - assert_eq!( - compute_aggregate_binary("string_agg", "n", None, &docs).unwrap(), - Value::String("alice,bob".into()) - ); - } - - #[test] - fn array_agg() { - let d1 = encode(&json!({"v": 1})); - let d2 = encode(&json!({"v": 2})); - let docs: Vec<&[u8]> = vec![&d1, &d2]; - let result = compute_aggregate_binary("array_agg", "v", None, &docs).unwrap(); - assert_eq!( - result, - Value::Array(vec![Value::Integer(1), Value::Integer(2),]) - ); - } - - #[test] - fn stddev_pop() { - let d1 = encode(&json!({"v": 2.0})); - let d2 = encode(&json!({"v": 4.0})); - let d3 = encode(&json!({"v": 4.0})); - let d4 = encode(&json!({"v": 4.0})); - let d5 = encode(&json!({"v": 5.0})); - let d6 = encode(&json!({"v": 5.0})); - let d7 = encode(&json!({"v": 7.0})); - let d8 = encode(&json!({"v": 9.0})); - let docs: Vec<&[u8]> = vec![&d1, &d2, &d3, &d4, &d5, &d6, &d7, &d8]; - let result = compute_aggregate_binary("stddev_pop", "v", None, &docs).unwrap(); - if let Value::Float(v) = result { - assert!((v - 2.0).abs() < 0.01); - } else { - panic!("expected Float"); - } - } - - #[test] - fn percentile_cont_median() { - let d1 = encode(&json!({"v": 1.0})); - let d2 = encode(&json!({"v": 2.0})); - let d3 = encode(&json!({"v": 3.0})); - let docs: Vec<&[u8]> = vec![&d1, &d2, &d3]; - assert_eq!( - compute_aggregate_binary("percentile_cont", "v", None, &docs).unwrap(), - Value::Float(2.0) - ); - } - - #[test] - fn missing_field_skipped() { - let d1 = encode(&json!({"v": 10})); - let d2 = encode(&json!({"other": 99})); - let d3 = encode(&json!({"v": 30})); - let docs: Vec<&[u8]> = vec![&d1, &d2, &d3]; - assert_eq!( - compute_aggregate_binary("sum", "v", None, &docs).unwrap(), - Value::Float(40.0) - ); - } - - #[test] - fn null_field_skipped_in_count_distinct() { - let d1 = encode(&json!({"v": "a"})); - let d2 = encode(&json!({"v": null})); - let d3 = encode(&json!({"v": "a"})); - let docs: Vec<&[u8]> = vec![&d1, &d2, &d3]; - assert_eq!( - compute_aggregate_binary("count_distinct", "v", None, &docs).unwrap(), - Value::Integer(1) - ); - } - - #[test] - fn array_agg_distinct() { - let d1 = encode(&json!({"v": 1})); - let d2 = encode(&json!({"v": 2})); - let d3 = encode(&json!({"v": 1})); - let docs: Vec<&[u8]> = vec![&d1, &d2, &d3]; - let result = compute_aggregate_binary("array_agg_distinct", "v", None, &docs).unwrap(); - assert_eq!( - result, - Value::Array(vec![Value::Integer(1), Value::Integer(2),]) - ); - } - - #[test] - fn sum_case_when_expression() { - let d1 = encode(&json!({"category": "tools"})); - let d2 = encode(&json!({"category": "books"})); - let d3 = encode(&json!({"category": "tools"})); - let docs: Vec<&[u8]> = vec![&d1, &d2, &d3]; - let expr = crate::expr::SqlExpr::Case { - operand: None, - when_thens: vec![( - crate::expr::SqlExpr::BinaryOp { - left: Box::new(crate::expr::SqlExpr::Column("category".into())), - op: crate::expr::BinaryOp::Eq, - right: Box::new(crate::expr::SqlExpr::Literal(Value::String("tools".into()))), - }, - crate::expr::SqlExpr::Literal(Value::Integer(1)), - )], - else_expr: Some(Box::new(crate::expr::SqlExpr::Literal(Value::Integer(0)))), - }; - - assert_eq!( - compute_aggregate_binary("sum", "*", Some(&expr), &docs).unwrap(), - Value::Float(2.0) - ); - } - - #[test] - fn approx_count_distinct_basic() { - let docs: Vec> = vec![ - encode(&json!({"region": "us"})), - encode(&json!({"region": "eu"})), - encode(&json!({"region": "us"})), - encode(&json!({"region": "ap"})), - ]; - let refs: Vec<&[u8]> = docs.iter().map(|d| d.as_slice()).collect(); - let result = - compute_aggregate_binary("approx_count_distinct", "region", None, &refs).unwrap(); - // HLL may not be exactly 3 but should be close. - if let Value::Integer(n) = result { - assert!((2..=4).contains(&n), "expected ~3 distinct, got {n}"); - } else { - panic!("expected Integer, got {result:?}"); - } - } - - #[test] - fn approx_percentile_basic() { - let docs: Vec> = (1..=100).map(|i| encode(&json!({"val": i}))).collect(); - let refs: Vec<&[u8]> = docs.iter().map(|d| d.as_slice()).collect(); - let result = compute_aggregate_binary("approx_percentile", "0.5:val", None, &refs).unwrap(); - if let Value::Float(f) = result { - assert!( - (f - 50.0).abs() < 10.0, - "p50 of 1..100 should be ~50, got {f}" - ); - } else { - panic!("expected Float, got {result:?}"); - } - } - - #[test] - fn approx_topk_basic() { - let mut docs: Vec> = Vec::new(); - for _ in 0..10 { - docs.push(encode(&json!({"cat": "a"}))); - } - for _ in 0..5 { - docs.push(encode(&json!({"cat": "b"}))); - } - for _ in 0..1 { - docs.push(encode(&json!({"cat": "c"}))); - } - let refs: Vec<&[u8]> = docs.iter().map(|d| d.as_slice()).collect(); - let result = compute_aggregate_binary("approx_topk", "3:cat", None, &refs).unwrap(); - if let Value::Array(arr) = result { - assert!(!arr.is_empty(), "should have top-k results"); - } else { - panic!("expected Array, got {result:?}"); - } - } - - #[test] - fn division_by_zero_in_expr_propagates() { - // `SUM(a / b)` over a document with `b = 0` must surface - // `EvalError::DivisionByZero` rather than folding the offending row - // to NULL and silently continuing the aggregate. - let d1 = encode(&json!({"a": 10, "b": 0})); - let docs: Vec<&[u8]> = vec![&d1]; - let expr = crate::expr::SqlExpr::BinaryOp { - left: Box::new(crate::expr::SqlExpr::Column("a".into())), - op: crate::expr::BinaryOp::Div, - right: Box::new(crate::expr::SqlExpr::Column("b".into())), - }; - let err = compute_aggregate_binary("sum", "*", Some(&expr), &docs).unwrap_err(); - assert_eq!(err, EvalError::DivisionByZero); - } -} diff --git a/nodedb-query/src/msgpack_scan/aggregate_helpers.rs b/nodedb-query/src/msgpack_scan/aggregate_helpers.rs index 3d82f37c6..904222530 100644 --- a/nodedb-query/src/msgpack_scan/aggregate_helpers.rs +++ b/nodedb-query/src/msgpack_scan/aggregate_helpers.rs @@ -3,7 +3,7 @@ //! Public helpers for streaming aggregate accumulators. //! //! These thin wrappers expose field-extraction primitives used by the -//! `handlers/aggregate.rs` streaming accumulator path in the `nodedb` crate. +//! streaming aggregate accumulators in the `nodedb` crate. //! Each function operates on a single raw MessagePack document byte slice and //! returns only the scalar value needed by the calling accumulator — no //! document bytes are retained after the call returns. @@ -12,7 +12,7 @@ use nodedb_types::Value; use crate::expr::{EvalError, SqlExpr}; use crate::msgpack_scan::field::extract_field; -use crate::msgpack_scan::reader::{read_f64, read_str, read_value}; +use crate::msgpack_scan::reader::{read_f64, read_numeric, read_str, read_value}; use crate::value_ops; // ── Expression evaluator ─────────────────────────────────────────────────── @@ -56,6 +56,32 @@ pub fn extract_f64( Ok(read_f64(doc, start)) } +/// Extract the SUM / AVG input of `field`, or of `expr` if provided, as the +/// number it contributes (`Integer`, `Float`, or `Decimal`). A raw field +/// and an expression result contribute by the same rule, +/// [`crate::numeric_sum::sum_input`]: a msgpack number read exactly, or a +/// numeric string read per [`crate::numeric_sum::numeric_text`] (a +/// `DECIMAL` cell is stored as its text). +/// Returns `Ok(None)` when nothing contributes, and +/// `Err(EvalError::DivisionByZero)` when `expr` divides/mods by zero. +#[inline] +pub fn extract_sum_value( + doc: &[u8], + field: &str, + expr: Option<&SqlExpr>, +) -> Result, EvalError> { + if let Some(expr) = expr { + return Ok(eval_expr(doc, expr)?.and_then(|v| crate::numeric_sum::sum_input(&v))); + } + let Some((start, _end)) = extract_field(doc, 0, field) else { + return Ok(None); + }; + Ok(match read_numeric(doc, start) { + Some(n) => Some(crate::numeric_sum::numeric_to_value(n)), + None => read_str(doc, start).and_then(crate::numeric_sum::numeric_text), + }) +} + /// Extract a display string from `field`, or evaluate `expr` if provided. /// Returns `Ok(None)` when the field is absent. pub fn extract_str( diff --git a/nodedb-query/src/msgpack_scan/compare.rs b/nodedb-query/src/msgpack_scan/compare.rs index b48122876..bea569e83 100644 --- a/nodedb-query/src/msgpack_scan/compare.rs +++ b/nodedb-query/src/msgpack_scan/compare.rs @@ -10,7 +10,8 @@ use std::hash::{BuildHasher, Hasher}; use nodedb_types::read_instant; -use crate::msgpack_scan::reader::{read_f64, read_i64, read_null, str_bounds}; +use crate::msgpack_scan::reader::{read_integer, read_null, read_numeric, str_bounds}; +use crate::numeric_cmp::cmp_numeric; /// Hash the raw bytes of a MessagePack value at `range` within `buf`. /// Uses a fast non-cryptographic hash suitable for hash joins and GROUP BY. @@ -49,7 +50,9 @@ pub fn hash_field_bytes_with( /// /// Comparison order: /// 1. Null < Bool < Number < Instant < String < Binary < Array < Map < Ext -/// 2. Within numbers: compare as f64 +/// 2. Within numbers: exact numeric order; integers (`uint64` included) +/// never round through f64, an integer against a float compares +/// exactly, and NaN sorts above every number and equals NaN /// 3. Within instants: by kind (UTC before naive), then signed epoch micros /// 4. Within strings: lexicographic on raw bytes (valid UTF-8 guarantees /// byte order = Unicode code-point order for ASCII/Latin-1) @@ -87,9 +90,10 @@ pub fn compare_field_bytes( a_val.cmp(&b_val) } RANK_NUMBER => { - // compare as f64 - match (read_f64(a_buf, a_off), read_f64(b_buf, b_off)) { - (Some(a), Some(b)) => a.partial_cmp(&b).unwrap_or(Ordering::Equal), + // Exact: integers never round through f64. NaN sorts above + // every number and equals NaN, so the order is total. + match (read_numeric(a_buf, a_off), read_numeric(b_buf, b_off)) { + (Some(a), Some(b)) => cmp_numeric(a, b), (Some(_), None) => Ordering::Greater, (None, Some(_)) => Ordering::Less, (None, None) => Ordering::Equal, @@ -128,10 +132,10 @@ pub fn compare_field_bytes( } } -/// Compare two numeric MessagePack values as i64. +/// Compare two integer MessagePack values exactly, `uint64` included. /// Useful when the caller knows both values are integers. pub fn compare_field_i64(a_buf: &[u8], a_off: usize, b_buf: &[u8], b_off: usize) -> Ordering { - match (read_i64(a_buf, a_off), read_i64(b_buf, b_off)) { + match (read_integer(a_buf, a_off), read_integer(b_buf, b_off)) { (Some(a), Some(b)) => a.cmp(&b), (Some(_), None) => Ordering::Greater, (None, Some(_)) => Ordering::Less, @@ -328,6 +332,83 @@ mod tests { ); } + fn cmp(a: &serde_json::Value, b: &serde_json::Value) -> Ordering { + let (a, b) = (encode(a), encode(b)); + compare_field_bytes(&a, val_range(&a), &b, val_range(&b)) + } + + #[test] + fn integers_past_two_pow_53_compare_exactly() { + let above = json!(9_007_199_254_740_993_i64); + let at = json!(9_007_199_254_740_992_i64); + assert_eq!(cmp(&above, &at), Ordering::Greater); + assert_eq!(cmp(&at, &above), Ordering::Less); + assert_eq!(cmp(&above, &above), Ordering::Equal); + } + + #[test] + fn nanosecond_timestamps_compare_exactly() { + let t1 = json!(1_700_000_000_000_000_001_i64); + let t2 = json!(1_700_000_000_000_000_002_i64); + assert_eq!(cmp(&t1, &t2), Ordering::Less); + assert_eq!(cmp(&t2, &t1), Ordering::Greater); + } + + #[test] + fn uint64_above_i64_max_orders_above_every_i64() { + assert_eq!(cmp(&json!(u64::MAX), &json!(i64::MAX)), Ordering::Greater); + assert_eq!(cmp(&json!(u64::MAX - 1), &json!(u64::MAX)), Ordering::Less); + assert_eq!(cmp(&json!(i64::MIN), &json!(u64::MAX)), Ordering::Less); + let (a, b) = (encode(&json!(u64::MAX)), encode(&json!(-1))); + assert_eq!(compare_field_i64(&a, 0, &b, 0), Ordering::Greater); + } + + #[test] + fn integer_against_float_compares_exactly() { + let above = json!(9_007_199_254_740_993_i64); + let float = json!(9_007_199_254_740_992.0_f64); + assert_eq!(cmp(&above, &float), Ordering::Greater); + assert_eq!(cmp(&float, &above), Ordering::Less); + assert_eq!(cmp(&json!(2), &json!(2.5)), Ordering::Less); + assert_eq!(cmp(&json!(3), &json!(3.0)), Ordering::Equal); + } + + /// NaN sorts above every number and equals NaN, so a sort over raw + /// cells is total. + #[test] + fn nan_cells_sort_above_every_number() { + let float = |f: f64| { + nodedb_types::value_to_msgpack(&nodedb_types::Value::Float(f)).expect("encode float") + }; + let nan = float(f64::NAN); + let inf = float(f64::INFINITY); + let max = encode(&json!(u64::MAX)); + assert_eq!( + compare_field_bytes(&nan, val_range(&nan), &inf, val_range(&inf)), + Ordering::Greater + ); + assert_eq!( + compare_field_bytes(&max, val_range(&max), &nan, val_range(&nan)), + Ordering::Less + ); + assert_eq!( + compare_field_bytes(&nan, val_range(&nan), &nan, val_range(&nan)), + Ordering::Equal + ); + let mut cells = [nan.clone(), encode(&json!(2)), float(-1.5), nan, max]; + cells.sort_by(|a, b| compare_field_bytes(a, val_range(a), b, val_range(b))); + assert_eq!(cells[0], float(-1.5)); + assert_eq!(cells[1], encode(&json!(2))); + assert_eq!(cells[2], encode(&json!(u64::MAX))); + } + + #[test] + fn numbers_keep_cross_rank_order() { + assert_eq!(cmp(&json!(true), &json!(u64::MAX)), Ordering::Less); + assert_eq!(cmp(&json!(u64::MAX), &json!("")), Ordering::Less); + assert_eq!(cmp(&json!(null), &json!(-1.5)), Ordering::Less); + } + #[test] fn compare_instants_negative_micros() { use nodedb_types::{InstantKind, write_instant}; diff --git a/nodedb-query/src/msgpack_scan/compare_decimal.rs b/nodedb-query/src/msgpack_scan/compare_decimal.rs new file mode 100644 index 000000000..1cea92300 --- /dev/null +++ b/nodedb-query/src/msgpack_scan/compare_decimal.rs @@ -0,0 +1,98 @@ +// SPDX-License-Identifier: Apache-2.0 + +//! Numeric comparison for a `DECIMAL` column's MessagePack cells. +//! +//! MessagePack has no decimal type. A `DECIMAL` cell is the string of its +//! canonical text, except an integral one above `i64::MAX`, which is a +//! `uint64`. One column therefore holds strings and integers side by side. +//! [`compare_field_bytes`] ranks numbers before strings and orders strings +//! by their bytes, so `5` sorts after `9223372036854775808` and `10` before +//! `9`. A caller that knows the column is `DECIMAL` compares through here. + +use std::cmp::Ordering; +use std::str::FromStr; + +use rust_decimal::Decimal; + +use super::compare::compare_field_bytes; +use super::reader::{read_f64, read_integer, read_str}; + +/// The decimal a MessagePack cell at `offset` stands for: a decimal string, +/// an integer, or a finite float. `None` for any other cell. +pub fn decimal_reading(buf: &[u8], offset: usize) -> Option { + if let Some(text) = read_str(buf, offset) { + return Decimal::from_str(text.trim()) + .or_else(|_| Decimal::from_scientific(text.trim())) + .ok(); + } + if let Some(integer) = read_integer(buf, offset) { + return Decimal::try_from_i128_with_scale(integer, 0).ok(); + } + read_f64(buf, offset).and_then(|float| Decimal::try_from(float).ok()) +} + +/// Order two `DECIMAL` cells by the numbers they stand for. +/// +/// A pair where either cell has no decimal reading falls back to +/// [`compare_field_bytes`], so a non-numeric cell keeps its usual rank. +pub fn compare_decimal_field_bytes( + a_buf: &[u8], + a_range: (usize, usize), + b_buf: &[u8], + b_range: (usize, usize), +) -> Ordering { + match ( + decimal_reading(a_buf, a_range.0), + decimal_reading(b_buf, b_range.0), + ) { + (Some(a), Some(b)) => a.cmp(&b), + _ => compare_field_bytes(a_buf, a_range, b_buf, b_range), + } +} + +#[cfg(test)] +mod tests { + use super::*; + + fn cell(value: &serde_json::Value) -> Vec { + nodedb_types::json_to_msgpack(value).expect("encode cell") + } + + fn cmp(a: &serde_json::Value, b: &serde_json::Value) -> Ordering { + let (a, b) = (cell(a), cell(b)); + compare_decimal_field_bytes(&a, (0, a.len()), &b, (0, b.len())) + } + + /// A small decimal stored as text orders below a `uint64` one. + #[test] + fn text_and_uint64_cells_order_numerically() { + let small = serde_json::json!("5"); + let mid = serde_json::json!(9_223_372_036_854_775_808_u64); + let max = serde_json::json!(u64::MAX); + assert_eq!(cmp(&small, &mid), Ordering::Less); + assert_eq!(cmp(&mid, &max), Ordering::Less); + assert_eq!(cmp(&max, &small), Ordering::Greater); + } + + /// Two text cells order by value, not by their bytes. + #[test] + fn text_cells_order_by_value() { + assert_eq!( + cmp(&serde_json::json!("10"), &serde_json::json!("9.5")), + Ordering::Greater + ); + assert_eq!( + cmp(&serde_json::json!("-2.50"), &serde_json::json!("-2.5")), + Ordering::Equal + ); + } + + /// A cell with no decimal reading keeps the generic rank. + #[test] + fn non_numeric_cell_falls_back() { + assert_eq!( + cmp(&serde_json::json!("abc"), &serde_json::json!("5")), + Ordering::Greater + ); + } +} diff --git a/nodedb-query/src/msgpack_scan/group_key.rs b/nodedb-query/src/msgpack_scan/group_key.rs index 3ed937692..2f1ef33fd 100644 --- a/nodedb-query/src/msgpack_scan/group_key.rs +++ b/nodedb-query/src/msgpack_scan/group_key.rs @@ -10,7 +10,7 @@ use nodedb_types::{NdbDateTime, read_instant}; use crate::expr::{EvalError, GroupKeySpec, SqlExpr}; use crate::msgpack_scan::field::extract_field; use crate::msgpack_scan::index::FieldIndex; -use crate::msgpack_scan::reader::{read_f64, read_i64, read_null, read_str}; +use crate::msgpack_scan::reader::{read_f64, read_integer, read_null, read_str}; /// Build a GROUP BY key string from raw msgpack bytes. /// @@ -113,7 +113,7 @@ fn append_value_at(buf: &mut String, doc: &[u8], start: usize, end: usize) { buf.push('"'); buf.push_str(&NdbDateTime::from_micros(micros).to_iso8601()); buf.push('"'); - } else if let Some(n) = read_i64(doc, start) { + } else if let Some(n) = read_integer(doc, start) { use std::fmt::Write; let _ = write!(buf, "{n}"); } else if let Some(n) = read_f64(doc, start) { @@ -191,6 +191,21 @@ mod tests { assert_eq!(key, "[200]"); } + /// A `uint64` above `i64::MAX` keys by its exact digits, and integers one + /// apart above 2^53 key apart. + #[test] + fn large_integers_key_exactly() { + let doc = encode(&json!({"v": u64::MAX})); + let key = build_group_key(&doc, &keys(&["v"])).unwrap(); + assert_eq!(key, "[18446744073709551615]"); + let a = encode(&json!({"v": 9_007_199_254_740_993_i64})); + let b = encode(&json!({"v": 9_007_199_254_740_992_i64})); + assert_ne!( + build_group_key(&a, &keys(&["v"])).unwrap(), + build_group_key(&b, &keys(&["v"])).unwrap() + ); + } + #[test] fn multiple_fields() { let doc = encode(&json!({"city": "ny", "year": 2024})); diff --git a/nodedb-query/src/msgpack_scan/mod.rs b/nodedb-query/src/msgpack_scan/mod.rs index 939c67a8d..c549ab0c8 100644 --- a/nodedb-query/src/msgpack_scan/mod.rs +++ b/nodedb-query/src/msgpack_scan/mod.rs @@ -6,9 +6,9 @@ //! `serde_json::Value` or `nodedb_types::Value`. Field extraction, numeric //! reads, comparisons, and hashing all work on raw byte offsets. -pub mod aggregate; pub mod aggregate_helpers; pub mod compare; +pub mod compare_decimal; pub mod field; pub mod filter; pub mod group_key; @@ -19,8 +19,8 @@ pub mod reader; pub mod sidecar; pub mod writer; -pub use aggregate::compute_aggregate_binary; pub use compare::{compare_field_bytes, hash_field_bytes}; +pub use compare_decimal::{compare_decimal_field_bytes, decimal_reading}; pub use field::{extract_field, extract_path}; pub use group_key::build_group_key; pub use index::FieldIndex; @@ -29,8 +29,8 @@ pub use kv_body::{ }; pub use kv_row::kv_row_msgpack; pub use reader::{ - array_header, map_header, read_bin_advance, read_bool, read_f64, read_i64, read_null, read_str, - read_str_advance, read_u32_advance, read_value, skip_value, + array_header, map_header, read_bin_advance, read_bool, read_f64, read_i64, read_integer, + read_null, read_str, read_str_advance, read_u32_advance, read_value, skip_value, }; pub use sidecar::{ SidecarEntry, SidecarFieldIndex, build_sidecar, field_index_from_sidecar, has_sidecar, diff --git a/nodedb-query/src/msgpack_scan/reader/mod.rs b/nodedb-query/src/msgpack_scan/reader/mod.rs index 1a19f8c94..a1cf22cc6 100644 --- a/nodedb-query/src/msgpack_scan/reader/mod.rs +++ b/nodedb-query/src/msgpack_scan/reader/mod.rs @@ -10,10 +10,10 @@ pub mod skip; pub mod tags; pub mod value; -pub(crate) use scalar::str_bounds; pub use scalar::{ - array_header, map_header, read_bin_advance, read_bool, read_f64, read_i64, read_null, read_str, - read_str_advance, read_u32_advance, + array_header, map_header, read_bin_advance, read_bool, read_f64, read_i64, read_integer, + read_null, read_str, read_str_advance, read_u32_advance, }; +pub(crate) use scalar::{read_numeric, str_bounds}; pub use skip::skip_value; pub use value::read_value; diff --git a/nodedb-query/src/msgpack_scan/reader/scalar.rs b/nodedb-query/src/msgpack_scan/reader/scalar.rs index 1067b7e04..20d959349 100644 --- a/nodedb-query/src/msgpack_scan/reader/scalar.rs +++ b/nodedb-query/src/msgpack_scan/reader/scalar.rs @@ -36,30 +36,44 @@ pub fn read_f64(buf: &[u8], offset: usize) -> Option { } /// Read an i64 from the value at `offset`. Handles all integer types. -/// Floats return `None` — use `read_f64` for those. +/// Floats return `None` — use `read_f64` for those. A `uint64` above +/// `i64::MAX` returns `None` — use `read_integer` for the exact value. pub fn read_i64(buf: &[u8], offset: usize) -> Option { + i64::try_from(read_integer(buf, offset)?).ok() +} + +/// Read any msgpack integer at `offset` exactly. `i128` holds every `i64` +/// and every `u64`. Floats and non-integers return `None`. +pub fn read_integer(buf: &[u8], offset: usize) -> Option { let tag = get(buf, offset)?; match tag { - 0x00..=0x7f => Some(tag as i64), - 0xe0..=0xff => Some((tag as i8) as i64), - UINT8 => Some(get(buf, offset + 1)? as i64), - UINT16 => Some(read_u16_be(buf, offset + 1)? as i64), - UINT32 => Some(read_u32_be(buf, offset + 1)? as i64), - UINT64 => { - let v = read_u64_be(buf, offset + 1)?; - Some(v as i64) - } - INT8 => Some(get(buf, offset + 1)? as i8 as i64), - INT16 => Some(read_u16_be(buf, offset + 1)? as i16 as i64), - INT32 => Some(read_u32_be(buf, offset + 1)? as i32 as i64), - INT64 => { - let v = read_u64_be(buf, offset + 1)?; - Some(v as i64) - } + 0x00..=0x7f => Some(i128::from(tag)), + 0xe0..=0xff => Some(i128::from(tag as i8)), + UINT8 => Some(i128::from(get(buf, offset + 1)?)), + UINT16 => Some(i128::from(read_u16_be(buf, offset + 1)?)), + UINT32 => Some(i128::from(read_u32_be(buf, offset + 1)?)), + UINT64 => Some(i128::from(read_u64_be(buf, offset + 1)?)), + INT8 => Some(i128::from(get(buf, offset + 1)? as i8)), + INT16 => Some(i128::from(read_u16_be(buf, offset + 1)? as i16)), + INT32 => Some(i128::from(read_u32_be(buf, offset + 1)? as i32)), + INT64 => Some(i128::from(read_u64_be(buf, offset + 1)? as i64)), _ => None, } } +/// Read a msgpack integer or float at `offset` without rounding an integer +/// through `f64`. +pub(crate) fn read_numeric(buf: &[u8], offset: usize) -> Option { + use crate::numeric_cmp::Numeric; + match read_integer(buf, offset) { + Some(i) => Some(Numeric::Int(i)), + None => match get(buf, offset)? { + FLOAT32 | FLOAT64 => read_f64(buf, offset).map(Numeric::Float), + _ => None, + }, + } +} + /// Read a string slice from the value at `offset`. Zero-copy — borrows /// directly from the input buffer. Returns `None` for non-string types /// or invalid UTF-8. @@ -244,6 +258,38 @@ mod tests { assert_eq!(read_i64(&buf, 0), Some(-500)); } + #[test] + fn uint64_above_i64_max_never_wraps() { + let mut buf = vec![UINT64]; + buf.extend_from_slice(&u64::MAX.to_be_bytes()); + assert_eq!(read_i64(&buf, 0), None); + assert_eq!(read_integer(&buf, 0), Some(i128::from(u64::MAX))); + + let mut edge = vec![UINT64]; + edge.extend_from_slice(&(i64::MAX as u64).to_be_bytes()); + assert_eq!(read_i64(&edge, 0), Some(i64::MAX)); + + let mut neg = vec![INT64]; + neg.extend_from_slice(&i64::MIN.to_be_bytes()); + assert_eq!(read_i64(&neg, 0), Some(i64::MIN)); + assert_eq!(read_integer(&neg, 0), Some(i128::from(i64::MIN))); + } + + #[test] + fn read_numeric_keeps_integers_exact() { + use crate::numeric_cmp::Numeric; + let buf = encode(&json!(9_007_199_254_740_993_i64)); + assert!(matches!( + read_numeric(&buf, 0), + Some(Numeric::Int(9_007_199_254_740_993)) + )); + let buf = encode(&json!(1.5)); + assert!(matches!(read_numeric(&buf, 0), Some(Numeric::Float(f)) if f == 1.5)); + let buf = encode(&json!("1")); + assert!(read_numeric(&buf, 0).is_none()); + assert!(read_integer(&buf, 0).is_none()); + } + #[test] fn read_str_fixstr() { let buf = encode(&json!("hi")); diff --git a/nodedb-query/src/msgpack_scan/reader/value.rs b/nodedb-query/src/msgpack_scan/reader/value.rs index 74b830d59..b91010c1f 100644 --- a/nodedb-query/src/msgpack_scan/reader/value.rs +++ b/nodedb-query/src/msgpack_scan/reader/value.rs @@ -29,9 +29,8 @@ pub fn read_value(buf: &[u8], offset: usize) -> Option { UINT32 => Some(nodedb_types::Value::Integer( read_u32_be(buf, offset + 1)? as i64 )), - UINT64 => Some(nodedb_types::Value::Integer( - read_u64_be(buf, offset + 1)? as i64 - )), + // Above `i64::MAX` this is a `Decimal`, never a wrapped negative. + UINT64 => Some(nodedb_types::Value::from_u64(read_u64_be(buf, offset + 1)?)), INT8 => Some(nodedb_types::Value::Integer( get(buf, offset + 1)? as i8 as i64 )), @@ -96,6 +95,26 @@ mod tests { assert_eq!(skip_value(&buf, 1), Some(11)); } + #[test] + fn read_value_uint64_above_i64_max_is_decimal() { + let mut buf = vec![UINT64]; + buf.extend_from_slice(&u64::MAX.to_be_bytes()); + assert_eq!( + read_value(&buf, 0), + Some(Value::Decimal(rust_decimal::Decimal::from(u64::MAX))) + ); + let mut edge = vec![UINT64]; + edge.extend_from_slice(&(i64::MAX as u64).to_be_bytes()); + assert_eq!(read_value(&edge, 0), Some(Value::Integer(i64::MAX))); + let buf = encode(&json!(9_223_372_036_854_775_808_u64)); + assert_eq!( + read_value(&buf, 0), + Some(Value::Decimal(rust_decimal::Decimal::from( + 9_223_372_036_854_775_808_u64 + ))) + ); + } + #[test] fn read_value_unknown_ext_is_none() { let buf = [FIXEXT8, 0x09, 0, 0, 0, 0, 0, 0, 0, 1]; diff --git a/nodedb-query/src/numeric_sum.rs b/nodedb-query/src/numeric_sum.rs new file mode 100644 index 000000000..64512d307 --- /dev/null +++ b/nodedb-query/src/numeric_sum.rs @@ -0,0 +1,669 @@ +// SPDX-License-Identifier: Apache-2.0 + +//! Exact SUM / AVG accumulation, shared by every aggregate path. +//! +//! The rule, applied the same way everywhere: +//! +//! - An integer input adds exactly into an `i128`. A `u64` above `i64::MAX` +//! and an integral `Decimal` are integer inputs. +//! - A fractional `Decimal` adds exactly into a `Decimal`. +//! - A float input adds into a Neumaier-compensated `f64`. Once the float +//! total is not finite, compensation stops and the total stays as IEEE +//! arithmetic gives it: an overflow is `Infinity`, `Infinity` plus +//! `-Infinity` is NaN, as in PostgreSQL. +//! - A numeric string is the number it spells: an integer, or a `Decimal` +//! when it has a fraction. A `DECIMAL` cell is stored as its text, so this +//! keeps a `DECIMAL` column exact. +//! - SUM with only integer inputs is an `Integer` when the total fits `i64`. +//! A larger total is a `Decimal`. +//! - SUM with a fractional `Decimal` input and no float input is a +//! `Decimal`: the exact integer part plus the exact decimal part. +//! - An exact total past the `Decimal` range is +//! [`EvalError::NumericOverflow`]. +//! - SUM with at least one float input is a `Float`: the exact parts plus +//! the float part. +//! - AVG is a `Float`. With no float input it divides the exact total by the +//! count, with no rounding of the total first. +//! - SUM and AVG over no input are NULL. +//! - Two partial accumulators merge without loss: the integer parts add +//! exactly, so a shard or spill merge gives the single-pass result. +//! +//! STDDEV and VARIANCE stay in `f64` (Welford); they are not exact sums. + +use nodedb_types::Value; +use rust_decimal::Decimal; +use rust_decimal::prelude::ToPrimitive; + +use crate::expr::EvalError; +use crate::numeric_cmp::{Numeric, parse_numeric_str}; + +/// Running exact SUM / AVG state. See the module docs for the rule. +#[derive(Debug, Clone, Default, PartialEq)] +pub struct ExactSum { + /// Exact total of the integer inputs. + int: i128, + /// Exact total of the fractional `Decimal` inputs. + frac: Decimal, + /// At least one fractional `Decimal` input was added. + has_decimal: bool, + /// An exact total (`int` or `frac`) left its range. + exact_overflow: bool, + /// Running total of the float inputs. + float: f64, + /// Neumaier compensation term: the rounding error lost from `float`. + /// The float total is `float + comp` while `float` is finite. + comp: f64, + /// At least one float input was added. + has_float: bool, + /// Number of inputs added. + count: u64, +} + +impl ExactSum { + pub fn new() -> Self { + Self::default() + } + + /// Number of inputs added. + pub fn count(&self) -> u64 { + self.count + } + + /// Add one integer input. + pub fn add_i64(&mut self, v: i64) { + self.add_int(i128::from(v)); + } + + /// Add one unsigned integer input. + pub fn add_u64(&mut self, v: u64) { + self.add_int(i128::from(v)); + } + + /// Add one float input. + pub fn add_f64(&mut self, v: f64) { + self.float_add(v); + self.has_float = true; + self.count += 1; + } + + /// Add a numeric `Value`: `Integer`, `Float`, or `Decimal`. Any other + /// value adds nothing. Returns whether the value was added. + pub fn add_value(&mut self, v: &Value) -> bool { + match v { + Value::Integer(i) => self.add_i64(*i), + Value::Float(f) => self.add_f64(*f), + Value::Decimal(d) => self.add_decimal(*d), + _ => return false, + } + true + } + + /// Add one `Decimal` input: an integral one into the integer total, any + /// other into the exact decimal total. + pub fn add_decimal(&mut self, d: Decimal) { + if d.fract().is_zero() + && let Some(i) = d.to_i128() + { + self.add_int(i); + return; + } + match self.frac.checked_add(d) { + Some(total) => self.frac = total, + None => self.exact_overflow = true, + } + self.has_decimal = true; + self.count += 1; + } + + /// Fold another partial accumulator into this one without loss. + pub fn merge(&mut self, other: &ExactSum) { + match self.int.checked_add(other.int) { + Some(total) => self.int = total, + None => self.exact_overflow = true, + } + match self.frac.checked_add(other.frac) { + Some(total) => self.frac = total, + None => self.exact_overflow = true, + } + self.exact_overflow |= other.exact_overflow; + self.has_decimal |= other.has_decimal; + self.float_add(other.float); + if self.float.is_finite() { + self.comp += other.comp; + } + self.has_float |= other.has_float; + self.count += other.count; + } + + /// SUM of the inputs. NULL for no input. + pub fn sum(&self) -> Result { + if self.count == 0 { + return Ok(Value::Null); + } + if self.exact_overflow { + return Err(EvalError::NumericOverflow { function: "sum" }); + } + if self.has_float { + return Ok(Value::Float(self.sum_f64())); + } + if self.has_decimal { + return self + .exact_decimal_total() + .map(Value::Decimal) + .ok_or(EvalError::NumericOverflow { function: "sum" }); + } + if let Ok(i) = i64::try_from(self.int) { + return Ok(Value::Integer(i)); + } + Decimal::try_from_i128_with_scale(self.int, 0) + .map(Value::Decimal) + .map_err(|_| EvalError::NumericOverflow { function: "sum" }) + } + + /// The total read as an `f64`, for a float-typed result. Rounds an + /// exact total above 2^53; [`Self::sum`] keeps it exact. `0.0` for no + /// input. A float total that is not finite (`Infinity`, `-Infinity`, + /// NaN) is the result alone. + pub fn sum_f64(&self) -> f64 { + if !self.float.is_finite() { + return self.float; + } + self.int as f64 + self.frac.as_f64() + (self.float + self.comp) + } + + /// The integer and decimal totals as one `Decimal`. `None` past the + /// `Decimal` range. + fn exact_decimal_total(&self) -> Option { + Decimal::try_from_i128_with_scale(self.int, 0) + .ok()? + .checked_add(self.frac) + } + + /// AVG of the inputs as a `Float`. NULL for no input. + pub fn avg(&self) -> Result { + match self.avg_f64()? { + Some(avg) => Ok(Value::Float(avg)), + None => Ok(Value::Null), + } + } + + /// AVG of the inputs as an `f64`. `None` for no input. + pub fn avg_f64(&self) -> Result, EvalError> { + if self.count == 0 { + return Ok(None); + } + if self.exact_overflow { + return Err(EvalError::NumericOverflow { function: "avg" }); + } + let n = i128::from(self.count); + let avg = if self.has_float { + self.sum_f64() / self.count as f64 + } else if self.has_decimal { + // The exact total divides as a `Decimal`, so only the quotient + // rounds. + let total = self + .exact_decimal_total() + .ok_or(EvalError::NumericOverflow { function: "avg" })?; + total + .checked_div(Decimal::from(self.count)) + .and_then(|avg| avg.to_f64()) + .ok_or(EvalError::NumericOverflow { function: "avg" })? + } else { + // Quotient and remainder keep the exact total: no rounding of a + // total above 2^53 before the division. + (self.int / n) as f64 + (self.int % n) as f64 / self.count as f64 + }; + Ok(Some(avg)) + } + + fn add_int(&mut self, v: i128) { + match self.int.checked_add(v) { + Some(total) => self.int = total, + None => self.exact_overflow = true, + } + self.count += 1; + } + + /// Neumaier summation step. A total that is not finite takes no + /// compensation: `Infinity` and NaN propagate as IEEE arithmetic gives + /// them. + fn float_add(&mut self, v: f64) { + let t = self.float + v; + if t.is_finite() { + if self.float.abs() >= v.abs() { + self.comp += (self.float - t) + v; + } else { + self.comp += (v - t) + self.float; + } + } + self.float = t; + } +} + +/// A SUM / AVG argument value as the number it contributes: a number as +/// itself, a numeric string as the number it spells (see [`numeric_text`]). +/// `None` for any other value, which contributes nothing. +pub fn sum_input(v: &Value) -> Option { + match v { + Value::Integer(_) | Value::Float(_) | Value::Decimal(_) => Some(v.clone()), + Value::String(s) => numeric_text(s), + _ => None, + } +} + +/// The number a numeric string spells, exactly where it can be: an integer +/// as [`numeric_to_value`] gives it, a fraction in decimal notation as a +/// `Decimal`, and any other numeric text (an exponent, a fraction past the +/// `Decimal` precision) as a `Float`. `None` for non-numeric text. +pub fn numeric_text(s: &str) -> Option { + parse_numeric_str(s).map(numeric_to_value) +} + +/// A numeric reading as a `Value`. An integer past `i64` is a `Decimal`. +/// An integer past the `Decimal` range is the nearest `Float`, the same +/// reading `f64` parsing gives it. +pub(crate) fn numeric_to_value(n: Numeric) -> Value { + match n { + Numeric::Int(i) => match i64::try_from(i) { + Ok(small) => Value::Integer(small), + Err(_) => Decimal::try_from_i128_with_scale(i, 0) + .map(Value::Decimal) + .unwrap_or(Value::Float(i as f64)), + }, + Numeric::Decimal(d) => Value::Decimal(d), + Numeric::Float(f) => Value::Float(f), + } +} + +/// Element count of the MessagePack form: the array `[int_hi, int_lo, frac, +/// has_decimal, exact_overflow, float, comp, has_float, count]`. The `i128` +/// total is split into its high `i64` and low `u64` halves, and `frac` is +/// the 16-byte `Decimal` serialization, so both round-trip exactly. +const MSGPACK_FIELDS: usize = 9; + +impl zerompk::ToMessagePack for ExactSum { + fn write(&self, writer: &mut W) -> zerompk::Result<()> { + writer.write_array_len(MSGPACK_FIELDS)?; + writer.write_i64((self.int >> 64) as i64)?; + writer.write_u64(self.int as u64)?; + writer.write_binary(&self.frac.serialize())?; + writer.write_boolean(self.has_decimal)?; + writer.write_boolean(self.exact_overflow)?; + writer.write_f64(self.float)?; + writer.write_f64(self.comp)?; + writer.write_boolean(self.has_float)?; + writer.write_u64(self.count) + } +} + +impl<'a> zerompk::FromMessagePack<'a> for ExactSum { + fn read>(reader: &mut R) -> zerompk::Result { + reader.check_array_len(MSGPACK_FIELDS)?; + let hi = reader.read_i64()?; + let lo = reader.read_u64()?; + let frac_bytes = reader.read_binary()?; + let frac_bytes: [u8; 16] = frac_bytes + .as_ref() + .try_into() + .map_err(|_| zerompk::Error::BufferTooSmall)?; + Ok(Self { + int: (i128::from(hi) << 64) | i128::from(lo), + frac: Decimal::deserialize(frac_bytes), + has_decimal: reader.read_boolean()?, + exact_overflow: reader.read_boolean()?, + float: reader.read_f64()?, + comp: reader.read_f64()?, + has_float: reader.read_boolean()?, + count: reader.read_u64()?, + }) + } +} + +#[cfg(test)] +mod tests { + use super::*; + + const ABOVE: i64 = 9_007_199_254_740_993; + const AT: i64 = 9_007_199_254_740_992; + + fn sum_of(values: &[Value]) -> Result { + let mut acc = ExactSum::new(); + for v in values { + acc.add_value(v); + } + acc.sum() + } + + #[test] + fn integers_past_two_pow_53_sum_exactly() { + let total = sum_of(&[Value::Integer(ABOVE), Value::Integer(AT)]).unwrap(); + assert_eq!(total, Value::Integer(18_014_398_509_481_985)); + } + + #[test] + fn nanosecond_timestamps_sum_exactly() { + let t = [ + Value::Integer(1_700_000_000_000_000_001), + Value::Integer(1_700_000_000_000_000_002), + ]; + assert_eq!( + sum_of(&t).unwrap(), + Value::Integer(3_400_000_000_000_000_003) + ); + } + + #[test] + fn total_past_i64_is_decimal() { + let total = sum_of(&[Value::Integer(i64::MAX), Value::Integer(i64::MAX)]).unwrap(); + assert_eq!( + total, + Value::Decimal(Decimal::from_i128_with_scale(2 * i128::from(i64::MAX), 0)) + ); + } + + #[test] + fn u64_above_i64_max_sums_exactly() { + let mut acc = ExactSum::new(); + acc.add_u64(u64::MAX); + acc.add_value(&Value::from_u64(u64::MAX)); + assert_eq!( + acc.sum().unwrap(), + Value::Decimal(Decimal::from_i128_with_scale(2 * i128::from(u64::MAX), 0)) + ); + let mut small = ExactSum::new(); + small.add_u64(u64::MAX); + small.add_i64(-1); + small.add_i64(i64::MIN); + assert_eq!(small.sum().unwrap(), Value::Integer(i64::MAX - 1)); + } + + #[test] + fn total_past_decimal_range_is_overflow_error() { + let mut acc = ExactSum::new(); + // Four times 2^94 is 2^96, past the 96-bit `Decimal` mantissa. + for _ in 0..4 { + acc.add_int(1i128 << 94); + } + assert_eq!( + acc.sum(), + Err(EvalError::NumericOverflow { function: "sum" }) + ); + } + + #[test] + fn i128_overflow_is_overflow_error() { + let mut acc = ExactSum::new(); + acc.add_int(i128::MAX); + acc.add_int(1); + assert_eq!( + acc.sum(), + Err(EvalError::NumericOverflow { function: "sum" }) + ); + assert_eq!( + acc.avg(), + Err(EvalError::NumericOverflow { function: "avg" }) + ); + } + + #[test] + fn mixed_int_float_is_float() { + let total = sum_of(&[Value::Integer(2), Value::Float(0.5)]).unwrap(); + assert_eq!(total, Value::Float(2.5)); + let total = sum_of(&[Value::Integer(ABOVE), Value::Float(1.0)]).unwrap(); + assert_eq!(total, Value::Float(ABOVE as f64 + 1.0)); + } + + #[test] + fn nan_input_makes_float_nan() { + let total = sum_of(&[Value::Integer(1), Value::Float(f64::NAN)]).unwrap(); + assert!(matches!(total, Value::Float(f) if f.is_nan())); + } + + #[test] + fn empty_is_null() { + let acc = ExactSum::new(); + assert_eq!(acc.sum().unwrap(), Value::Null); + assert_eq!(acc.avg().unwrap(), Value::Null); + assert_eq!(sum_of(&[Value::String("x".into())]).unwrap(), Value::Null); + } + + #[test] + fn avg_uses_the_exact_total() { + let mut acc = ExactSum::new(); + acc.add_i64(ABOVE); + acc.add_i64(AT); + // Exact mean is 2^53 + 0.5; the f64 nearest is 2^53. + assert_eq!(acc.avg().unwrap(), Value::Float(AT as f64)); + let mut big = ExactSum::new(); + big.add_i64(i64::MAX); + big.add_i64(i64::MAX); + assert_eq!(big.avg().unwrap(), Value::Float(i64::MAX as f64)); + let mut ints = ExactSum::new(); + ints.add_i64(10); + ints.add_i64(20); + ints.add_i64(25); + assert_eq!(ints.avg_f64().unwrap(), Some(55.0 / 3.0)); + } + + #[test] + fn partial_merge_keeps_exactness() { + let mut a = ExactSum::new(); + a.add_i64(ABOVE); + let mut b = ExactSum::new(); + b.add_i64(AT); + b.add_i64(i64::MAX); + a.merge(&b); + let mut single = ExactSum::new(); + for v in [ABOVE, AT, i64::MAX] { + single.add_i64(v); + } + assert_eq!(a.sum().unwrap(), single.sum().unwrap()); + assert_eq!(a.count(), 3); + } + + fn float_sum(values: &[f64]) -> f64 { + let mut acc = ExactSum::new(); + for v in values { + acc.add_f64(*v); + } + match acc.sum().unwrap() { + Value::Float(f) => f, + other => panic!("expected a float SUM, got {other:?}"), + } + } + + /// Float overflow is `Infinity` and `Infinity` plus `-Infinity` is NaN, + /// as in PostgreSQL. + #[test] + fn non_finite_float_totals_follow_ieee() { + assert_eq!(float_sum(&[1e308, 1e308]), f64::INFINITY); + assert_eq!(float_sum(&[-1e308, -1e308]), f64::NEG_INFINITY); + assert_eq!(float_sum(&[f64::INFINITY, 1.0]), f64::INFINITY); + assert_eq!(float_sum(&[1.0, f64::NEG_INFINITY]), f64::NEG_INFINITY); + assert!(float_sum(&[f64::INFINITY, f64::NEG_INFINITY]).is_nan()); + assert!(float_sum(&[f64::NAN, f64::INFINITY]).is_nan()); + let mut avg = ExactSum::new(); + avg.add_f64(1e308); + avg.add_f64(1e308); + assert_eq!(avg.avg().unwrap(), Value::Float(f64::INFINITY)); + let mut mixed = ExactSum::new(); + mixed.add_i64(5); + mixed.add_f64(f64::INFINITY); + assert_eq!(mixed.sum().unwrap(), Value::Float(f64::INFINITY)); + } + + /// Neumaier compensation keeps a small term that a large pair cancels. + #[test] + fn compensation_survives_large_cancelling_terms() { + assert_eq!(float_sum(&[1.0, 1e100, 1.0, -1e100]), 2.0); + assert_eq!(float_sum(&[0.1; 10]), 1.0); + } + + #[test] + fn merge_keeps_non_finite_totals() { + let mut a = ExactSum::new(); + a.add_f64(1e308); + let mut b = ExactSum::new(); + b.add_f64(1e308); + a.merge(&b); + assert_eq!(a.sum().unwrap(), Value::Float(f64::INFINITY)); + let mut pos = ExactSum::new(); + pos.add_f64(f64::INFINITY); + let mut neg = ExactSum::new(); + neg.add_f64(f64::NEG_INFINITY); + pos.merge(&neg); + assert!(matches!(pos.sum().unwrap(), Value::Float(f) if f.is_nan())); + let mut small = ExactSum::new(); + small.add_f64(1.0); + let mut big = ExactSum::new(); + big.add_f64(1e100); + big.add_f64(1.0); + big.add_f64(-1e100); + small.merge(&big); + assert_eq!(small.sum().unwrap(), Value::Float(2.0)); + } + + #[test] + fn msgpack_round_trip_keeps_non_finite_floats() { + for v in [f64::NAN, f64::INFINITY, f64::NEG_INFINITY] { + let mut acc = ExactSum::new(); + acc.add_f64(v); + let bytes = zerompk::to_msgpack_vec(&acc).unwrap(); + let back: ExactSum = zerompk::from_msgpack(&bytes).unwrap(); + assert_eq!(back.count(), 1); + match (back.sum().unwrap(), acc.sum().unwrap()) { + (Value::Float(x), Value::Float(y)) => { + assert!(x == y || (x.is_nan() && y.is_nan())); + } + other => panic!("expected float sums, got {other:?}"), + } + } + } + + #[test] + fn msgpack_round_trip_keeps_every_part() { + let mut big = ExactSum::new(); + big.add_u64(u64::MAX); + big.add_u64(u64::MAX); + let mut negative = ExactSum::new(); + negative.add_i64(i64::MIN); + negative.add_i64(i64::MIN); + negative.add_f64(0.1); + let mut overflow = ExactSum::new(); + overflow.add_int(i128::MAX); + overflow.add_int(1); + for acc in [ExactSum::new(), big, negative, overflow] { + let bytes = zerompk::to_msgpack_vec(&acc).unwrap(); + let back: ExactSum = zerompk::from_msgpack(&bytes).unwrap(); + assert_eq!(back, acc); + assert_eq!(back.sum(), acc.sum()); + } + } + + #[test] + fn coerced_strings_stay_exact() { + let mut acc = ExactSum::new(); + let text = sum_input(&Value::String("9007199254740993".into())).unwrap(); + assert_eq!(text, Value::Integer(ABOVE)); + assert!(acc.add_value(&text)); + assert!(acc.add_value(&Value::Integer(AT))); + assert!(!acc.add_value(&Value::String("1".into()))); + assert_eq!(acc.sum().unwrap(), Value::Integer(18_014_398_509_481_985)); + assert_eq!( + sum_input(&Value::String("18446744073709551615".into())), + Some(Value::Decimal(Decimal::from(u64::MAX))) + ); + assert_eq!(sum_input(&Value::String("x".into())), None); + assert_eq!(sum_input(&Value::Bool(true)), None); + } + + fn dec(text: &str) -> Decimal { + Decimal::from_str_exact(text).unwrap() + } + + #[test] + fn fractional_decimals_sum_exactly() { + let tenths = vec![Value::Decimal(dec("0.1")); 3]; + assert_eq!(sum_of(&tenths).unwrap(), Value::Decimal(dec("0.3"))); + let mixed = [Value::Decimal(dec("0.25")), Value::Integer(i64::MAX)]; + assert_eq!( + sum_of(&mixed).unwrap(), + Value::Decimal(Decimal::from(i64::MAX) + dec("0.25")) + ); + let with_float = [Value::Decimal(dec("0.5")), Value::Float(0.25)]; + assert_eq!(sum_of(&with_float).unwrap(), Value::Float(0.75)); + } + + #[test] + fn decimal_avg_divides_the_exact_total() { + let mut acc = ExactSum::new(); + acc.add_value(&Value::Decimal(dec("0.1"))); + acc.add_value(&Value::Decimal(dec("0.2"))); + assert_eq!(acc.avg().unwrap(), Value::Float(0.15)); + } + + #[test] + fn a_decimal_total_past_the_range_is_overflow_error() { + let mut acc = ExactSum::new(); + acc.add_value(&Value::Decimal(Decimal::MAX - dec("0.5"))); + acc.add_value(&Value::Decimal(Decimal::MAX - dec("0.5"))); + assert_eq!( + acc.sum(), + Err(EvalError::NumericOverflow { function: "sum" }) + ); + let mut joined = ExactSum::new(); + joined.add_value(&Value::Decimal(Decimal::MAX)); + joined.add_value(&Value::Decimal(dec("0.5"))); + assert_eq!( + joined.sum(), + Err(EvalError::NumericOverflow { function: "sum" }) + ); + } + + #[test] + fn numeric_text_is_exact_where_it_can_be() { + assert_eq!(numeric_text("12"), Some(Value::Integer(12))); + assert_eq!(numeric_text("0.1"), Some(Value::Decimal(dec("0.1")))); + assert_eq!(numeric_text("1e3"), Some(Value::Float(1000.0))); + assert_eq!(numeric_text("x"), None); + let mut acc = ExactSum::new(); + for text in ["0.1", "0.2"] { + assert!(acc.add_value(&sum_input(&Value::String(text.into())).unwrap())); + } + assert_eq!(acc.sum().unwrap(), Value::Decimal(dec("0.3"))); + } + + #[test] + fn decimal_parts_merge_and_round_trip() { + let mut a = ExactSum::new(); + a.add_value(&Value::Decimal(dec("0.1"))); + let mut b = ExactSum::new(); + b.add_value(&Value::Decimal(dec("0.2"))); + b.add_i64(1); + a.merge(&b); + assert_eq!(a.sum().unwrap(), Value::Decimal(dec("1.3"))); + let bytes = zerompk::to_msgpack_vec(&a).unwrap(); + let back: ExactSum = zerompk::from_msgpack(&bytes).unwrap(); + assert_eq!(back, a); + } + + #[test] + fn sum_inputs_sum_exactly() { + let mut acc = ExactSum::new(); + for v in [ + Value::Integer(ABOVE), + Value::from_u64(u64::MAX), + Value::String("2".into()), + ] { + assert!(acc.add_value(&sum_input(&v).unwrap())); + } + assert_eq!(sum_input(&Value::Null), None); + assert_eq!( + acc.sum().unwrap(), + Value::Decimal(Decimal::from_i128_with_scale( + i128::from(ABOVE) + i128::from(u64::MAX) + 2, + 0 + )) + ); + } +} diff --git a/nodedb-query/src/ts_functions/percentile.rs b/nodedb-query/src/ts_functions/percentile.rs index c16c954aa..7874ca670 100644 --- a/nodedb-query/src/ts_functions/percentile.rs +++ b/nodedb-query/src/ts_functions/percentile.rs @@ -12,7 +12,7 @@ pub fn ts_percentile_exact(values: &[f64], p: f64) -> Option { if sorted.is_empty() { return None; } - sorted.sort_unstable_by(|a, b| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal)); + sorted.sort_unstable_by(|a, b| crate::numeric_cmp::cmp_f64(*a, *b)); let n = sorted.len(); if n == 1 { diff --git a/nodedb-query/src/value_ops.rs b/nodedb-query/src/value_ops.rs index 4d2226e58..c91d65c09 100644 --- a/nodedb-query/src/value_ops.rs +++ b/nodedb-query/src/value_ops.rs @@ -10,6 +10,8 @@ use std::cmp::Ordering; use nodedb_types::Value; +use crate::numeric_cmp::{Numeric, cmp_numeric, decimal_reading, numeric_eq, parse_numeric_str}; + /// Coerce a Value to f64. /// /// - Integer/Float: direct conversion @@ -33,23 +35,45 @@ pub fn value_to_f64(v: &Value, coerce_bool: bool) -> Option { /// Whether either side is a typed instant, so the pair compares by epoch /// microseconds through `Value::partial_cmp_coerced` (an ISO string on the /// other side is parsed). -fn involves_instant(a: &Value, b: &Value) -> bool { +pub fn involves_instant(a: &Value, b: &Value) -> bool { a.as_instant().is_some() || b.as_instant().is_some() } +/// `v` read as a number: integers, floats, decimals, numeric strings, and +/// bools (`true` = 1). Integers, decimals, integer text and fractional +/// decimal text stay exact. +fn numeric_reading(v: &Value) -> Option { + match v { + Value::Integer(i) => Some(Numeric::Int(i128::from(*i))), + Value::Bool(b) => Some(Numeric::Int(i128::from(*b))), + Value::String(s) => parse_numeric_str(s), + Value::Float(f) => Some(Numeric::Float(*f)), + Value::Decimal(d) => Some(decimal_reading(d)), + _ => None, + } +} + +/// Numeric order of `a` and `b` when both have a numeric reading (number, +/// decimal, numeric string, bool). `None` when a side is not numeric. The +/// order is total: NaN sorts above every number and equals NaN. +pub fn numeric_order(a: &Value, b: &Value) -> Option { + let (na, nb) = (numeric_reading(a)?, numeric_reading(b)?); + Some(cmp_numeric(na, nb)) +} + /// Compare two Values with type coercion, for a predicate. /// /// A typed instant compares by epoch microseconds against another instant /// or an ISO-8601 string, and against anything else has no order: `None`, /// so `WHERE at >= 5` over an instant matches nothing. Otherwise numeric -/// comparison first (with bool coercion; `None` for a NaN), then string -/// comparison of the display forms. +/// comparison first (with bool coercion, exact for integers and decimals, +/// NaN above every number), then string comparison of the display forms. pub fn partial_compare_values(a: &Value, b: &Value) -> Option { if involves_instant(a, b) { return a.partial_cmp_coerced(b); } - if let (Some(na), Some(nb)) = (value_to_f64(a, true), value_to_f64(b, true)) { - return na.partial_cmp(&nb); + if let Some(order) = numeric_order(a, b) { + return Some(order); } let sa = value_to_display_string(a); let sb = value_to_display_string(b); @@ -66,11 +90,36 @@ pub fn compare_values(a: &Value, b: &Value) -> Ordering { partial_compare_values(a, b).unwrap_or(Ordering::Equal) } +/// Order of two non-NULL values for an ORDER BY over output rows (window +/// ORDER BY, post-aggregate ORDER BY). Text orders by its bytes, so `"10"` +/// sorts before `"9"`. Booleans sort `false` first. Any other pair takes +/// [`compare_values`]: numbers by their exact reading, NaN above every +/// number. The caller places NULLs. +pub fn compare_sort_values(a: &Value, b: &Value) -> Ordering { + match (a, b) { + (Value::String(x), Value::String(y)) => x.cmp(y), + (Value::Bool(x), Value::Bool(y)) => x.cmp(y), + _ => compare_values(a, b), + } +} + +/// Whether two values are peers under an ORDER BY: both NULL, or both +/// non-NULL and ordered `Equal` by [`compare_sort_values`]. NaN is a peer +/// of NaN. +pub fn sort_peers(a: &Value, b: &Value) -> bool { + match (a.is_null(), b.is_null()) { + (true, true) => true, + (false, false) => compare_sort_values(a, b) == Ordering::Equal, + _ => false, + } +} + /// Check equality with type coercion. /// /// A typed instant equals another instant or an ISO-8601 string with the -/// same epoch microseconds. Handles `"5" == 5` by coercing both sides to -/// f64 when one is a number and the other is a numeric string. +/// same epoch microseconds. Handles `"5" == 5` by reading both sides as +/// numbers when each has a numeric reading. Every numeric pair compares +/// exactly, and NaN equals NaN. pub fn coerced_eq(a: &Value, b: &Value) -> bool { if a == b { return true; @@ -78,10 +127,10 @@ pub fn coerced_eq(a: &Value, b: &Value) -> bool { if involves_instant(a, b) { return a.eq_coerced(b); } - if let (Some(af), Some(bf)) = (value_to_f64(a, true), value_to_f64(b, true)) { - return (af - bf).abs() < f64::EPSILON; + match (numeric_reading(a), numeric_reading(b)) { + (Some(na), Some(nb)) => numeric_eq(na, nb), + _ => false, } - false } /// Check if a Value is truthy (for boolean contexts). @@ -165,6 +214,168 @@ mod tests { ); } + /// `2^53 + 1` and `2^53` collapse to one `f64`. Integers compare exactly. + #[test] + fn integers_past_two_pow_53_compare_exactly() { + let above = Value::Integer(9_007_199_254_740_993); + let at = Value::Integer(9_007_199_254_740_992); + assert!(!coerced_eq(&above, &at)); + assert_eq!(partial_compare_values(&above, &at), Some(Ordering::Greater)); + assert_eq!(partial_compare_values(&at, &above), Some(Ordering::Less)); + // Nanosecond timestamps one tick apart. + let t0 = Value::Integer(1_700_000_000_000_000_001); + let t1 = Value::Integer(1_700_000_000_000_000_002); + assert!(!coerced_eq(&t0, &t1)); + assert_eq!(partial_compare_values(&t0, &t1), Some(Ordering::Less)); + assert_eq!( + partial_compare_values(&Value::Integer(i64::MAX), &Value::Integer(i64::MAX - 1)), + Some(Ordering::Greater) + ); + } + + #[test] + fn integer_against_float_compares_without_rounding() { + let above = Value::Integer(9_007_199_254_740_993); + let float = Value::Float(9_007_199_254_740_992.0); + assert!(!coerced_eq(&above, &float)); + assert_eq!( + partial_compare_values(&above, &float), + Some(Ordering::Greater) + ); + assert_eq!(partial_compare_values(&float, &above), Some(Ordering::Less)); + assert_eq!( + partial_compare_values(&Value::Integer(2), &Value::Float(2.5)), + Some(Ordering::Less) + ); + assert_eq!( + partial_compare_values(&Value::Integer(-2), &Value::Float(-2.5)), + Some(Ordering::Greater) + ); + assert!(coerced_eq(&Value::Integer(3), &Value::Float(3.0))); + // `i64::MAX` rounds up to `2^63` as an `f64`, so it is below that float. + const TWO_POW_63: f64 = 9_223_372_036_854_775_808.0; + assert_eq!( + partial_compare_values(&Value::Integer(i64::MAX), &Value::Float(TWO_POW_63)), + Some(Ordering::Less) + ); + assert_eq!( + partial_compare_values(&Value::Integer(i64::MIN), &Value::Float(-TWO_POW_63)), + Some(Ordering::Equal) + ); + assert_eq!( + partial_compare_values(&Value::Integer(i64::MIN), &Value::Float(f64::NEG_INFINITY)), + Some(Ordering::Greater) + ); + assert_eq!( + partial_compare_values(&Value::Integer(0), &Value::Float(f64::NAN)), + Some(Ordering::Less) + ); + } + + /// NaN sorts above every number and equals NaN, so a sort is total. + #[test] + fn nan_sorts_above_every_number() { + let nan = Value::Float(f64::NAN); + assert_eq!( + compare_values(&nan, &Value::Float(f64::INFINITY)), + Ordering::Greater + ); + assert_eq!( + compare_values(&Value::Integer(i64::MAX), &nan), + Ordering::Less + ); + assert_eq!( + compare_values(&Value::Decimal(rust_decimal::Decimal::MAX), &nan), + Ordering::Less + ); + assert_eq!(compare_values(&nan, &nan), Ordering::Equal); + assert!(coerced_eq(&nan, &Value::Float(f64::NAN))); + let mut values = [ + nan.clone(), + Value::Integer(3), + Value::Float(f64::NEG_INFINITY), + nan.clone(), + Value::Float(0.5), + ]; + values.sort_by(compare_values); + assert_eq!(values[0], Value::Float(f64::NEG_INFINITY)); + assert_eq!(values[1], Value::Float(0.5)); + assert_eq!(values[2], Value::Integer(3)); + assert!(matches!(values[3], Value::Float(f) if f.is_nan())); + assert!(matches!(values[4], Value::Float(f) if f.is_nan())); + } + + /// Fractional decimals and decimal text compare exactly, never through + /// `f64`. + #[test] + fn fractional_decimals_compare_exactly() { + use std::str::FromStr; + let low = rust_decimal::Decimal::from_str("12345678901234567.01").unwrap(); + let high = rust_decimal::Decimal::from_str("12345678901234567.02").unwrap(); + assert_eq!( + partial_compare_values(&Value::Decimal(low), &Value::Decimal(high)), + Some(Ordering::Less) + ); + assert!(!coerced_eq(&Value::Decimal(low), &Value::Decimal(high))); + let low_text = Value::String("12345678901234567.01".into()); + let high_text = Value::String("12345678901234567.02".into()); + assert_eq!(compare_values(&high_text, &low_text), Ordering::Greater); + assert!(!coerced_eq(&low_text, &high_text)); + assert!(coerced_eq(&low_text, &Value::Decimal(low))); + let half = rust_decimal::Decimal::from_str("9007199254740993.5").unwrap(); + assert_eq!( + partial_compare_values( + &Value::Decimal(half), + &Value::Integer(9_007_199_254_740_993) + ), + Some(Ordering::Greater) + ); + assert_eq!( + partial_compare_values( + &Value::Integer(9_007_199_254_740_994), + &Value::Decimal(half) + ), + Some(Ordering::Greater) + ); + } + + /// A `u64` above `i64::MAX` is an integral `Decimal`, which compares + /// exactly against integers instead of rounding through `f64`. + #[test] + fn integral_decimals_compare_exactly() { + let u_max = Value::from_u64(u64::MAX); + let u_max_less_one = Value::from_u64(u64::MAX - 1); + assert_eq!( + partial_compare_values(&u_max_less_one, &u_max), + Some(Ordering::Less) + ); + assert!(!coerced_eq(&u_max_less_one, &u_max)); + assert_eq!( + partial_compare_values(&u_max, &Value::Integer(i64::MAX)), + Some(Ordering::Greater) + ); + let two_pow_63 = Value::from_u64(1 << 63); + assert_eq!( + partial_compare_values(&two_pow_63, &Value::Integer(i64::MAX)), + Some(Ordering::Greater) + ); + assert_eq!( + partial_compare_values( + &Value::Decimal(rust_decimal::Decimal::new(15, 1)), + &Value::Integer(1) + ), + Some(Ordering::Greater) + ); + } + + #[test] + fn integer_strings_compare_exactly_against_integers() { + let text = Value::String("9007199254740993".into()); + assert!(!coerced_eq(&text, &Value::Integer(9_007_199_254_740_992))); + assert!(coerced_eq(&text, &Value::Integer(9_007_199_254_740_993))); + assert!(coerced_eq(&Value::String("5.0".into()), &Value::Integer(5))); + } + #[test] fn instants_compare_by_micros_against_instants_and_iso_strings() { let earlier = Value::NaiveDateTime(nodedb_types::NdbDateTime::from_micros( diff --git a/nodedb-query/src/window/aggregate.rs b/nodedb-query/src/window/aggregate.rs index 4ad8a384e..c091e6a87 100644 --- a/nodedb-query/src/window/aggregate.rs +++ b/nodedb-query/src/window/aggregate.rs @@ -9,15 +9,21 @@ //! - All other frame combinations → per-row frame evaluator that computes the //! concrete `[start_idx, end_idx]` slice for every row and aggregates over //! it. +//! +//! A float result keeps NaN and ±Infinity, as in PostgreSQL. + +use nodedb_types::Value; use super::arg::{ArgValues, arg_at, eval_arg_values}; +use super::extremum::value_replaces; use super::frame::{build_peer_groups, evaluate_frame_bounds}; use super::helpers::{as_f64, set_window_col}; use super::running::running_aggregate; use super::spec::{FrameBound, WindowFuncSpec}; +use crate::numeric_sum::{ExactSum, sum_input}; pub(super) fn apply_aggregate_window( - rows: &mut [(String, serde_json::Value)], + rows: &mut [(String, Value)], indices: &[usize], spec: &WindowFuncSpec, ) -> Result<(), crate::expr::EvalError> { @@ -49,7 +55,7 @@ pub(super) fn apply_aggregate_window( /// 2. Aggregate the evaluated argument over `indices[start_idx..=end_idx]`. /// 3. Write the result back under `spec.alias`. fn per_row_aggregate( - rows: &mut [(String, serde_json::Value)], + rows: &mut [(String, Value)], indices: &[usize], spec: &WindowFuncSpec, arg_values: &ArgValues, @@ -61,11 +67,11 @@ fn per_row_aggregate( // Extract order-by values for RANGE numeric offsets. let order_expr = spec.order_by.first().map(|(expr, _)| expr); - let order_values: Vec = indices + let order_values: Vec = indices .iter() .map(|&i| match order_expr { - Some(expr) => super::helpers::eval_expr_on_json(expr, &rows[i].1), - None => Ok(serde_json::Value::Null), + Some(expr) => expr.eval(&rows[i].1), + None => Ok(Value::Null), }) .collect::, _>>()?; @@ -82,14 +88,14 @@ fn per_row_aggregate( .map(|pos| as_f64(&arg_at(arg_values, pos))) .collect(); - let results: Vec = (0..len) + let results: Vec = (0..len) .map(|pos| { let (start_idx, end_idx) = evaluate_frame_bounds(&spec.frame, pos, len, &order_values, &peer_groups); aggregate_slice(&all_vals, arg_values, spec, start_idx, end_idx) }) - .collect(); + .collect::>()?; for (pos, result) in results.into_iter().enumerate() { let row_idx = indices[pos]; @@ -98,6 +104,19 @@ fn per_row_aggregate( Ok(()) } +/// The exact SUM / AVG state of the frame slice `[start_idx, end_idx]`. +fn frame_sum(arg_values: &ArgValues, start_idx: usize, end_idx: usize) -> ExactSum { + let mut acc = ExactSum::new(); + if let Some(values) = arg_values { + for v in values.get(start_idx..=end_idx).unwrap_or(&[]) { + if let Some(n) = sum_input(v) { + acc.add_value(&n); + } + } + } + acc +} + /// Aggregate the evaluated argument over the frame slice /// `[start_idx, end_idx]` (partition positions, not row indices). fn aggregate_slice( @@ -106,57 +125,43 @@ fn aggregate_slice( spec: &WindowFuncSpec, start_idx: usize, end_idx: usize, -) -> serde_json::Value { - let slice_vals: Vec = all_vals[start_idx..=end_idx] - .iter() - .filter_map(|v| *v) - .collect(); - - match spec.func_name.as_str() { - "sum" => { - let rt = crate::simd_agg::ts_runtime(); - serde_json::json!((rt.sum_f64)(&slice_vals)) - } +) -> Result { + Ok(match spec.func_name.as_str() { + "sum" => frame_sum(arg_values, start_idx, end_idx).sum()?, // `COUNT(*)` counts frame rows; `COUNT(expr)` counts the rows whose // argument is non-NULL, so a NULL argument is excluded rather than // inflating the count. "count" => match arg_values { - None => serde_json::json!(end_idx - start_idx + 1), - Some(values) => serde_json::json!( + None => Value::from_u64((end_idx - start_idx + 1) as u64), + Some(values) => Value::from_u64( values[start_idx..=end_idx] .iter() .filter(|v| !v.is_null()) - .count() + .count() as u64, ), }, - "avg" => { - if slice_vals.is_empty() { - serde_json::Value::Null - } else { - let rt = crate::simd_agg::ts_runtime(); - serde_json::json!((rt.sum_f64)(&slice_vals) / slice_vals.len() as f64) - } - } - "min" => { - if slice_vals.is_empty() { - serde_json::Value::Null - } else { - let rt = crate::simd_agg::ts_runtime(); - serde_json::json!((rt.min_f64)(&slice_vals)) - } - } - "max" => { - if slice_vals.is_empty() { - serde_json::Value::Null - } else { - let rt = crate::simd_agg::ts_runtime(); - serde_json::json!((rt.max_f64)(&slice_vals)) + "avg" => frame_sum(arg_values, start_idx, end_idx).avg()?, + "min" | "max" => { + let want_max = spec.func_name == "max"; + let Some(values) = arg_values else { + return Ok(Value::Null); + }; + let frame_values = values.get(start_idx..=end_idx).unwrap_or(&[]); + let mut best: Option<&Value> = None; + for (numeric, candidate) in all_vals[start_idx..=end_idx].iter().zip(frame_values) { + if numeric.is_none() { + continue; + } + if value_replaces(candidate, best, want_max) { + best = Some(candidate); + } } + best.cloned().unwrap_or(Value::Null) } "first_value" => arg_at(arg_values, start_idx), "last_value" => arg_at(arg_values, end_idx), - _ => serde_json::Value::Null, - } + _ => Value::Null, + }) } #[cfg(test)] @@ -164,11 +169,20 @@ mod tests { use super::super::spec::{FrameBound, WindowFrame, WindowFuncSpec}; use super::apply_aggregate_window; use crate::expr::SqlExpr; + use nodedb_types::Value; use serde_json::json; - fn numbered(n: usize) -> Vec<(String, serde_json::Value)> { + fn v(j: serde_json::Value) -> Value { + Value::from(j) + } + + fn res(row: &(String, Value)) -> Value { + row.1.get("result").cloned().unwrap_or(Value::Null) + } + + fn numbered(n: usize) -> Vec<(String, Value)> { (1..=n) - .map(|i| (i.to_string(), json!({ "n": i as i64 }))) + .map(|i| (i.to_string(), v(json!({ "n": i as i64 })))) .collect() } @@ -228,11 +242,11 @@ mod tests { // row 2 (n=3): sum of [2,3,4] = 9 // row 3 (n=4): sum of [3,4,5] = 12 // row 4 (n=5): sum of [4,5] = 9 - assert_eq!(rows[0].1["result"], json!(3.0)); - assert_eq!(rows[1].1["result"], json!(6.0)); - assert_eq!(rows[2].1["result"], json!(9.0)); - assert_eq!(rows[3].1["result"], json!(12.0)); - assert_eq!(rows[4].1["result"], json!(9.0)); + assert_eq!(res(&rows[0]), Value::Integer(3)); + assert_eq!(res(&rows[1]), Value::Integer(6)); + assert_eq!(res(&rows[2]), Value::Integer(9)); + assert_eq!(res(&rows[3]), Value::Integer(12)); + assert_eq!(res(&rows[4]), Value::Integer(9)); } #[test] @@ -245,11 +259,11 @@ mod tests { rows_frame(FrameBound::UnboundedPreceding, FrameBound::CurrentRow), ); apply_aggregate_window(&mut rows, &indices, &spec).unwrap(); - assert_eq!(rows[0].1["result"], json!(1.0)); - assert_eq!(rows[1].1["result"], json!(3.0)); - assert_eq!(rows[2].1["result"], json!(6.0)); - assert_eq!(rows[3].1["result"], json!(10.0)); - assert_eq!(rows[4].1["result"], json!(15.0)); + assert_eq!(res(&rows[0]), Value::Integer(1)); + assert_eq!(res(&rows[1]), Value::Integer(3)); + assert_eq!(res(&rows[2]), Value::Integer(6)); + assert_eq!(res(&rows[3]), Value::Integer(10)); + assert_eq!(res(&rows[4]), Value::Integer(15)); } #[test] @@ -265,9 +279,9 @@ mod tests { // row 0: sum 1+2+3+4+5=15 // row 1: sum 2+3+4+5=14 // ... - assert_eq!(rows[0].1["result"], json!(15.0)); - assert_eq!(rows[1].1["result"], json!(14.0)); - assert_eq!(rows[4].1["result"], json!(5.0)); + assert_eq!(res(&rows[0]), Value::Integer(15)); + assert_eq!(res(&rows[1]), Value::Integer(14)); + assert_eq!(res(&rows[4]), Value::Integer(5)); } // ── RANGE ───────────────────────────────────────────────────────────────── @@ -276,10 +290,10 @@ mod tests { fn range_unbounded_preceding_current_row_with_ties() { // Values: n in [1, 1, 2, 3] — two rows with n=1 share same frame. let mut rows = vec![ - ("a".into(), json!({"n": 1i64})), - ("b".into(), json!({"n": 1i64})), - ("c".into(), json!({"n": 2i64})), - ("d".into(), json!({"n": 3i64})), + ("a".into(), v(json!({"n": 1i64}))), + ("b".into(), v(json!({"n": 1i64}))), + ("c".into(), v(json!({"n": 2i64}))), + ("d".into(), v(json!({"n": 3i64}))), ]; let indices: Vec = (0..4).collect(); // Use the fast running path (RANGE UNBOUNDED PRECEDING TO CURRENT ROW) @@ -292,13 +306,229 @@ mod tests { ); apply_aggregate_window(&mut rows, &indices, &spec).unwrap(); // Row a (n=1, pos=0): CURRENT ROW expands to last peer at pos=1, sum=1+1=2 - assert_eq!(rows[0].1["result"], json!(2.0)); + assert_eq!(res(&rows[0]), Value::Integer(2)); // Row b (n=1, pos=1): same - assert_eq!(rows[1].1["result"], json!(2.0)); + assert_eq!(res(&rows[1]), Value::Integer(2)); // Row c (n=2): sum=1+1+2=4 - assert_eq!(rows[2].1["result"], json!(4.0)); + assert_eq!(res(&rows[2]), Value::Integer(4)); // Row d (n=3): sum=1+1+2+3=7 - assert_eq!(rows[3].1["result"], json!(7.0)); + assert_eq!(res(&rows[3]), Value::Integer(7)); + } + + // ── SUM / AVG exactness ─────────────────────────────────────────────────── + + #[test] + fn frame_and_running_sum_keep_integers_above_2_pow_53_exact() { + let vals = [ + Value::Integer(9_007_199_254_740_993), + Value::Integer(9_007_199_254_740_992), + ]; + let indices: Vec = (0..2).collect(); + + let mut rows = keyed(&vals); + apply_aggregate_window(&mut rows, &indices, &make_spec("sum", "v", whole())).unwrap(); + assert_eq!(res(&rows[0]), Value::Integer(18_014_398_509_481_985)); + + let running = range_frame(FrameBound::UnboundedPreceding, FrameBound::CurrentRow); + let mut rows = keyed(&vals); + apply_aggregate_window(&mut rows, &indices, &make_spec("sum", "v", running)).unwrap(); + assert_eq!( + results(&rows), + vec![ + Value::Integer(9_007_199_254_740_993), + Value::Integer(18_014_398_509_481_985) + ] + ); + + let mut rows = keyed(&vals); + apply_aggregate_window(&mut rows, &indices, &make_spec("avg", "v", whole())).unwrap(); + assert_eq!(res(&rows[0]), Value::Float(9_007_199_254_740_992.0)); + } + + #[test] + fn frame_sum_past_i64_and_u64_is_exact() { + let vals = [ + Value::Integer(i64::MAX), + Value::from_u64(u64::MAX), + Value::from_u64(u64::MAX), + ]; + let indices: Vec = (0..3).collect(); + let mut rows = keyed(&vals); + apply_aggregate_window(&mut rows, &indices, &make_spec("sum", "v", whole())).unwrap(); + let want = i128::from(i64::MAX) + 2 * i128::from(u64::MAX); + assert_eq!( + res(&rows[0]), + Value::Decimal(rust_decimal::Decimal::from_i128_with_scale(want, 0)) + ); + } + + #[test] + fn frame_sum_mixed_int_float_is_float_and_empty_is_null() { + let vals = [Value::Integer(2), Value::Float(0.5), Value::Null]; + let indices: Vec = (0..3).collect(); + let mut rows = keyed(&vals); + apply_aggregate_window(&mut rows, &indices, &make_spec("sum", "v", whole())).unwrap(); + assert_eq!(res(&rows[0]), Value::Float(2.5)); + + let one_row = rows_frame(FrameBound::CurrentRow, FrameBound::CurrentRow); + let mut rows = keyed(&vals); + apply_aggregate_window(&mut rows, &indices, &make_spec("sum", "v", one_row)).unwrap(); + assert_eq!(res(&rows[2]), Value::Null); + } + + /// Float overflow is `Infinity` and `Infinity` plus `-Infinity` is NaN; + /// both reach the window column as floats, as in PostgreSQL. + #[test] + fn frame_and_running_sum_keep_non_finite_floats() { + let indices: Vec = (0..2).collect(); + let overflow = [Value::Float(1e308), Value::Float(1e308)]; + + let mut rows = keyed(&overflow); + apply_aggregate_window(&mut rows, &indices, &make_spec("sum", "v", whole())).unwrap(); + assert_eq!(res(&rows[0]), Value::Float(f64::INFINITY)); + + let running = range_frame(FrameBound::UnboundedPreceding, FrameBound::CurrentRow); + let mut rows = keyed(&overflow); + apply_aggregate_window(&mut rows, &indices, &make_spec("sum", "v", running)).unwrap(); + assert_eq!( + results(&rows), + vec![Value::Float(1e308), Value::Float(f64::INFINITY)] + ); + + let cancel = [Value::Float(f64::INFINITY), Value::Float(f64::NEG_INFINITY)]; + let mut rows = keyed(&cancel); + apply_aggregate_window(&mut rows, &indices, &make_spec("avg", "v", whole())).unwrap(); + assert!(matches!(res(&rows[0]), Value::Float(f) if f.is_nan())); + } + + // ── MIN / MAX exactness ─────────────────────────────────────────────────── + + /// Rows with order key `n` = 1..; `v` carries the given values. + fn keyed(vals: &[Value]) -> Vec<(String, Value)> { + vals.iter() + .enumerate() + .map(|(i, val)| { + let doc = std::collections::HashMap::from([ + ("n".to_string(), Value::Integer(i as i64 + 1)), + ("v".to_string(), val.clone()), + ]); + (i.to_string(), Value::Object(doc)) + }) + .collect() + } + + fn results(rows: &[(String, Value)]) -> Vec { + rows.iter().map(res).collect() + } + + fn whole() -> WindowFrame { + rows_frame( + FrameBound::UnboundedPreceding, + FrameBound::UnboundedFollowing, + ) + } + + #[test] + fn frame_min_max_keep_integers_above_2_pow_53_exact() { + let vals = [ + Value::Integer(9_007_199_254_740_993), + Value::Integer(9_007_199_254_740_992), + ]; + let indices: Vec = (0..2).collect(); + + let mut rows = keyed(&vals); + apply_aggregate_window(&mut rows, &indices, &make_spec("min", "v", whole())).unwrap(); + assert_eq!(res(&rows[0]), Value::Integer(9_007_199_254_740_992)); + + let mut rows = keyed(&vals); + apply_aggregate_window(&mut rows, &indices, &make_spec("max", "v", whole())).unwrap(); + assert_eq!(res(&rows[0]), Value::Integer(9_007_199_254_740_993)); + } + + #[test] + fn frame_min_max_over_nanosecond_timestamps() { + let vals = [ + Value::Integer(1_700_000_000_000_000_002), + Value::Integer(1_700_000_000_000_000_001), + Value::Integer(1_700_000_000_000_000_003), + ]; + let indices: Vec = (0..3).collect(); + + let mut rows = keyed(&vals); + apply_aggregate_window(&mut rows, &indices, &make_spec("min", "v", whole())).unwrap(); + assert_eq!(res(&rows[2]), Value::Integer(1_700_000_000_000_000_001)); + + let mut rows = keyed(&vals); + apply_aggregate_window(&mut rows, &indices, &make_spec("max", "v", whole())).unwrap(); + assert_eq!(res(&rows[2]), Value::Integer(1_700_000_000_000_000_003)); + } + + #[test] + fn frame_min_max_mixed_int_float_return_original_type() { + let vals = [ + Value::Integer(9_007_199_254_740_993), + Value::Float(9_007_199_254_740_992.0), + Value::Float(1.5), + Value::Integer(2), + ]; + let indices: Vec = (0..4).collect(); + + let mut rows = keyed(&vals); + apply_aggregate_window(&mut rows, &indices, &make_spec("min", "v", whole())).unwrap(); + assert_eq!(res(&rows[0]), Value::Float(1.5)); + + let mut rows = keyed(&vals); + apply_aggregate_window(&mut rows, &indices, &make_spec("max", "v", whole())).unwrap(); + assert_eq!(res(&rows[0]), Value::Integer(9_007_199_254_740_993)); + } + + /// NaN sorts above every number: MAX over a frame with NaN is NaN, and + /// MIN skips it. + #[test] + fn frame_min_max_place_nan_above_every_number() { + let vals = [Value::Float(f64::NAN), Value::Integer(3), Value::Float(1.5)]; + let indices: Vec = (0..3).collect(); + + let mut rows = keyed(&vals); + apply_aggregate_window(&mut rows, &indices, &make_spec("max", "v", whole())).unwrap(); + assert!(matches!(res(&rows[0]), Value::Float(f) if f.is_nan())); + + let mut rows = keyed(&vals); + apply_aggregate_window(&mut rows, &indices, &make_spec("min", "v", whole())).unwrap(); + assert_eq!(res(&rows[0]), Value::Float(1.5)); + } + + #[test] + fn running_min_max_keep_integers_exact() { + let vals = [ + Value::Integer(9_007_199_254_740_993), + Value::Integer(9_007_199_254_740_992), + Value::Integer(9_007_199_254_740_995), + ]; + let indices: Vec = (0..3).collect(); + let running = || range_frame(FrameBound::UnboundedPreceding, FrameBound::CurrentRow); + + let mut rows = keyed(&vals); + apply_aggregate_window(&mut rows, &indices, &make_spec("min", "v", running())).unwrap(); + assert_eq!( + results(&rows), + vec![ + Value::Integer(9_007_199_254_740_993), + Value::Integer(9_007_199_254_740_992), + Value::Integer(9_007_199_254_740_992), + ] + ); + + let mut rows = keyed(&vals); + apply_aggregate_window(&mut rows, &indices, &make_spec("max", "v", running())).unwrap(); + assert_eq!( + results(&rows), + vec![ + Value::Integer(9_007_199_254_740_993), + Value::Integer(9_007_199_254_740_993), + Value::Integer(9_007_199_254_740_995), + ] + ); } // ── GROUPS ──────────────────────────────────────────────────────────────── @@ -307,11 +537,11 @@ mod tests { fn groups_1_preceding_1_following_sum() { // Values: [1, 1, 2, 3, 3] — groups [0, 0, 1, 2, 2] let mut rows = vec![ - ("a".into(), json!({"n": 1i64})), - ("b".into(), json!({"n": 1i64})), - ("c".into(), json!({"n": 2i64})), - ("d".into(), json!({"n": 3i64})), - ("e".into(), json!({"n": 3i64})), + ("a".into(), v(json!({"n": 1i64}))), + ("b".into(), v(json!({"n": 1i64}))), + ("c".into(), v(json!({"n": 2i64}))), + ("d".into(), v(json!({"n": 3i64}))), + ("e".into(), v(json!({"n": 3i64}))), ]; let indices: Vec = (0..5).collect(); let spec = make_spec( @@ -321,14 +551,14 @@ mod tests { ); apply_aggregate_window(&mut rows, &indices, &spec).unwrap(); // pos=0 (group 0): frame → groups 0..=1 → rows 0..=2 → sum=1+1+2=4 - assert_eq!(rows[0].1["result"], json!(4.0)); + assert_eq!(res(&rows[0]), Value::Integer(4)); // pos=1 (group 0): same frame - assert_eq!(rows[1].1["result"], json!(4.0)); + assert_eq!(res(&rows[1]), Value::Integer(4)); // pos=2 (group 1): frame → groups 0..=2 → rows 0..=4 → sum=1+1+2+3+3=10 - assert_eq!(rows[2].1["result"], json!(10.0)); + assert_eq!(res(&rows[2]), Value::Integer(10)); // pos=3 (group 2): frame → groups 1..=2 → rows 2..=4 → sum=2+3+3=8 - assert_eq!(rows[3].1["result"], json!(8.0)); + assert_eq!(res(&rows[3]), Value::Integer(8)); // pos=4 (group 2): same - assert_eq!(rows[4].1["result"], json!(8.0)); + assert_eq!(res(&rows[4]), Value::Integer(8)); } } diff --git a/nodedb-query/src/window/arg.rs b/nodedb-query/src/window/arg.rs index 8084912bf..4562ee954 100644 --- a/nodedb-query/src/window/arg.rs +++ b/nodedb-query/src/window/arg.rs @@ -10,7 +10,8 @@ //! projection or WHERE clause, and a non-column argument yields its computed //! value rather than a NULL placeholder. -use super::helpers::eval_expr_on_json; +use nodedb_types::Value; + use super::spec::WindowFuncSpec; use crate::expr::{EvalError, SqlExpr}; @@ -19,7 +20,7 @@ use crate::expr::{EvalError, SqlExpr}; /// /// `None` means the function was called without that argument — `COUNT(*)`, /// `ROW_NUMBER()` — which is distinct from an argument that evaluated to NULL. -pub(super) type ArgValues = Option>; +pub(super) type ArgValues = Option>; /// Evaluate argument `idx` of `spec` once for every row in the partition. /// @@ -27,7 +28,7 @@ pub(super) type ArgValues = Option>; /// when the frame evaluator revisits rows, and gives every aggregate the same /// values the frame bounds are computed against. pub(super) fn eval_arg_values( - rows: &[(String, serde_json::Value)], + rows: &[(String, Value)], indices: &[usize], spec: &WindowFuncSpec, idx: usize, @@ -37,18 +38,18 @@ pub(super) fn eval_arg_values( }; let values = indices .iter() - .map(|&i| eval_expr_on_json(expr, &rows[i].1)) + .map(|&i| expr.eval(&rows[i].1)) .collect::, _>>()?; Ok(Some(values)) } /// Value of argument `idx` at partition position `pos`, or NULL when the /// function was called without that argument. -pub(super) fn arg_at(values: &ArgValues, pos: usize) -> serde_json::Value { +pub(super) fn arg_at(values: &ArgValues, pos: usize) -> Value { values .as_ref() .and_then(|v| v.get(pos).cloned()) - .unwrap_or(serde_json::Value::Null) + .unwrap_or(Value::Null) } /// Resolve a constant integer argument — the `LAG`/`LEAD` offset, `NTILE` @@ -69,12 +70,12 @@ pub(super) fn const_usize_arg(spec: &WindowFuncSpec, idx: usize, default: usize) /// Resolve the constant `default` argument of `LAG`/`LEAD` — the value /// returned when the offset falls outside the partition. -pub(super) fn const_default_arg(spec: &WindowFuncSpec, idx: usize) -> serde_json::Value { +pub(super) fn const_default_arg(spec: &WindowFuncSpec, idx: usize) -> Value { spec.args .get(idx) .and_then(|e| match e { - SqlExpr::Literal(v) => Some(serde_json::Value::from(v.clone())), + SqlExpr::Literal(v) => Some(v.clone()), _ => None, }) - .unwrap_or(serde_json::Value::Null) + .unwrap_or(Value::Null) } diff --git a/nodedb-query/src/window/eval.rs b/nodedb-query/src/window/eval.rs index ad741cdf3..7a5e23f5c 100644 --- a/nodedb-query/src/window/eval.rs +++ b/nodedb-query/src/window/eval.rs @@ -13,10 +13,11 @@ use super::spec::WindowFuncSpec; /// Evaluate window functions over sorted, partitioned rows. /// -/// `rows` is the sorted result set. Each row is a `(doc_id, serde_json::Value)`. -/// The same rows are mutated in place with window columns appended to each -/// document. The row array keeps its input order; each spec's partitions are -/// ordered by that spec's own ORDER BY, independent of the row array order. +/// `rows` is the result set. Each row is a `(doc_id, Value::Object)`. The +/// same rows are mutated in place with window columns appended to each +/// document. A window result keeps NaN and ±Infinity. The row array keeps +/// its input order; each spec's partitions are ordered by that spec's own +/// ORDER BY, independent of the row array order. /// /// Unknown window function names must be rejected by the planner before /// reaching this dispatcher; an unrecognised name here is an internal bug @@ -26,7 +27,7 @@ use super::spec::WindowFuncSpec; /// expression propagates as `Err(EvalError::DivisionByZero)` rather than /// folding to NULL. pub fn evaluate_window_functions( - rows: &mut [(String, serde_json::Value)], + rows: &mut [(String, nodedb_types::Value)], specs: &[WindowFuncSpec], ) -> Result<(), crate::expr::EvalError> { for spec in specs { @@ -68,39 +69,56 @@ mod tests { use super::super::spec::{WindowFrame, WindowFuncSpec}; use super::evaluate_window_functions; use crate::expr::SqlExpr; + use nodedb_types::Value; use serde_json::json; - fn make_rows() -> Vec<(String, serde_json::Value)> { + type Row = (String, Value); + + fn row(id: &str, doc: serde_json::Value) -> Row { + (id.to_string(), Value::from(doc)) + } + + /// Column `name` of `row` in its JSON form, NULL when absent. + fn col(row: &Row, name: &str) -> serde_json::Value { + serde_json::Value::from(row.1.get(name).cloned().unwrap_or(Value::Null)) + } + + fn make_rows() -> Vec { vec![ - ( - "1".into(), - json!({"dept": "eng", "salary": 100, "name": "Alice"}), - ), - ( - "2".into(), - json!({"dept": "eng", "salary": 120, "name": "Bob"}), - ), - ( - "3".into(), - json!({"dept": "eng", "salary": 90, "name": "Carol"}), - ), - ( - "4".into(), - json!({"dept": "sales", "salary": 80, "name": "Dave"}), - ), - ( - "5".into(), - json!({"dept": "sales", "salary": 110, "name": "Eve"}), - ), + row("1", json!({"dept": "eng", "salary": 100, "name": "Alice"})), + row("2", json!({"dept": "eng", "salary": 120, "name": "Bob"})), + row("3", json!({"dept": "eng", "salary": 90, "name": "Carol"})), + row("4", json!({"dept": "sales", "salary": 80, "name": "Dave"})), + row("5", json!({"dept": "sales", "salary": 110, "name": "Eve"})), ] } - fn numbered(n: usize) -> Vec<(String, serde_json::Value)> { + fn numbered(n: usize) -> Vec { (1..=n) - .map(|i| (i.to_string(), json!({ "n": i }))) + .map(|i| row(&i.to_string(), json!({ "n": i }))) .collect() } + fn peers() -> Vec { + vec![ + row("a", json!({"n": 1})), + row("b", json!({"n": 1})), + row("c", json!({"n": 2})), + row("d", json!({"n": 3})), + ] + } + + fn ordered_by_n(alias: &str, func: &str) -> WindowFuncSpec { + WindowFuncSpec { + alias: alias.into(), + func_name: func.into(), + args: vec![], + partition_by: vec![], + order_by: vec![(SqlExpr::Column("n".into()), true)], + frame: WindowFrame::default(), + } + } + #[test] fn row_number_single_partition() { let mut rows = make_rows(); @@ -113,8 +131,8 @@ mod tests { frame: WindowFrame::default(), }; evaluate_window_functions(&mut rows, &[spec]).unwrap(); - assert_eq!(rows[0].1["rn"], json!(1)); - assert_eq!(rows[4].1["rn"], json!(5)); + assert_eq!(col(&rows[0], "rn"), json!(1)); + assert_eq!(col(&rows[4], "rn"), json!(5)); } #[test] @@ -129,10 +147,10 @@ mod tests { frame: WindowFrame::default(), }; evaluate_window_functions(&mut rows, &[spec]).unwrap(); - assert_eq!(rows[0].1["rn"], json!(1)); - assert_eq!(rows[2].1["rn"], json!(3)); - assert_eq!(rows[3].1["rn"], json!(1)); - assert_eq!(rows[4].1["rn"], json!(2)); + assert_eq!(col(&rows[0], "rn"), json!(1)); + assert_eq!(col(&rows[2], "rn"), json!(3)); + assert_eq!(col(&rows[3], "rn"), json!(1)); + assert_eq!(col(&rows[4], "rn"), json!(2)); } #[test] @@ -149,120 +167,122 @@ mod tests { evaluate_window_functions(&mut rows, &[spec]).unwrap(); // The frame runs in salary order within each dept, not in row // arrival order: eng = Carol(90) → Alice(100) → Bob(120). - assert_eq!(rows[0].1["running_total"], json!(190.0)); - assert_eq!(rows[1].1["running_total"], json!(310.0)); - assert_eq!(rows[2].1["running_total"], json!(90.0)); - assert_eq!(rows[3].1["running_total"], json!(80.0)); - assert_eq!(rows[4].1["running_total"], json!(190.0)); + // Integer salaries total exactly as integers. + assert_eq!(col(&rows[0], "running_total"), json!(190)); + assert_eq!(col(&rows[1], "running_total"), json!(310)); + assert_eq!(col(&rows[2], "running_total"), json!(90)); + assert_eq!(col(&rows[3], "running_total"), json!(80)); + assert_eq!(col(&rows[4], "running_total"), json!(190)); + } + + /// A float overflow reaches the window column as `Infinity`, not NULL. + #[test] + fn running_sum_keeps_infinity() { + let mut rows = vec![ + ( + "a".to_string(), + Value::Object( + [ + ("n".to_string(), Value::Integer(1)), + ("x".to_string(), Value::Float(1e308)), + ] + .into(), + ), + ), + ( + "b".to_string(), + Value::Object( + [ + ("n".to_string(), Value::Integer(2)), + ("x".to_string(), Value::Float(1e308)), + ] + .into(), + ), + ), + ]; + let mut spec = ordered_by_n("total", "sum"); + spec.args = vec![SqlExpr::Column("x".into())]; + evaluate_window_functions(&mut rows, &[spec]).unwrap(); + assert_eq!(rows[0].1.get("total"), Some(&Value::Float(1e308))); + assert_eq!(rows[1].1.get("total"), Some(&Value::Float(f64::INFINITY))); + } + + /// NaN order keys are peers of each other and sort after every number. + #[test] + fn nan_order_keys_rank_last_as_peers() { + let doc = |n: f64| Value::Object([("n".to_string(), Value::Float(n))].into()); + let mut rows = vec![ + ("a".to_string(), doc(f64::NAN)), + ("b".to_string(), doc(2.0)), + ("c".to_string(), doc(f64::NAN)), + ("d".to_string(), doc(f64::INFINITY)), + ]; + evaluate_window_functions(&mut rows, &[ordered_by_n("rnk", "rank")]).unwrap(); + assert_eq!(col(&rows[1], "rnk"), json!(1)); + assert_eq!(col(&rows[3], "rnk"), json!(2)); + assert_eq!(col(&rows[0], "rnk"), json!(3)); + assert_eq!(col(&rows[2], "rnk"), json!(3)); } #[test] fn percent_rank_distinct_keys() { let mut rows = numbered(5); - let spec = WindowFuncSpec { - alias: "pr".into(), - func_name: "percent_rank".into(), - args: vec![], - partition_by: vec![], - order_by: vec![(SqlExpr::Column("n".into()), true)], - frame: WindowFrame::default(), - }; - evaluate_window_functions(&mut rows, &[spec]).unwrap(); - assert_eq!(rows[0].1["pr"], json!(0.0)); - assert_eq!(rows[1].1["pr"], json!(0.25)); - assert_eq!(rows[2].1["pr"], json!(0.5)); - assert_eq!(rows[3].1["pr"], json!(0.75)); - assert_eq!(rows[4].1["pr"], json!(1.0)); + evaluate_window_functions(&mut rows, &[ordered_by_n("pr", "percent_rank")]).unwrap(); + assert_eq!(col(&rows[0], "pr"), json!(0.0)); + assert_eq!(col(&rows[1], "pr"), json!(0.25)); + assert_eq!(col(&rows[2], "pr"), json!(0.5)); + assert_eq!(col(&rows[3], "pr"), json!(0.75)); + assert_eq!(col(&rows[4], "pr"), json!(1.0)); } #[test] fn percent_rank_with_peers() { // Peers share the leader's rank, so [1, 1, 2, 3] yields ranks // 1, 1, 3, 4 → percent_rank = 0, 0, 2/3, 3/3. - let mut rows = vec![ - ("a".into(), json!({"n": 1})), - ("b".into(), json!({"n": 1})), - ("c".into(), json!({"n": 2})), - ("d".into(), json!({"n": 3})), - ]; - let spec = WindowFuncSpec { - alias: "pr".into(), - func_name: "percent_rank".into(), - args: vec![], - partition_by: vec![], - order_by: vec![(SqlExpr::Column("n".into()), true)], - frame: WindowFrame::default(), - }; - evaluate_window_functions(&mut rows, &[spec]).unwrap(); - assert_eq!(rows[0].1["pr"], json!(0.0)); - assert_eq!(rows[1].1["pr"], json!(0.0)); - assert_eq!(rows[2].1["pr"], json!(2.0 / 3.0)); - assert_eq!(rows[3].1["pr"], json!(1.0)); + let mut rows = peers(); + evaluate_window_functions(&mut rows, &[ordered_by_n("pr", "percent_rank")]).unwrap(); + assert_eq!(col(&rows[0], "pr"), json!(0.0)); + assert_eq!(col(&rows[1], "pr"), json!(0.0)); + assert_eq!(col(&rows[2], "pr"), json!(2.0 / 3.0)); + assert_eq!(col(&rows[3], "pr"), json!(1.0)); } #[test] fn cume_dist_distinct_keys() { let mut rows = numbered(5); - let spec = WindowFuncSpec { - alias: "cd".into(), - func_name: "cume_dist".into(), - args: vec![], - partition_by: vec![], - order_by: vec![(SqlExpr::Column("n".into()), true)], - frame: WindowFrame::default(), - }; - evaluate_window_functions(&mut rows, &[spec]).unwrap(); - assert_eq!(rows[0].1["cd"], json!(0.2)); - assert_eq!(rows[1].1["cd"], json!(0.4)); - assert_eq!(rows[2].1["cd"], json!(0.6)); - assert_eq!(rows[3].1["cd"], json!(0.8)); - assert_eq!(rows[4].1["cd"], json!(1.0)); + evaluate_window_functions(&mut rows, &[ordered_by_n("cd", "cume_dist")]).unwrap(); + assert_eq!(col(&rows[0], "cd"), json!(0.2)); + assert_eq!(col(&rows[1], "cd"), json!(0.4)); + assert_eq!(col(&rows[2], "cd"), json!(0.6)); + assert_eq!(col(&rows[3], "cd"), json!(0.8)); + assert_eq!(col(&rows[4], "cd"), json!(1.0)); } #[test] fn cume_dist_with_peers() { - let mut rows = vec![ - ("a".into(), json!({"n": 1})), - ("b".into(), json!({"n": 1})), - ("c".into(), json!({"n": 2})), - ("d".into(), json!({"n": 3})), - ]; - let spec = WindowFuncSpec { - alias: "cd".into(), - func_name: "cume_dist".into(), - args: vec![], - partition_by: vec![], - order_by: vec![(SqlExpr::Column("n".into()), true)], - frame: WindowFrame::default(), - }; - evaluate_window_functions(&mut rows, &[spec]).unwrap(); + let mut rows = peers(); + evaluate_window_functions(&mut rows, &[ordered_by_n("cd", "cume_dist")]).unwrap(); // Peers share value of last peer's position / N. - assert_eq!(rows[0].1["cd"], json!(0.5)); - assert_eq!(rows[1].1["cd"], json!(0.5)); - assert_eq!(rows[2].1["cd"], json!(0.75)); - assert_eq!(rows[3].1["cd"], json!(1.0)); + assert_eq!(col(&rows[0], "cd"), json!(0.5)); + assert_eq!(col(&rows[1], "cd"), json!(0.5)); + assert_eq!(col(&rows[2], "cd"), json!(0.75)); + assert_eq!(col(&rows[3], "cd"), json!(1.0)); } #[test] fn nth_value_returns_nth_then_holds() { let mut rows = numbered(5); - let spec = WindowFuncSpec { - alias: "nv".into(), - func_name: "nth_value".into(), - args: vec![ - SqlExpr::Column("n".into()), - SqlExpr::Literal(nodedb_types::Value::Integer(2)), - ], - partition_by: vec![], - order_by: vec![(SqlExpr::Column("n".into()), true)], - frame: WindowFrame::default(), - }; + let mut spec = ordered_by_n("nv", "nth_value"); + spec.args = vec![ + SqlExpr::Column("n".into()), + SqlExpr::Literal(Value::Integer(2)), + ]; evaluate_window_functions(&mut rows, &[spec]).unwrap(); - assert_eq!(rows[0].1["nv"], json!(null)); - assert_eq!(rows[1].1["nv"], json!(2)); - assert_eq!(rows[2].1["nv"], json!(2)); - assert_eq!(rows[3].1["nv"], json!(2)); - assert_eq!(rows[4].1["nv"], json!(2)); + assert_eq!(col(&rows[0], "nv"), json!(null)); + assert_eq!(col(&rows[1], "nv"), json!(2)); + assert_eq!(col(&rows[2], "nv"), json!(2)); + assert_eq!(col(&rows[3], "nv"), json!(2)); + assert_eq!(col(&rows[4], "nv"), json!(2)); } #[test] @@ -280,16 +300,16 @@ mod tests { frame: WindowFrame::default(), }; evaluate_window_functions(&mut rows, &[spec]).unwrap(); - assert_eq!(rows[0].1["name"], json!("Alice")); - assert_eq!(rows[1].1["name"], json!("Bob")); - assert_eq!(rows[2].1["name"], json!("Carol")); - assert_eq!(rows[3].1["name"], json!("Dave")); - assert_eq!(rows[4].1["name"], json!("Eve")); - assert_eq!(rows[0].1["rnk"], json!(2)); // Alice, salary 100 - assert_eq!(rows[1].1["rnk"], json!(1)); // Bob, salary 120 - assert_eq!(rows[2].1["rnk"], json!(3)); // Carol, salary 90 - assert_eq!(rows[3].1["rnk"], json!(2)); // Dave, salary 80 - assert_eq!(rows[4].1["rnk"], json!(1)); // Eve, salary 110 + assert_eq!(col(&rows[0], "name"), json!("Alice")); + assert_eq!(col(&rows[1], "name"), json!("Bob")); + assert_eq!(col(&rows[2], "name"), json!("Carol")); + assert_eq!(col(&rows[3], "name"), json!("Dave")); + assert_eq!(col(&rows[4], "name"), json!("Eve")); + assert_eq!(col(&rows[0], "rnk"), json!(2)); // Alice, salary 100 + assert_eq!(col(&rows[1], "rnk"), json!(1)); // Bob, salary 120 + assert_eq!(col(&rows[2], "rnk"), json!(3)); // Carol, salary 90 + assert_eq!(col(&rows[3], "rnk"), json!(2)); // Dave, salary 80 + assert_eq!(col(&rows[4], "rnk"), json!(1)); // Eve, salary 110 } #[test] diff --git a/nodedb-query/src/window/extremum.rs b/nodedb-query/src/window/extremum.rs new file mode 100644 index 000000000..464d42885 --- /dev/null +++ b/nodedb-query/src/window/extremum.rs @@ -0,0 +1,93 @@ +// SPDX-License-Identifier: Apache-2.0 + +//! Extremum selection for MIN / MAX aggregates. +//! +//! MIN / MAX keep the original value and return it unchanged. An integer +//! stays an integer, a float stays a float. A numeric pair compares by its +//! exact numeric reading, so integers above 2^53 and fractional decimals +//! never round through `f64`. NaN sorts above every number, as in +//! PostgreSQL: MAX over a NaN input is NaN, and MIN skips NaN unless every +//! input is NaN. A non-numeric pair uses the coerced value order. + +use std::cmp::Ordering; + +use nodedb_types::Value; + +/// Whether the `Value` `candidate` replaces `current` as the MIN (`want_max` +/// false) or MAX (`want_max` true). An empty `current` is always replaced. +pub fn value_replaces(candidate: &Value, current: Option<&Value>, want_max: bool) -> bool { + let Some(current) = current else { + return true; + }; + let order = crate::value_ops::numeric_order(candidate, current) + .unwrap_or_else(|| crate::value_ops::compare_values(candidate, current)); + let wanted = if want_max { + Ordering::Greater + } else { + Ordering::Less + }; + order == wanted +} + +/// Direction of an extremum function: `Some(false)` for MIN, `Some(true)` +/// for MAX, `None` for any other function. +pub fn extremum_direction(func_name: &str) -> Option { + match func_name { + "min" => Some(false), + "max" => Some(true), + _ => None, + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn integers_above_2_pow_53_compare_exactly() { + let a = Value::Integer(9_007_199_254_740_993); + let b = Value::Integer(9_007_199_254_740_992); + assert!(value_replaces(&b, Some(&a), false)); + assert!(!value_replaces(&b, Some(&a), true)); + assert!(value_replaces(&a, Some(&b), true)); + } + + /// NaN is above every number: it wins MAX and loses MIN. + #[test] + fn value_nan_is_the_largest_extreme() { + let one = Value::Integer(1); + let nan = Value::Float(f64::NAN); + assert!(!value_replaces(&nan, Some(&one), false)); + assert!(value_replaces(&nan, Some(&one), true)); + assert!(value_replaces(&one, Some(&nan), false)); + assert!(!value_replaces(&one, Some(&nan), true)); + assert!(!value_replaces(&nan, Some(&nan), true)); + let inf = Value::Float(f64::INFINITY); + assert!(value_replaces(&nan, Some(&inf), true)); + } + + /// MAX over decimal text one hundredth apart past `2^53` picks the + /// larger, whatever the input order. + #[test] + fn decimal_text_max_compares_exactly() { + let low = Value::String("12345678901234567.01".into()); + let high = Value::String("12345678901234567.02".into()); + assert!(value_replaces(&high, Some(&low), true)); + assert!(!value_replaces(&low, Some(&high), true)); + assert!(value_replaces(&low, Some(&high), false)); + let low_dec = Value::Decimal(rust_decimal::Decimal::new(1_234_567_890_123_456_701, 2)); + let high_dec = Value::Decimal(rust_decimal::Decimal::new(1_234_567_890_123_456_702, 2)); + assert!(value_replaces(&high_dec, Some(&low_dec), true)); + assert!(!value_replaces(&low_dec, Some(&high_dec), true)); + assert!(value_replaces(&high, Some(&low_dec), true)); + } + + #[test] + fn value_mixed_int_float_compare_exactly() { + let big = Value::Integer(9_007_199_254_740_993); + let float = Value::Float(9_007_199_254_740_992.0); + assert!(value_replaces(&big, Some(&float), true)); + assert!(!value_replaces(&big, Some(&float), false)); + assert!(value_replaces(&float, Some(&big), false)); + } +} diff --git a/nodedb-query/src/window/frame.rs b/nodedb-query/src/window/frame.rs index 90cd2dd64..6e011620d 100644 --- a/nodedb-query/src/window/frame.rs +++ b/nodedb-query/src/window/frame.rs @@ -7,8 +7,11 @@ //! concrete `[start_idx, end_idx]` inclusive range into the partition's index //! array. +use nodedb_types::Value; + use super::helpers::as_f64; use super::spec::{FrameBound, WindowFrame}; +use crate::value_ops::sort_peers; /// Resolve the frame for row at position `pos` in a partition of `len` rows. /// @@ -30,7 +33,7 @@ pub(super) fn evaluate_frame_bounds( frame: &WindowFrame, pos: usize, len: usize, - order_values: &[serde_json::Value], + order_values: &[Value], peer_groups: &[usize], ) -> (usize, usize) { match frame.mode.as_str() { @@ -45,14 +48,14 @@ pub(super) fn evaluate_frame_bounds( /// Build a per-row peer-group index array for a partition. /// -/// Two rows are in the same peer group when they share the same order-by -/// value. The returned vec has the same length as the partition; each element -/// is the zero-based group index of that row. -pub(super) fn build_peer_groups(order_values: &[serde_json::Value]) -> Vec { +/// Two rows are in the same peer group when their order-by values are peers +/// (see [`sort_peers`]). The returned vec has the same length as the +/// partition; each element is the zero-based group index of that row. +pub(super) fn build_peer_groups(order_values: &[Value]) -> Vec { let mut groups = Vec::with_capacity(order_values.len()); let mut current_group = 0usize; for (i, val) in order_values.iter().enumerate() { - if i > 0 && val != &order_values[i - 1] { + if i > 0 && !sort_peers(val, &order_values[i - 1]) { current_group += 1; } groups.push(current_group); @@ -85,7 +88,7 @@ fn range_bounds( end: &FrameBound, pos: usize, len: usize, - order_values: &[serde_json::Value], + order_values: &[Value], ) -> (usize, usize) { let current_val = order_values.get(pos).and_then(as_f64); @@ -98,7 +101,7 @@ fn range_bound_to_idx( bound: &FrameBound, pos: usize, len: usize, - order_values: &[serde_json::Value], + order_values: &[Value], current_val: Option, is_start: bool, ) -> usize { @@ -108,17 +111,22 @@ fn range_bound_to_idx( FrameBound::CurrentRow => { // Peer-aware: for start bound, go back to first peer; // for end bound, advance to last peer. + let peers = |other: usize| match (order_values.get(other), order_values.get(pos)) { + (Some(a), Some(b)) => sort_peers(a, b), + (None, None) => true, + _ => false, + }; if is_start { - // Scan backward to find the first row with the same value. + // Scan backward to find the first peer of the current row. let mut idx = pos; - while idx > 0 && order_values.get(idx - 1) == order_values.get(pos) { + while idx > 0 && peers(idx - 1) { idx -= 1; } idx } else { - // Scan forward to find the last row with the same value. + // Scan forward to find the last peer of the current row. let mut idx = pos; - while idx + 1 < len && order_values.get(idx + 1) == order_values.get(pos) { + while idx + 1 < len && peers(idx + 1) { idx += 1; } idx @@ -208,7 +216,6 @@ fn groups_bound_to_group( mod tests { use super::*; use crate::window::spec::{FrameBound, WindowFrame}; - use serde_json::json; fn range_frame(start: FrameBound, end: FrameBound) -> WindowFrame { WindowFrame { @@ -234,8 +241,8 @@ mod tests { } } - fn num_vals(ns: &[i64]) -> Vec { - ns.iter().map(|&n| json!(n)).collect() + fn num_vals(ns: &[i64]) -> Vec { + ns.iter().map(|&n| Value::Integer(n)).collect() } // ROWS ─────────────────────────────────────────────────────────────────── diff --git a/nodedb-query/src/window/helpers.rs b/nodedb-query/src/window/helpers.rs index 487d30c4f..3e2559d02 100644 --- a/nodedb-query/src/window/helpers.rs +++ b/nodedb-query/src/window/helpers.rs @@ -1,10 +1,16 @@ // SPDX-License-Identifier: Apache-2.0 -//! Shared helpers for window-function evaluation. +//! Shared helpers for window-function evaluation over document rows. +//! +//! A row is `(id, Value::Object)`. `Value` holds every result as computed, +//! NaN and ±Infinity included, which a JSON number cannot. use std::collections::HashMap; +use nodedb_types::Value; + use crate::expr::types::SqlExpr; +use crate::value_ops::{compare_sort_values, sort_peers, value_to_display_string}; /// Group row indices by partition key, preserving first-seen partition order, /// then sort each partition's indices by the spec's ORDER BY. @@ -16,7 +22,7 @@ use crate::expr::types::SqlExpr; /// propagates as `Err(EvalError::DivisionByZero)` rather than being folded to /// NULL. pub(super) fn build_partitions( - rows: &[(String, serde_json::Value)], + rows: &[(String, Value)], partition_by: &[SqlExpr], order_by: &[(SqlExpr, bool)], ) -> Result>, crate::expr::EvalError> { @@ -29,7 +35,7 @@ pub(super) fn build_partitions( for (i, (_id, doc)) in rows.iter().enumerate() { let key: String = partition_by .iter() - .map(|expr| eval_expr_on_json(expr, doc).map(|v| v.to_string())) + .map(|expr| expr.eval(doc).map(|v| partition_key_part(&v))) .collect::, _>>()? .join("\x00"); let entry = groups.entry(key.clone()).or_default(); @@ -43,12 +49,12 @@ pub(super) fn build_partitions( }; if !order_by.is_empty() { - let mut keys: Vec> = Vec::with_capacity(rows.len()); + let mut keys: Vec> = Vec::with_capacity(rows.len()); for (_id, doc) in rows.iter() { keys.push( order_by .iter() - .map(|(expr, _)| eval_expr_on_json(expr, doc)) + .map(|(expr, _)| expr.eval(doc)) .collect::, _>>()?, ); } @@ -61,6 +67,12 @@ pub(super) fn build_partitions( Ok(partitions) } +/// One PARTITION BY value as a key fragment. The type name keeps the text +/// `"1"` apart from the integer `1`. +fn partition_key_part(v: &Value) -> String { + format!("{}:{}", v.type_name(), value_to_display_string(v)) +} + /// Decide NULL placement for one ORDER BY column, shared by every window /// evaluator's `compare_order_keys`. /// @@ -94,8 +106,8 @@ pub(super) fn null_order( /// Compare two rows' pre-evaluated ORDER BY keys. fn compare_order_keys( - a: &[serde_json::Value], - b: &[serde_json::Value], + a: &[Value], + b: &[Value], order_by: &[(SqlExpr, bool)], ) -> std::cmp::Ordering { use std::cmp::Ordering; @@ -104,7 +116,7 @@ fn compare_order_keys( continue; }; let ord = null_order(va.is_null(), vb.is_null(), *ascending).unwrap_or_else(|| { - let c = crate::json_expr::compare_json(va, vb); + let c = compare_sort_values(va, vb); if *ascending { c } else { c.reverse() } }); if ord != Ordering::Equal { @@ -114,46 +126,30 @@ fn compare_order_keys( Ordering::Equal } -pub(super) fn set_window_col(row: &mut serde_json::Value, alias: &str, val: serde_json::Value) { - if let serde_json::Value::Object(map) = row { +pub(super) fn set_window_col(row: &mut Value, alias: &str, val: Value) { + if let Value::Object(map) = row { map.insert(alias.to_string(), val); } } -/// Evaluate a `SqlExpr` against a serde_json document, returning a serde_json value. -/// -/// A division/modulo-by-zero in a PARTITION BY / ORDER BY expression is -/// surfaced as `Err(EvalError::DivisionByZero)` — the same -/// statement-failure treatment WHERE/projection expressions get — rather than -/// being folded to `NULL`. The `Result` is threaded through every window -/// function that reaches this evaluator up to `evaluate_window_functions`. -pub(super) fn eval_expr_on_json( - expr: &SqlExpr, - doc: &serde_json::Value, -) -> Result { - crate::json_expr::eval_expr_on_json(expr, doc) -} - -pub(super) fn as_f64(v: &serde_json::Value) -> Option { - match v { - serde_json::Value::Number(n) => n.as_f64(), - serde_json::Value::String(s) => s.parse().ok(), - _ => None, - } +/// Numeric view of a value for RANGE offsets and MIN / MAX filtering: a +/// number or numeric text. `None` for any other value. +pub(super) fn as_f64(v: &Value) -> Option { + crate::value_ops::value_to_f64(v, false) } /// Returns true when row at index `b` has the same ORDER BY key as row at /// index `a` (used by peer-aware ranking like RANK and PERCENT_RANK). pub(super) fn order_keys_equal( - rows: &[(String, serde_json::Value)], + rows: &[(String, Value)], a: usize, b: usize, order_by: &[(SqlExpr, bool)], ) -> Result { for (expr, _) in order_by { - let va = eval_expr_on_json(expr, &rows[a].1)?; - let vb = eval_expr_on_json(expr, &rows[b].1)?; - if va != vb { + let va = expr.eval(&rows[a].1)?; + let vb = expr.eval(&rows[b].1)?; + if !sort_peers(&va, &vb) { return Ok(false); } } diff --git a/nodedb-query/src/window/mod.rs b/nodedb-query/src/window/mod.rs index f1e4f5dda..9ffc30cf4 100644 --- a/nodedb-query/src/window/mod.rs +++ b/nodedb-query/src/window/mod.rs @@ -8,6 +8,7 @@ pub mod aggregate; pub mod arg; pub mod eval; +pub mod extremum; pub mod frame; pub mod helpers; pub mod offset; diff --git a/nodedb-query/src/window/offset.rs b/nodedb-query/src/window/offset.rs index 9d945a606..76c51166b 100644 --- a/nodedb-query/src/window/offset.rs +++ b/nodedb-query/src/window/offset.rs @@ -2,12 +2,14 @@ //! Offset window functions: lag, lead, nth_value. +use nodedb_types::Value; + use super::arg::{arg_at, const_default_arg, const_usize_arg, eval_arg_values}; use super::helpers::set_window_col; use super::spec::WindowFuncSpec; pub(super) fn apply_lag( - rows: &mut [(String, serde_json::Value)], + rows: &mut [(String, Value)], indices: &[usize], spec: &WindowFuncSpec, ) -> Result<(), crate::expr::EvalError> { @@ -27,7 +29,7 @@ pub(super) fn apply_lag( } pub(super) fn apply_lead( - rows: &mut [(String, serde_json::Value)], + rows: &mut [(String, Value)], indices: &[usize], spec: &WindowFuncSpec, ) -> Result<(), crate::expr::EvalError> { @@ -52,7 +54,7 @@ pub(super) fn apply_lead( /// rows of each partition return NULL and rows from the n'th onward return /// the value of `expr` at the n'th row. pub(super) fn apply_nth_value( - rows: &mut [(String, serde_json::Value)], + rows: &mut [(String, Value)], indices: &[usize], spec: &WindowFuncSpec, ) -> Result<(), crate::expr::EvalError> { @@ -63,7 +65,7 @@ pub(super) fn apply_nth_value( let val = if pos + 1 >= n { arg_at(&arg_values, n - 1) } else { - serde_json::Value::Null + Value::Null }; set_window_col(&mut rows[i].1, &spec.alias, val); } diff --git a/nodedb-query/src/window/ranking.rs b/nodedb-query/src/window/ranking.rs index fd0555ef8..3bc053d58 100644 --- a/nodedb-query/src/window/ranking.rs +++ b/nodedb-query/src/window/ranking.rs @@ -3,23 +3,26 @@ //! Ranking and distribution window functions: row_number, rank, dense_rank, //! ntile, percent_rank, cume_dist. +use nodedb_types::Value; + use crate::expr::SqlExpr; use super::helpers::{order_keys_equal, set_window_col}; use super::spec::WindowFuncSpec; -pub(super) fn apply_row_number( - rows: &mut [(String, serde_json::Value)], - indices: &[usize], - alias: &str, -) { +/// A row count or rank as an integer value. +fn count_value(n: usize) -> Value { + Value::from_u64(n as u64) +} + +pub(super) fn apply_row_number(rows: &mut [(String, Value)], indices: &[usize], alias: &str) { for (rank, &i) in indices.iter().enumerate() { - set_window_col(&mut rows[i].1, alias, serde_json::json!(rank + 1)); + set_window_col(&mut rows[i].1, alias, count_value(rank + 1)); } } pub(super) fn apply_rank( - rows: &mut [(String, serde_json::Value)], + rows: &mut [(String, Value)], indices: &[usize], alias: &str, order_by: &[(SqlExpr, bool)], @@ -28,23 +31,19 @@ pub(super) fn apply_rank( return Ok(()); } let mut current_rank = 1; - set_window_col(&mut rows[indices[0]].1, alias, serde_json::json!(1)); + set_window_col(&mut rows[indices[0]].1, alias, count_value(1)); for pos in 1..indices.len() { if !order_keys_equal(rows, indices[pos - 1], indices[pos], order_by)? { current_rank = pos + 1; } - set_window_col( - &mut rows[indices[pos]].1, - alias, - serde_json::json!(current_rank), - ); + set_window_col(&mut rows[indices[pos]].1, alias, count_value(current_rank)); } Ok(()) } pub(super) fn apply_dense_rank( - rows: &mut [(String, serde_json::Value)], + rows: &mut [(String, Value)], indices: &[usize], alias: &str, order_by: &[(SqlExpr, bool)], @@ -53,26 +52,18 @@ pub(super) fn apply_dense_rank( return Ok(()); } let mut current_rank = 1; - set_window_col(&mut rows[indices[0]].1, alias, serde_json::json!(1)); + set_window_col(&mut rows[indices[0]].1, alias, count_value(1)); for pos in 1..indices.len() { if !order_keys_equal(rows, indices[pos - 1], indices[pos], order_by)? { current_rank += 1; } - set_window_col( - &mut rows[indices[pos]].1, - alias, - serde_json::json!(current_rank), - ); + set_window_col(&mut rows[indices[pos]].1, alias, count_value(current_rank)); } Ok(()) } -pub(super) fn apply_ntile( - rows: &mut [(String, serde_json::Value)], - indices: &[usize], - spec: &WindowFuncSpec, -) { +pub(super) fn apply_ntile(rows: &mut [(String, Value)], indices: &[usize], spec: &WindowFuncSpec) { let n = spec .args .first() @@ -92,14 +83,14 @@ pub(super) fn apply_ntile( for (pos, &i) in indices.iter().enumerate() { // Integer division distributes rows as evenly as possible (PostgreSQL semantics). let bucket = (pos * n / total) + 1; - set_window_col(&mut rows[i].1, &spec.alias, serde_json::json!(bucket)); + set_window_col(&mut rows[i].1, &spec.alias, count_value(bucket)); } } /// PostgreSQL `percent_rank()` — `(rank - 1) / (partition_rows - 1)`. Single- /// row partitions return 0. Peer rows share their leader's value. pub(super) fn apply_percent_rank( - rows: &mut [(String, serde_json::Value)], + rows: &mut [(String, Value)], indices: &[usize], alias: &str, order_by: &[(SqlExpr, bool)], @@ -109,19 +100,19 @@ pub(super) fn apply_percent_rank( return Ok(()); } if total == 1 { - set_window_col(&mut rows[indices[0]].1, alias, serde_json::json!(0.0)); + set_window_col(&mut rows[indices[0]].1, alias, Value::Float(0.0)); return Ok(()); } let denom = (total - 1) as f64; let mut current_rank = 1usize; - set_window_col(&mut rows[indices[0]].1, alias, serde_json::json!(0.0)); + set_window_col(&mut rows[indices[0]].1, alias, Value::Float(0.0)); for pos in 1..total { if !order_keys_equal(rows, indices[pos - 1], indices[pos], order_by)? { current_rank = pos + 1; } let pr = (current_rank - 1) as f64 / denom; - set_window_col(&mut rows[indices[pos]].1, alias, serde_json::json!(pr)); + set_window_col(&mut rows[indices[pos]].1, alias, Value::Float(pr)); } Ok(()) } @@ -130,7 +121,7 @@ pub(super) fn apply_percent_rank( /// Peer rows (equal ORDER BY keys) share the same value, taken from the last /// peer's position. pub(super) fn apply_cume_dist( - rows: &mut [(String, serde_json::Value)], + rows: &mut [(String, Value)], indices: &[usize], alias: &str, order_by: &[(SqlExpr, bool)], @@ -151,7 +142,7 @@ pub(super) fn apply_cume_dist( } let cd = group_end as f64 / denom; for pos in group_start..group_end { - set_window_col(&mut rows[indices[pos]].1, alias, serde_json::json!(cd)); + set_window_col(&mut rows[indices[pos]].1, alias, Value::Float(cd)); } group_start = group_end; } diff --git a/nodedb-query/src/window/running.rs b/nodedb-query/src/window/running.rs index 40cdb9feb..b22dbfad3 100644 --- a/nodedb-query/src/window/running.rs +++ b/nodedb-query/src/window/running.rs @@ -11,9 +11,13 @@ //! computed at the *last* peer in the group (i.e., they all include each //! other). This matches PostgreSQL behaviour for RANGE CURRENT ROW. +use nodedb_types::Value; + use super::arg::{ArgValues, arg_at}; -use super::helpers::{as_f64, order_keys_equal, set_window_col}; +use super::extremum::{extremum_direction, value_replaces}; +use super::helpers::{order_keys_equal, set_window_col}; use super::spec::WindowFuncSpec; +use crate::numeric_sum::{ExactSum, sum_input}; /// Apply a peer-aware running aggregate over a sorted partition. /// @@ -21,7 +25,7 @@ use super::spec::WindowFuncSpec; /// `arg_values` holds the function argument already evaluated once per /// partition position; `None` is the no-argument form (`COUNT(*)`). pub(super) fn running_aggregate( - rows: &mut [(String, serde_json::Value)], + rows: &mut [(String, Value)], indices: &[usize], spec: &WindowFuncSpec, arg_values: &ArgValues, @@ -34,10 +38,12 @@ pub(super) fn running_aggregate( // Accumulate state incrementally row-by-row, but defer writing results // until the end of each peer group (so all peers see the group's final // value). We track where the current peer group started. - let mut running_sum = 0.0f64; + // SUM / AVG total exactly per `ExactSum`. + let mut running_sum = ExactSum::new(); let mut running_count = 0u64; - let mut running_min: Option = None; - let mut running_max: Option = None; + // MIN / MAX hold the original argument value, compared exactly. + let extremum_dir = extremum_direction(&spec.func_name); + let mut running_extreme: Option = None; // Indices of rows belonging to the *current* peer group (deferred write). let mut peer_start = 0usize; @@ -45,11 +51,13 @@ pub(super) fn running_aggregate( for pos in 0..len { let i = indices[pos]; let val = arg_at(arg_values, pos); - if let Some(n) = as_f64(&val) { - running_sum += n; + if sum_input(&val).is_some_and(|n| running_sum.add_value(&n)) { running_count += 1; - running_min = Some(running_min.map_or(n, |m: f64| m.min(n))); - running_max = Some(running_max.map_or(n, |m: f64| m.max(n))); + if let Some(want_max) = extremum_dir + && value_replaces(&val, running_extreme.as_ref(), want_max) + { + running_extreme = Some(val); + } } else if spec.func_name == "count" && (arg_values.is_none() || !val.is_null()) { // `COUNT(*)` counts every row; `COUNT(expr)` counts rows whose // argument is non-NULL, including non-numeric values. @@ -63,24 +71,13 @@ pub(super) fn running_aggregate( if is_last_in_group { // Compute the result at the end of this peer group. let result = match spec.func_name.as_str() { - "sum" => serde_json::json!(running_sum), - "count" => serde_json::json!(running_count), - "avg" => { - if running_count > 0 { - serde_json::json!(running_sum / running_count as f64) - } else { - serde_json::Value::Null - } - } - "min" => running_min - .map(|m| serde_json::json!(m)) - .unwrap_or(serde_json::Value::Null), - "max" => running_max - .map(|m| serde_json::json!(m)) - .unwrap_or(serde_json::Value::Null), + "sum" => running_sum.sum()?, + "count" => Value::from_u64(running_count), + "avg" => running_sum.avg()?, + "min" | "max" => running_extreme.clone().unwrap_or(Value::Null), "first_value" => arg_at(arg_values, 0), "last_value" => arg_at(arg_values, pos), - _ => serde_json::Value::Null, + _ => Value::Null, }; // Write the *same* result to every row in the peer group. diff --git a/nodedb-query/src/window/value_agg.rs b/nodedb-query/src/window/value_agg.rs index 616c503bb..b75b010f8 100644 --- a/nodedb-query/src/window/value_agg.rs +++ b/nodedb-query/src/window/value_agg.rs @@ -7,9 +7,10 @@ use std::collections::HashMap; use nodedb_types::Value; +use super::extremum::{extremum_direction, value_replaces}; use super::spec::{FrameBound, WindowFrame, WindowFuncSpec}; use super::value_eval::{WindowError, cmp_values, eval_arg_for_row, order_keys_equal_v, set_cell}; -use crate::simd_agg; +use crate::numeric_sum::ExactSum; pub(super) fn apply_v_aggregate( rows: &mut [Vec], @@ -52,10 +53,12 @@ fn apply_v_running_aggregate( return Ok(()); } - let mut running_sum = 0.0f64; + // SUM / AVG total exactly per `ExactSum`. + let mut running_sum = ExactSum::new(); let mut running_count = 0u64; - let mut running_min: Option = None; - let mut running_max: Option = None; + // MIN / MAX hold the original argument value, compared exactly. + let extremum_dir = extremum_direction(&spec.func_name); + let mut running_extreme: Option = None; let mut peer_start = 0usize; for pos in 0..len { @@ -65,12 +68,17 @@ fn apply_v_running_aggregate( None => Value::Null, }; - if let Some(n) = val.as_f64() { - running_sum += n; + if is_window_number(&val) { + running_sum.add_value(&val); running_count += 1; - running_min = Some(running_min.map_or(n, |m: f64| m.min(n))); - running_max = Some(running_max.map_or(n, |m: f64| m.max(n))); - } else if spec.func_name == "count" { + if let Some(want_max) = extremum_dir + && value_replaces(&val, running_extreme.as_ref(), want_max) + { + running_extreme = Some(val); + } + } else if spec.func_name == "count" && (spec.args.is_empty() || !val.is_null()) { + // `COUNT(*)` counts every row; `COUNT(expr)` counts rows whose + // argument is non-NULL, including non-numeric values. running_count += 1; } @@ -88,17 +96,10 @@ fn apply_v_running_aggregate( }; let result = match spec.func_name.as_str() { - "sum" => Value::Float(running_sum), + "sum" => running_sum.sum()?, "count" => Value::Integer(running_count as i64), - "avg" => { - if running_count > 0 { - Value::Float(running_sum / running_count as f64) - } else { - Value::Null - } - } - "min" => running_min.map(Value::Float).unwrap_or(Value::Null), - "max" => running_max.map(Value::Float).unwrap_or(Value::Null), + "avg" => running_sum.avg()?, + "min" | "max" => running_extreme.clone().unwrap_or(Value::Null), "first_value" => first_val, "last_value" => last_val, _ => Value::Null, @@ -140,29 +141,20 @@ fn apply_v_per_row_aggregate( Vec::new() }; - let all_vals: Vec> = indices + let arg_vals: Vec = indices .iter() .map(|&i| match rows.get(i) { - Some(row) => Ok(eval_arg(spec, row, column_index)?.as_f64()), - None => Ok(None), + Some(row) => eval_arg(spec, row, column_index), + None => Ok(Value::Null), }) .collect::, WindowError>>()?; - let results: Vec = (0..len) .map(|pos| { let (start_idx, end_idx) = evaluate_v_frame_bounds(&spec.frame, pos, len, &order_values, &peer_groups); - aggregate_v_slice( - &all_vals, - indices, - rows, - column_index, - spec, - start_idx, - end_idx, - ) + aggregate_v_slice(&arg_vals, spec, start_idx, end_idx) }) - .collect::, WindowError>>()?; + .collect::>()?; for (pos, result) in results.into_iter().enumerate() { set_cell(rows, indices[pos], write_col, result); @@ -170,74 +162,57 @@ fn apply_v_per_row_aggregate( Ok(()) } +/// A window SUM / AVG / MIN / MAX input: an integer, float, or decimal. +fn is_window_number(v: &Value) -> bool { + matches!(v, Value::Integer(_) | Value::Float(_) | Value::Decimal(_)) +} + +/// Aggregate the evaluated argument over the frame slice +/// `[start_idx, end_idx]` (partition positions). `arg_vals` holds the +/// argument per partition position. fn aggregate_v_slice( - all_vals: &[Option], - indices: &[usize], - rows: &[Vec], - column_index: &HashMap, + arg_vals: &[Value], spec: &WindowFuncSpec, start_idx: usize, end_idx: usize, ) -> Result { - let slice_vals: Vec = all_vals[start_idx..=end_idx] - .iter() - .filter_map(|v| *v) - .collect(); + let frame_values = arg_vals.get(start_idx..=end_idx).unwrap_or(&[]); let slice_count = end_idx - start_idx + 1; - - let result = match spec.func_name.as_str() { - "sum" => { - let rt = simd_agg::ts_runtime(); - Value::Float((rt.sum_f64)(&slice_vals)) + let frame_sum = || { + let mut acc = ExactSum::new(); + for v in frame_values { + acc.add_value(v); } - "count" => Value::Integer(slice_count as i64), - "avg" => { - if slice_vals.is_empty() { - Value::Null - } else { - let rt = simd_agg::ts_runtime(); - Value::Float((rt.sum_f64)(&slice_vals) / slice_vals.len() as f64) - } - } - "min" => { - if slice_vals.is_empty() { - Value::Null + acc + }; + + Ok(match spec.func_name.as_str() { + "sum" => frame_sum().sum()?, + // `COUNT(*)` counts frame rows; `COUNT(expr)` counts the rows whose + // argument is non-NULL. + "count" => { + let counted = if spec.args.is_empty() { + slice_count } else { - let rt = simd_agg::ts_runtime(); - Value::Float((rt.min_f64)(&slice_vals)) - } + frame_values.iter().filter(|v| !v.is_null()).count() + }; + Value::Integer(counted as i64) } - "max" => { - if slice_vals.is_empty() { - Value::Null - } else { - let rt = simd_agg::ts_runtime(); - Value::Float((rt.max_f64)(&slice_vals)) + "avg" => frame_sum().avg()?, + "min" | "max" => { + let want_max = spec.func_name == "max"; + let mut best: Option<&Value> = None; + for candidate in frame_values.iter().filter(|v| is_window_number(v)) { + if value_replaces(candidate, best, want_max) { + best = Some(candidate); + } } + best.cloned().unwrap_or(Value::Null) } - "first_value" => match indices.get(start_idx).and_then(|&i| rows.get(i)) { - Some(row) => eval_arg_for_row( - spec.args - .first() - .unwrap_or(&crate::expr::types::SqlExpr::Literal(Value::Null)), - row, - column_index, - )?, - None => Value::Null, - }, - "last_value" => match indices.get(end_idx).and_then(|&i| rows.get(i)) { - Some(row) => eval_arg_for_row( - spec.args - .first() - .unwrap_or(&crate::expr::types::SqlExpr::Literal(Value::Null)), - row, - column_index, - )?, - None => Value::Null, - }, + "first_value" => arg_vals.get(start_idx).cloned().unwrap_or(Value::Null), + "last_value" => arg_vals.get(end_idx).cloned().unwrap_or(Value::Null), _ => Value::Null, - }; - Ok(result) + }) } fn build_v_peer_groups(order_values: &[Value]) -> Vec { @@ -545,6 +520,267 @@ mod tests { assert!((rows[0][1].as_f64().unwrap() - 3.0).abs() < 1e-9); } + fn whole_partition() -> WindowFrame { + frame( + "rows", + FrameBound::UnboundedPreceding, + FrameBound::UnboundedFollowing, + ) + } + + #[test] + fn frame_min_max_keep_integers_above_2_pow_53_exact() { + let cols = ci(&["v"]); + let vals = [9_007_199_254_740_993, 9_007_199_254_740_992]; + + let mut rows = rows_v(&vals); + run_agg( + &mut rows, + &cols, + &agg_spec("min", whole_partition(), vec![]), + ); + assert_eq!(rows[0][1], Value::Integer(9_007_199_254_740_992)); + + let mut rows = rows_v(&vals); + run_agg( + &mut rows, + &cols, + &agg_spec("max", whole_partition(), vec![]), + ); + assert_eq!(rows[0][1], Value::Integer(9_007_199_254_740_993)); + } + + #[test] + fn sum_keeps_integers_above_2_pow_53_exact() { + let cols = ci(&["v"]); + let vals = [9_007_199_254_740_993, 9_007_199_254_740_992]; + + let mut rows = rows_v(&vals); + run_agg( + &mut rows, + &cols, + &agg_spec("sum", whole_partition(), vec![]), + ); + assert_eq!(rows[0][1], Value::Integer(18_014_398_509_481_985)); + + let mut rows = rows_v(&[9_007_199_254_740_993, 1]); + let running = agg_spec("sum", WindowFrame::default(), vec![(col("v"), false)]); + run_agg(&mut rows, &cols, &running); + assert_eq!(rows[0][1], Value::Integer(9_007_199_254_740_993)); + assert_eq!(rows[1][1], Value::Integer(9_007_199_254_740_994)); + + let mut rows = rows_v(&vals); + run_agg( + &mut rows, + &cols, + &agg_spec("avg", whole_partition(), vec![]), + ); + assert_eq!(rows[0][1], Value::Float(9_007_199_254_740_992.0)); + } + + #[test] + fn sum_past_i64_is_decimal_and_u64_counts() { + let cols = ci(&["v"]); + let mut rows = vec![ + vec![Value::Integer(i64::MAX)], + vec![Value::from_u64(u64::MAX)], + ]; + run_agg( + &mut rows, + &cols, + &agg_spec("sum", whole_partition(), vec![]), + ); + assert_eq!( + rows[0][1], + Value::Decimal(rust_decimal::Decimal::from_i128_with_scale( + i128::from(i64::MAX) + i128::from(u64::MAX), + 0 + )) + ); + + let mut rows = vec![ + vec![Value::Integer(i64::MAX)], + vec![Value::from_u64(u64::MAX)], + ]; + run_agg( + &mut rows, + &cols, + &agg_spec("max", whole_partition(), vec![]), + ); + assert_eq!(rows[0][1], Value::from_u64(u64::MAX)); + } + + #[test] + fn sum_mixed_int_float_is_float_and_nan_is_the_largest_extreme() { + let cols = ci(&["v"]); + let mixed = || vec![vec![Value::Integer(2)], vec![Value::Float(0.5)]]; + let mut rows = mixed(); + run_agg( + &mut rows, + &cols, + &agg_spec("sum", whole_partition(), vec![]), + ); + assert_eq!(rows[0][1], Value::Float(2.5)); + + let with_nan = || { + vec![ + vec![Value::Float(f64::NAN)], + vec![Value::Integer(4)], + vec![Value::Integer(9)], + ] + }; + let mut rows = with_nan(); + run_agg( + &mut rows, + &cols, + &agg_spec("min", whole_partition(), vec![]), + ); + assert_eq!(rows[0][1], Value::Integer(4)); + // NaN sorts above every number, as in PostgreSQL. + let mut rows = with_nan(); + run_agg( + &mut rows, + &cols, + &agg_spec("max", whole_partition(), vec![]), + ); + assert!(matches!(rows[0][1], Value::Float(f) if f.is_nan())); + } + + #[test] + fn frame_min_max_over_nanosecond_timestamps() { + let cols = ci(&["v"]); + let ts = [ + 1_700_000_000_000_000_001, + 1_700_000_000_000_000_003, + 1_700_000_000_000_000_002, + ]; + + let mut rows = rows_v(&ts); + run_agg( + &mut rows, + &cols, + &agg_spec("min", whole_partition(), vec![]), + ); + assert_eq!(rows[1][1], Value::Integer(1_700_000_000_000_000_001)); + + let mut rows = rows_v(&ts); + run_agg( + &mut rows, + &cols, + &agg_spec("max", whole_partition(), vec![]), + ); + assert_eq!(rows[1][1], Value::Integer(1_700_000_000_000_000_003)); + } + + #[test] + fn frame_min_max_mixed_int_float_return_original_type() { + let cols = ci(&["v"]); + let mixed = || { + vec![ + vec![Value::Integer(9_007_199_254_740_993)], + vec![Value::Float(9_007_199_254_740_992.0)], + vec![Value::Float(1.5)], + vec![Value::Integer(2)], + ] + }; + + let mut rows = mixed(); + run_agg( + &mut rows, + &cols, + &agg_spec("min", whole_partition(), vec![]), + ); + assert_eq!(rows[0][1], Value::Float(1.5)); + + let mut rows = mixed(); + run_agg( + &mut rows, + &cols, + &agg_spec("max", whole_partition(), vec![]), + ); + assert_eq!(rows[0][1], Value::Integer(9_007_199_254_740_993)); + } + + #[test] + fn running_min_max_keep_integers_exact() { + // Default frame, distinct order keys `k`: running extreme per row. + let cols = ci(&["v", "k"]); + let keyed = || { + vec![ + vec![Value::Integer(9_007_199_254_740_993), Value::Integer(1)], + vec![Value::Integer(9_007_199_254_740_992), Value::Integer(2)], + vec![Value::Integer(9_007_199_254_740_995), Value::Integer(3)], + ] + }; + let order = vec![(col("k"), true)]; + + let mut rows = keyed(); + run_agg( + &mut rows, + &cols, + &agg_spec("min", WindowFrame::default(), order.clone()), + ); + let mins: Vec = rows.iter().map(|r| r[2].clone()).collect(); + assert_eq!( + mins, + vec![ + Value::Integer(9_007_199_254_740_993), + Value::Integer(9_007_199_254_740_992), + Value::Integer(9_007_199_254_740_992), + ] + ); + + let mut rows = keyed(); + run_agg( + &mut rows, + &cols, + &agg_spec("max", WindowFrame::default(), order), + ); + let maxes: Vec = rows.iter().map(|r| r[2].clone()).collect(); + assert_eq!( + maxes, + vec![ + Value::Integer(9_007_199_254_740_993), + Value::Integer(9_007_199_254_740_993), + Value::Integer(9_007_199_254_740_995), + ] + ); + } + + #[test] + fn count_expr_skips_null_arguments() { + let cols = ci(&["v"]); + let rows_with_null = || { + vec![ + vec![Value::Integer(1)], + vec![Value::Null], + vec![Value::String("x".into())], + ] + }; + + let mut rows = rows_with_null(); + run_agg( + &mut rows, + &cols, + &agg_spec("count", whole_partition(), vec![]), + ); + assert_eq!(rows[0][1], Value::Integer(2)); + + let mut rows = rows_with_null(); + let mut star = agg_spec("count", whole_partition(), vec![]); + star.args.clear(); + run_agg(&mut rows, &cols, &star); + assert_eq!(rows[0][1], Value::Integer(3)); + + let mut rows = rows_with_null(); + run_agg( + &mut rows, + &cols, + &agg_spec("count", WindowFrame::default(), vec![]), + ); + assert_eq!(rows[2][1], Value::Integer(2)); + } + #[test] fn first_and_last_value() { let cols = ci(&["v"]); diff --git a/nodedb-query/src/window/value_eval.rs b/nodedb-query/src/window/value_eval.rs index b4d1cbabb..efbf0ed78 100644 --- a/nodedb-query/src/window/value_eval.rs +++ b/nodedb-query/src/window/value_eval.rs @@ -126,10 +126,8 @@ pub(super) fn order_keys_equal_v( // ── Argument evaluation (pub(super) for value_agg) ──────────────────────────── -/// Evaluate a window-function argument expression against one row. -/// -/// This is the Value-native counterpart of `window::helpers::eval_expr_on_json`. -/// A division/modulo-by-zero surfaces as +/// Evaluate a window-function argument expression against one column-major +/// row. A division/modulo-by-zero surfaces as /// `Err(EvalError::DivisionByZero)` — which the value-path callers convert into /// `WindowError` via `?` — rather than being folded to `NULL`. pub(super) fn eval_arg_for_row( diff --git a/nodedb-sql/src/ddl_ast/collection_type/kv.rs b/nodedb-sql/src/ddl_ast/collection_type/kv.rs index 0661e557c..7fb59a8a1 100644 --- a/nodedb-sql/src/ddl_ast/collection_type/kv.rs +++ b/nodedb-sql/src/ddl_ast/collection_type/kv.rs @@ -38,7 +38,8 @@ pub(crate) fn build_kv_collection_type( ColumnDef::nullable(name.clone(), column_type) } else { ColumnDef::required(name.clone(), column_type) - }; + } + .with_declared_width(&bare_type); if is_pk { col = col.with_primary_key(); } diff --git a/nodedb-sql/src/ddl_ast/collection_type/strict.rs b/nodedb-sql/src/ddl_ast/collection_type/strict.rs index 93ad55c32..4c63010a0 100644 --- a/nodedb-sql/src/ddl_ast/collection_type/strict.rs +++ b/nodedb-sql/src/ddl_ast/collection_type/strict.rs @@ -38,7 +38,8 @@ pub(crate) fn build_strict_schema( ColumnDef::nullable(name.clone(), column_type) } else { ColumnDef::required(name.clone(), column_type) - }; + } + .with_declared_width(&bare_type); if is_pk { col = col.with_primary_key(); } @@ -114,3 +115,41 @@ pub(crate) fn build_strict_schema( }) } } + +#[cfg(test)] +mod tests { + use nodedb_types::columnar::{FloatWidth, IntWidth}; + + use super::*; + + #[test] + fn schema_records_each_declared_numeric_width() { + let columns: Vec<(String, String)> = [ + ("id", "TEXT PRIMARY KEY"), + ("s", "SMALLINT NOT NULL"), + ("i", "INT"), + ("b", "BIGINT"), + ("r", "REAL"), + ("d", "DOUBLE PRECISION"), + ("t", "TEXT"), + ] + .into_iter() + .map(|(name, declared)| (name.to_string(), declared.to_string())) + .collect(); + let schema = build_strict_schema(&columns, false).expect("valid schema"); + let width = |name: &str| { + let column = schema + .columns + .iter() + .find(|c| c.name == name) + .expect("column present"); + (column.int_width, column.float_width) + }; + assert_eq!(width("s"), (Some(IntWidth::I16), None)); + assert_eq!(width("i"), (Some(IntWidth::I32), None)); + assert_eq!(width("b"), (Some(IntWidth::I64), None)); + assert_eq!(width("r"), (None, Some(FloatWidth::F32))); + assert_eq!(width("d"), (None, Some(FloatWidth::F64))); + assert_eq!(width("t"), (None, None)); + } +} diff --git a/nodedb-sql/src/ddl_ast/graph_parse/edge_predicate.rs b/nodedb-sql/src/ddl_ast/graph_parse/edge_predicate.rs new file mode 100644 index 000000000..a3ed7d8c3 --- /dev/null +++ b/nodedb-sql/src/ddl_ast/graph_parse/edge_predicate.rs @@ -0,0 +1,504 @@ +// SPDX-License-Identifier: Apache-2.0 + +//! `EDGE WHERE `: the edge-property filter of GRAPH TRAVERSE and +//! GRAPH PATH. +//! +//! The clause is last in the statement. The text before it is tokenized +//! alone, so a property named like a DSL keyword (`depth`, `label`) never +//! reads as a clause. The predicate parses as a PostgreSQL expression and +//! maps onto `MetadataFilter`: +//! +//! - `=` `<>` `!=` `>` `>=` `<` `<=` between a property name and a literal, +//! on either side. +//! - `IN (…)`, `NOT IN (…)`, `IS NULL`, `IS NOT NULL`. +//! - `AND`, `OR`, `NOT`, parentheses. Bare `TRUE` admits, bare `FALSE` +//! rejects. +//! +//! Literals: `NULL`, `TRUE`/`FALSE`, single-quoted strings, integers that +//! fit `i64`, finite floats. A property name is one identifier. Anything +//! else is a parse error. + +use nodedb_types::Value; +use nodedb_types::filter::MetadataFilter; +use sqlparser::ast::{BinaryOperator, Expr, UnaryOperator, Value as SqlValue}; +use sqlparser::dialect::PostgreSqlDialect; +use sqlparser::parser::Parser; +use sqlparser::tokenizer::Token; + +use crate::error::SqlError; +use crate::parser::normalize::normalize_ident; + +/// Split `sql` at its top-level `EDGE WHERE`. Quoted literals, quoted +/// identifiers and `{…}` object literals never split. +pub(super) fn split_edge_where(sql: &str) -> Result<(&str, Option<&str>), SqlError> { + let bytes = sql.as_bytes(); + let mut i = 0; + let mut depth = 0usize; + let mut prev: Option<(usize, usize)> = None; + while i < bytes.len() { + match bytes[i] { + b'\'' | b'"' => { + i = skip_quoted(bytes, i); + prev = None; + } + b'{' => { + depth += 1; + i += 1; + prev = None; + } + b'}' => { + depth = depth.saturating_sub(1); + i += 1; + prev = None; + } + c if c.is_ascii_alphanumeric() || c == b'_' => { + let start = i; + while i < bytes.len() && (bytes[i].is_ascii_alphanumeric() || bytes[i] == b'_') { + i += 1; + } + if depth == 0 { + let after_edge = + prev.is_some_and(|(s, e)| sql[s..e].eq_ignore_ascii_case("EDGE")); + if after_edge && sql[start..i].eq_ignore_ascii_case("WHERE") { + let head = &sql[..prev.map_or(start, |(s, _)| s)]; + let predicate = sql[i..].trim().trim_end_matches(';').trim_end(); + if predicate.is_empty() { + return Err(parse_err("EDGE WHERE requires a predicate".to_owned())); + } + return Ok((head, Some(predicate))); + } + prev = Some((start, i)); + } + } + c if c.is_ascii_whitespace() => i += 1, + _ => { + i += 1; + prev = None; + } + } + } + Ok((sql, None)) +} + +/// The index just past the quoted run opening at `start`. A doubled quote +/// is an escaped quote inside the run. +fn skip_quoted(bytes: &[u8], start: usize) -> usize { + let quote = bytes[start]; + let mut j = start + 1; + while j < bytes.len() { + if bytes[j] == quote { + if bytes.get(j + 1) == Some("e) { + j += 2; + continue; + } + return j + 1; + } + j += 1; + } + j +} + +/// Parse the predicate text after `EDGE WHERE` into its AND-ed filters. +pub(super) fn parse_edge_predicate(text: &str) -> Result, SqlError> { + let dialect = PostgreSqlDialect {}; + let mut parser = Parser::new(&dialect) + .try_with_sql(text) + .map_err(|e| parse_err(format!("EDGE WHERE: {e}")))?; + let expr = parser + .parse_expr() + .map_err(|e| parse_err(format!("EDGE WHERE: {e}")))?; + parser + .expect_token(&Token::EOF) + .map_err(|e| parse_err(format!("EDGE WHERE: {e}")))?; + Ok(match convert(&expr)? { + MetadataFilter::And(children) => children, + single => vec![single], + }) +} + +fn convert(expr: &Expr) -> Result { + match expr { + Expr::Nested(inner) => convert(inner), + Expr::BinaryOp { + left, + op: BinaryOperator::And, + right, + } => { + let mut children = Vec::new(); + for side in [convert(left)?, convert(right)?] { + match side { + MetadataFilter::And(inner) => children.extend(inner), + other => children.push(other), + } + } + Ok(MetadataFilter::And(children)) + } + Expr::BinaryOp { + left, + op: BinaryOperator::Or, + right, + } => Ok(MetadataFilter::Or(vec![convert(left)?, convert(right)?])), + Expr::UnaryOp { + op: UnaryOperator::Not, + expr, + } => Ok(MetadataFilter::Not(Box::new(convert(expr)?))), + Expr::BinaryOp { left, op, right } => comparison(left, op, right), + Expr::InList { + expr, + list, + negated, + } => { + let field = field(expr)?; + let values = list.iter().map(literal).collect::, _>>()?; + Ok(if *negated { + MetadataFilter::NotIn { field, values } + } else { + MetadataFilter::In { field, values } + }) + } + Expr::IsNull(inner) => Ok(MetadataFilter::Eq { + field: field(inner)?, + value: Value::Null, + }), + Expr::IsNotNull(inner) => Ok(MetadataFilter::Ne { + field: field(inner)?, + value: Value::Null, + }), + Expr::Value(v) => match &v.value { + SqlValue::Boolean(true) => Ok(MetadataFilter::And(Vec::new())), + SqlValue::Boolean(false) => Ok(MetadataFilter::Or(Vec::new())), + _ => Err(parse_err(format!( + "EDGE WHERE term '{expr}' is not a predicate" + ))), + }, + other => Err(parse_err(format!("EDGE WHERE does not support '{other}'"))), + } +} + +fn comparison(left: &Expr, op: &BinaryOperator, right: &Expr) -> Result { + let (field, value, op) = match (field(left), field(right)) { + (Ok(_), Ok(_)) => { + return Err(parse_err(format!( + "EDGE WHERE compares a property with a literal, found '{left} {op} {right}'" + ))); + } + (Ok(f), Err(_)) => (f, literal(right)?, op.clone()), + (Err(_), Ok(f)) => (f, literal(left)?, mirror(op)), + (Err(e), Err(_)) => return Err(e), + }; + Ok(match op { + BinaryOperator::Eq => MetadataFilter::Eq { field, value }, + BinaryOperator::NotEq => MetadataFilter::Ne { field, value }, + BinaryOperator::Gt => MetadataFilter::Gt { field, value }, + BinaryOperator::GtEq => MetadataFilter::Gte { field, value }, + BinaryOperator::Lt => MetadataFilter::Lt { field, value }, + BinaryOperator::LtEq => MetadataFilter::Lte { field, value }, + other => { + return Err(parse_err(format!( + "EDGE WHERE does not support operator '{other}'" + ))); + } + }) +} + +/// The operator that keeps `literal OP field` true as `field OP' literal`. +fn mirror(op: &BinaryOperator) -> BinaryOperator { + match op { + BinaryOperator::Gt => BinaryOperator::Lt, + BinaryOperator::GtEq => BinaryOperator::LtEq, + BinaryOperator::Lt => BinaryOperator::Gt, + BinaryOperator::LtEq => BinaryOperator::GtEq, + other => other.clone(), + } +} + +fn field(expr: &Expr) -> Result { + match expr { + Expr::Identifier(ident) => Ok(normalize_ident(ident)), + Expr::Nested(inner) => field(inner), + other => Err(parse_err(format!( + "EDGE WHERE expects a property name, found '{other}': quote names with \"…\"" + ))), + } +} + +fn literal(expr: &Expr) -> Result { + match expr { + Expr::Nested(inner) => literal(inner), + Expr::Value(v) => match &v.value { + SqlValue::Null => Ok(Value::Null), + SqlValue::Boolean(b) => Ok(Value::Bool(*b)), + SqlValue::SingleQuotedString(s) => Ok(Value::String(s.clone())), + SqlValue::Number(n, _) => number(n), + other => Err(parse_err(format!( + "EDGE WHERE does not support literal '{other}'" + ))), + }, + Expr::UnaryOp { + op: UnaryOperator::Minus, + expr: inner, + } => match inner.as_ref() { + Expr::Value(v) => match &v.value { + SqlValue::Number(n, _) => number(&format!("-{n}")), + other => Err(parse_err(format!("EDGE WHERE cannot negate '{other}'"))), + }, + other => Err(parse_err(format!("EDGE WHERE cannot negate '{other}'"))), + }, + other => Err(parse_err(format!( + "EDGE WHERE expects a literal, found '{other}'" + ))), + } +} + +/// An integer when `text` has no `.` or exponent and fits `i64`, else a +/// finite float. +fn number(text: &str) -> Result { + if !text.contains(['.', 'e', 'E']) + && let Ok(i) = text.parse::() + { + return Ok(Value::Integer(i)); + } + text.parse::() + .ok() + .filter(|f| f.is_finite()) + .map(Value::Float) + .ok_or_else(|| parse_err(format!("EDGE WHERE number '{text}' is not a finite number"))) +} + +fn parse_err(detail: String) -> SqlError { + SqlError::Parse { detail } +} + +#[cfg(test)] +mod tests { + use super::super::entry::try_parse; + use super::*; + use crate::ddl_ast::statement::{GraphStmt, NodedbStatement}; + + fn predicate_of(sql: &str) -> Vec { + match try_parse(sql) + .expect("graph DSL") + .expect("well-formed statement") + { + NodedbStatement::Graph(GraphStmt::GraphTraverse { edge_predicate, .. }) + | NodedbStatement::Graph(GraphStmt::GraphPath { edge_predicate, .. }) => edge_predicate, + other => panic!("expected TRAVERSE or PATH, got {other:?}"), + } + } + + fn traverse(predicate: &str) -> Vec { + predicate_of(&format!( + "GRAPH TRAVERSE IN 'g' FROM 'a' DEPTH 2 EDGE WHERE {predicate}" + )) + } + + fn parse_error(sql: &str) -> String { + match try_parse(sql).expect("graph DSL") { + Err(SqlError::Parse { detail }) => detail, + other => panic!("expected a parse error for {sql}, got {other:?}"), + } + } + + fn gt(field: &str, value: Value) -> MetadataFilter { + MetadataFilter::Gt { + field: field.into(), + value, + } + } + + #[test] + fn two_filters_are_the_and_ed_list() { + assert_eq!( + traverse("score > 5 AND active = TRUE"), + vec![ + gt("score", Value::Integer(5)), + MetadataFilter::Eq { + field: "active".into(), + value: Value::Bool(true), + }, + ] + ); + } + + #[test] + fn a_literal_on_the_left_mirrors_the_operator() { + assert_eq!(traverse("5 < score"), vec![gt("score", Value::Integer(5))]); + assert_eq!( + traverse("5 = score"), + vec![MetadataFilter::Eq { + field: "score".into(), + value: Value::Integer(5), + }] + ); + } + + #[test] + fn in_not_in_and_null_tests() { + assert_eq!( + traverse("kind IN ('road', 'rail') AND tag NOT IN (1, 2)"), + vec![ + MetadataFilter::In { + field: "kind".into(), + values: vec![Value::String("road".into()), Value::String("rail".into())], + }, + MetadataFilter::NotIn { + field: "tag".into(), + values: vec![Value::Integer(1), Value::Integer(2)], + }, + ] + ); + assert_eq!( + traverse("missing IS NULL AND present IS NOT NULL"), + vec![ + MetadataFilter::Eq { + field: "missing".into(), + value: Value::Null, + }, + MetadataFilter::Ne { + field: "present".into(), + value: Value::Null, + }, + ] + ); + } + + #[test] + fn or_not_and_parentheses_nest() { + assert_eq!( + traverse("NOT (closed = TRUE) AND (a = 1 OR b = 2)"), + vec![ + MetadataFilter::Not(Box::new(MetadataFilter::Eq { + field: "closed".into(), + value: Value::Bool(true), + })), + MetadataFilter::Or(vec![ + MetadataFilter::Eq { + field: "a".into(), + value: Value::Integer(1), + }, + MetadataFilter::Eq { + field: "b".into(), + value: Value::Integer(2), + }, + ]), + ] + ); + assert_eq!(traverse("TRUE"), Vec::::new()); + assert_eq!(traverse("FALSE"), vec![MetadataFilter::Or(Vec::new())]); + } + + #[test] + fn quoted_names_and_escaped_literals() { + assert_eq!( + traverse(r#""we""ird" = 'O''Reilly'"#), + vec![MetadataFilter::Eq { + field: "we\"ird".into(), + value: Value::String("O'Reilly".into()), + }] + ); + // An unquoted name folds to lower case, a quoted one keeps its case. + assert_eq!( + traverse(r#"Score > 1 AND "Score" > 2"#), + vec![ + gt("score", Value::Integer(1)), + gt("Score", Value::Integer(2)) + ] + ); + } + + #[test] + fn numbers_keep_their_kind() { + assert_eq!( + traverse("n = -9223372036854775808"), + vec![MetadataFilter::Eq { + field: "n".into(), + value: Value::Integer(i64::MIN), + }] + ); + assert_eq!(traverse("n > 1e-7"), vec![gt("n", Value::Float(1e-7))]); + assert_eq!(traverse("n > 2.5"), vec![gt("n", Value::Float(2.5))]); + } + + #[test] + fn keyword_shaped_properties_do_not_hijack_clauses() { + match try_parse( + "GRAPH TRAVERSE IN 'g' FROM 'a' DEPTH 3 LABEL 'L' EDGE WHERE depth > 7 AND label = 'x'", + ) + .expect("graph DSL") + .expect("well-formed") + { + NodedbStatement::Graph(GraphStmt::GraphTraverse { + depth, + edge_labels, + edge_predicate, + .. + }) => { + assert_eq!(depth, 3); + assert_eq!(edge_labels, vec!["L".to_string()]); + assert_eq!( + edge_predicate, + vec![ + gt("depth", Value::Integer(7)), + MetadataFilter::Eq { + field: "label".into(), + value: Value::String("x".into()), + }, + ] + ); + } + other => panic!("expected GraphTraverse, got {other:?}"), + } + } + + #[test] + fn edge_where_inside_a_quoted_node_id_does_not_split() { + match try_parse("GRAPH TRAVERSE IN 'g' FROM 'x EDGE WHERE y' DEPTH 1") + .expect("graph DSL") + .expect("well-formed") + { + NodedbStatement::Graph(GraphStmt::GraphTraverse { + start, + edge_predicate, + .. + }) => { + assert_eq!(start, "x EDGE WHERE y"); + assert!(edge_predicate.is_empty()); + } + other => panic!("expected GraphTraverse, got {other:?}"), + } + } + + #[test] + fn graph_path_takes_a_predicate() { + assert_eq!( + predicate_of("GRAPH PATH IN 'g' FROM 'a' TO 'z' MAX_DEPTH 6 EDGE WHERE score > 5;"), + vec![gt("score", Value::Integer(5))] + ); + } + + #[test] + fn unsupported_predicates_are_parse_errors() { + for predicate in [ + "score > other", + "name LIKE 'a%'", + "lower(name) = 'a'", + "score > 5 extra", + "a.b = 1", + ] { + parse_error(&format!( + "GRAPH TRAVERSE IN 'g' FROM 'a' EDGE WHERE {predicate}" + )); + } + let detail = parse_error("GRAPH TRAVERSE IN 'g' FROM 'a' EDGE WHERE "); + assert!(detail.contains("requires a predicate"), "{detail}"); + } + + #[test] + fn other_graph_statements_refuse_edge_where() { + let detail = parse_error("GRAPH NEIGHBORS IN 'g' OF 'a' EDGE WHERE score > 1"); + assert!( + detail.contains("GRAPH NEIGHBORS does not accept EDGE WHERE"), + "{detail}" + ); + } +} diff --git a/nodedb-sql/src/ddl_ast/graph_parse/entry.rs b/nodedb-sql/src/ddl_ast/graph_parse/entry.rs index 1976639ac..f49018ae3 100644 --- a/nodedb-sql/src/ddl_ast/graph_parse/entry.rs +++ b/nodedb-sql/src/ddl_ast/graph_parse/entry.rs @@ -3,7 +3,7 @@ //! Graph DSL entry point. use super::super::statement::{GraphStmt, NodedbStatement}; -use super::{tokenizer, variants}; +use super::{edge_predicate, tokenizer, variants}; use crate::error::SqlError; /// Parse a graph DSL statement. @@ -28,9 +28,36 @@ pub fn try_parse(sql: &str) -> Option> { return None; } - let toks = tokenizer::tokenize(trimmed); + Some(parse_graph(trimmed, &upper)) +} + +/// Parse a statement known to start with `GRAPH `. A trailing `EDGE WHERE` +/// predicate splits off first, so the clause text before it tokenizes alone. +fn parse_graph(trimmed: &str, upper: &str) -> Result { + let (head, predicate_text) = edge_predicate::split_edge_where(trimmed)?; + let predicate = match predicate_text { + Some(text) => edge_predicate::parse_edge_predicate(text)?, + None => Vec::new(), + }; + let toks = tokenizer::tokenize(head); + + if upper.starts_with("GRAPH TRAVERSE ") { + return variants::parse_traverse(&toks, predicate); + } + if upper.starts_with("GRAPH PATH ") { + return variants::parse_path(&toks, predicate); + } + if predicate_text.is_some() { + return Err(SqlError::Parse { + detail: format!( + "{} does not accept EDGE WHERE: only GRAPH TRAVERSE and GRAPH PATH filter on \ + edge properties", + graph_command(upper) + ), + }); + } - let parsed = if upper.starts_with("GRAPH INSERT EDGE ") { + if upper.starts_with("GRAPH INSERT EDGE ") { variants::parse_insert_edge(&toks) } else if upper.starts_with("GRAPH DELETE EDGE ") { variants::parse_delete_edge(&toks) @@ -38,12 +65,8 @@ pub fn try_parse(sql: &str) -> Option> { variants::parse_set_labels(&toks, false) } else if upper.starts_with("GRAPH UNLABEL ") { variants::parse_set_labels(&toks, true) - } else if upper.starts_with("GRAPH TRAVERSE ") { - variants::parse_traverse(&toks) } else if upper.starts_with("GRAPH NEIGHBORS ") { variants::parse_neighbors(&toks) - } else if upper.starts_with("GRAPH PATH ") { - variants::parse_path(&toks) } else if upper.starts_with("GRAPH ALGO ") { variants::parse_algo(&toks) } else if upper.starts_with("GRAPH RAG FUSION ") { @@ -54,9 +77,21 @@ pub fn try_parse(sql: &str) -> Option> { Err(SqlError::Parse { detail: "unrecognised GRAPH command".to_owned(), }) - }; + } +} - Some(parsed) +/// The command words of a `GRAPH` statement, for an error message: +/// `GRAPH NEIGHBORS`, `GRAPH INSERT EDGE`, `GRAPH RAG FUSION`. +fn graph_command(upper: &str) -> String { + const MULTI_WORD: [&str; 3] = ["GRAPH INSERT EDGE", "GRAPH DELETE EDGE", "GRAPH RAG FUSION"]; + if let Some(command) = MULTI_WORD.iter().find(|c| upper.starts_with(*c)) { + return (*command).to_owned(); + } + upper + .split_whitespace() + .take(2) + .collect::>() + .join(" ") } #[cfg(test)] @@ -174,18 +209,108 @@ mod tests { src, dst, max_depth, - edge_label, + edge_labels, + edge_predicate, }) => { + assert!(edge_predicate.is_empty()); assert_eq!(collection, "docs"); assert_eq!(src, "a"); assert_eq!(dst, "b"); assert_eq!(max_depth, 5); - assert_eq!(edge_label.as_deref(), Some("l")); + assert_eq!(edge_labels, vec!["l".to_string()]); } other => panic!("expected GraphPath, got {other:?}"), } } + fn labels_of(sql: &str) -> Vec { + match parsed(sql) { + NodedbStatement::Graph( + GraphStmt::GraphTraverse { edge_labels, .. } + | GraphStmt::GraphNeighbors { edge_labels, .. } + | GraphStmt::GraphPath { edge_labels, .. }, + ) => edge_labels, + other => panic!("expected a graph walk statement, got {other:?}"), + } + } + + fn ab() -> Vec { + vec!["a".to_string(), "b".to_string()] + } + + #[test] + fn label_list_parses_on_every_walk_statement() { + assert_eq!( + labels_of("GRAPH TRAVERSE IN 'c' FROM 'x' DEPTH 2 LABEL 'a', 'b' DIRECTION out"), + ab() + ); + assert_eq!( + labels_of("GRAPH NEIGHBORS IN 'c' OF 'x' LABEL 'a', 'b' DIRECTION both"), + ab() + ); + assert_eq!( + labels_of("GRAPH PATH IN 'c' FROM 'x' TO 'y' MAX_DEPTH 3 LABEL 'a', 'b'"), + ab() + ); + } + + #[test] + fn parenthesised_label_list_parses() { + assert_eq!( + labels_of("GRAPH TRAVERSE IN 'c' FROM 'x' LABEL ('a','b')"), + ab() + ); + assert_eq!( + labels_of("GRAPH NEIGHBORS IN 'c' OF 'x' LABEL ('a', 'b')"), + ab() + ); + assert_eq!( + labels_of("GRAPH PATH IN 'c' FROM 'x' TO 'y' LABEL ('a','b')"), + ab() + ); + } + + #[test] + fn omitted_label_clause_is_an_empty_set() { + assert!(labels_of("GRAPH TRAVERSE IN 'c' FROM 'x' DEPTH 2").is_empty()); + assert!(labels_of("GRAPH NEIGHBORS IN 'c' OF 'x'").is_empty()); + assert!(labels_of("GRAPH PATH IN 'c' FROM 'x' TO 'y'").is_empty()); + } + + #[test] + fn label_clause_without_quoted_labels_is_refused() { + for sql in [ + "GRAPH TRAVERSE IN 'c' FROM 'x' LABEL DIRECTION out", + "GRAPH TRAVERSE IN 'c' FROM 'x' LABEL knows", + "GRAPH NEIGHBORS IN 'c' OF 'x' LABEL", + "GRAPH PATH IN 'c' FROM 'x' TO 'y' LABEL knows", + ] { + let error = try_parse(sql) + .expect("graph DSL") + .expect_err("a LABEL clause with no quoted label must be refused"); + assert!( + error.to_string().contains("LABEL"), + "the error must name the LABEL clause: {error}" + ); + } + } + + #[test] + fn quoted_keyword_label_does_not_shadow_the_keyword() { + let stmt = parsed("GRAPH TRAVERSE IN 'c' FROM 'x' LABEL 'DIRECTION', 'b' DIRECTION in"); + match stmt { + NodedbStatement::Graph(GraphStmt::GraphTraverse { + edge_labels, + direction, + .. + }) => { + assert_eq!(edge_labels, vec!["DIRECTION".to_string(), "b".to_string()]); + assert_eq!(direction, GraphDirection::In); + } + other => panic!("expected GraphTraverse, got {other:?}"), + } + } + #[test] fn parse_graph_labels_list() { let stmt = parsed("GRAPH LABEL 'alice' AS 'Person', 'User'"); diff --git a/nodedb-sql/src/ddl_ast/graph_parse/helpers.rs b/nodedb-sql/src/ddl_ast/graph_parse/helpers.rs index c109438f3..14afef8af 100644 --- a/nodedb-sql/src/ddl_ast/graph_parse/helpers.rs +++ b/nodedb-sql/src/ddl_ast/graph_parse/helpers.rs @@ -20,12 +20,12 @@ pub(super) fn quoted_after(toks: &[Tok<'_>], keyword: &str) -> Option { } } -pub(super) fn quoted_list_after(toks: &[Tok<'_>], keyword: &str) -> Vec { - let Some(pos) = find_keyword(toks, keyword) else { - return Vec::new(); - }; - toks[pos + 1..] - .iter() +/// Collect the run of quoted tokens at the head of `toks`. +/// +/// The run ends at the first unquoted token. The tokenizer drops `,`, `(` +/// and `)`, so `'a', 'b'` and `('a', 'b')` give the same run. +fn quoted_run(toks: &[Tok<'_>]) -> Vec { + toks.iter() .map_while(|t| match t { Tok::Quoted(s) => Some(s.clone().into_owned()), _ => None, @@ -33,6 +33,30 @@ pub(super) fn quoted_list_after(toks: &[Tok<'_>], keyword: &str) -> Vec .collect() } +pub(super) fn quoted_list_after(toks: &[Tok<'_>], keyword: &str) -> Vec { + find_keyword(toks, keyword) + .map(|pos| quoted_run(&toks[pos + 1..])) + .unwrap_or_default() +} + +/// Read an optional `LABEL '