From 7e5c028422c49284e9cd71e3206453a852bbecb1 Mon Sep 17 00:00:00 2001 From: Farhan Syah Date: Thu, 8 Oct 2026 16:51:29 +0800 Subject: [PATCH 01/24] fix(columnar): refuse corrupt cells and segments instead of reading NULL Each column's block layout is recorded in the segment footer and the reader decodes by it, rather than inferring the layout from the codec. A memtable cell, WAL row or string block whose bytes do not decode is a typed error, never a NULL or a skipped row. Read paths that skipped an unreadable flushed segment or timeseries partition now refuse the read and file a black-box report through an optional diagnostics feature. Memtable appends and batch inserts apply whole or not at all, and a flush encodes the segment before draining the memtable, so an encode error leaves every row readable. --- Cargo.lock | 5 + nodedb-columnar/Cargo.toml | 9 + nodedb-columnar/src/compaction/segment.rs | 92 ++++- nodedb-columnar/src/delete_bitmap.rs | 20 + nodedb-columnar/src/diag/context.rs | 96 +++++ nodedb-columnar/src/diag/inert.rs | 17 + nodedb-columnar/src/diag/mod.rs | 26 ++ nodedb-columnar/src/diag/recording.rs | 78 ++++ nodedb-columnar/src/error.rs | 27 ++ nodedb-columnar/src/format.rs | 33 ++ nodedb-columnar/src/lib.rs | 10 +- .../src/materialize_rows/extract.rs | 105 ----- nodedb-columnar/src/materialize_rows/mod.rs | 6 +- nodedb-columnar/src/materialize_rows/rows.rs | 6 +- .../src/memtable/column_data/access.rs | 140 ++++++- .../src/memtable/column_data/mod.rs | 1 + .../src/memtable/column_data/push.rs | 300 +++++++------- .../src/memtable/column_data/push_ref.rs | 115 ++++++ .../src/memtable/column_data/types.rs | 18 + nodedb-columnar/src/memtable/iter.rs | 43 +- nodedb-columnar/src/memtable/mutation.rs | 142 ++++++- nodedb-columnar/src/mutation/batch.rs | 220 +++++++++++ nodedb-columnar/src/mutation/engine.rs | 61 ++- nodedb-columnar/src/mutation/mod.rs | 2 + nodedb-columnar/src/mutation/snapshot.rs | 51 ++- nodedb-columnar/src/mutation/truncate.rs | 10 +- nodedb-columnar/src/mutation/write.rs | 332 +++++++++++----- nodedb-columnar/src/reader/block_decode.rs | 186 +++++---- nodedb-columnar/src/reader/cell.rs | 355 +++++++++++++++++ nodedb-columnar/src/reader/mod.rs | 2 + nodedb-columnar/src/reader/segment_reader.rs | 8 +- nodedb-columnar/src/wal_record.rs | 238 +++++++----- nodedb-columnar/src/writer/block.rs | 38 +- nodedb-columnar/src/writer/segment_writer.rs | 5 +- nodedb-columnar/src/writer/stats.rs | 178 +++++++-- nodedb/Cargo.toml | 2 +- .../backup/restore/columnar_reissue.rs | 5 +- .../columnar_checkpoint/geometry_restore.rs | 226 +++++++---- .../data/executor/columnar_checkpoint/load.rs | 5 +- .../dispatch/meta_retention/columnar_plain.rs | 367 +++++++++++------- .../handlers/columnar_read/convert.rs | 56 +-- .../handlers/columnar_read/flushed_segment.rs | 206 ++++++++++ .../columnar_read/materialize_scan.rs | 126 +++--- .../columnar_read/materialize_scan_ts.rs | 311 +++++++-------- .../executor/handlers/columnar_read/mod.rs | 1 + .../handlers/columnar_read/scan_flushed.rs | 80 +--- .../executor/handlers/columnar_write/flush.rs | 318 ++++++++++----- .../handlers/columnar_write/read_prior.rs | 105 ++--- .../handlers/columnar_write/row_ingest.rs | 222 +++++++---- .../executor/handlers/timeseries/aggregate.rs | 27 +- .../executor/handlers/timeseries/cell_read.rs | 159 ++++++++ .../timeseries/ingest_resolved_returning.rs | 88 +++++ .../data/executor/handlers/timeseries/mod.rs | 3 + .../handlers/timeseries/partition_read.rs | 124 ++++++ .../timeseries/raw_scan/partition_scan.rs | 63 ++- .../handlers/timeseries/raw_scan/row_emit.rs | 69 ++-- .../data/executor/handlers/timeseries_wal.rs | 3 +- .../transaction/stage_write/stage_rls.rs | 13 +- nodedb/src/data/executor/wal_replay_all.rs | 1 + nodedb/src/diag/context/columnar.rs | 88 +++++ nodedb/src/diag/recording/columnar.rs | 66 ++++ .../timeseries/grouped_scan/partition.rs | 295 ++++++++------ nodedb/src/error_from_columnar.rs | 114 ++++++ 63 files changed, 4483 insertions(+), 1635 deletions(-) create mode 100644 nodedb-columnar/src/diag/context.rs create mode 100644 nodedb-columnar/src/diag/inert.rs create mode 100644 nodedb-columnar/src/diag/mod.rs create mode 100644 nodedb-columnar/src/diag/recording.rs delete mode 100644 nodedb-columnar/src/materialize_rows/extract.rs create mode 100644 nodedb-columnar/src/memtable/column_data/push_ref.rs create mode 100644 nodedb-columnar/src/mutation/batch.rs create mode 100644 nodedb-columnar/src/reader/cell.rs create mode 100644 nodedb/src/data/executor/handlers/columnar_read/flushed_segment.rs create mode 100644 nodedb/src/data/executor/handlers/timeseries/cell_read.rs create mode 100644 nodedb/src/data/executor/handlers/timeseries/ingest_resolved_returning.rs create mode 100644 nodedb/src/data/executor/handlers/timeseries/partition_read.rs create mode 100644 nodedb/src/diag/context/columnar.rs create mode 100644 nodedb/src/diag/recording/columnar.rs create mode 100644 nodedb/src/error_from_columnar.rs 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/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..c824f35e9 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,15 @@ 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 +114,15 @@ 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 +136,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/Cargo.toml b/nodedb/Cargo.toml index 97d185c73..c9cf82312 100644 --- a/nodedb/Cargo.toml +++ b/nodedb/Cargo.toml @@ -23,7 +23,7 @@ path = "src/lib.rs" [dependencies] # Internal foundation crates nodedb-types = { workspace = true } -nodedb-columnar = { workspace = true } +nodedb-columnar = { workspace = true, features = ["diagnostics"] } nodedb-bridge = { workspace = true } nodedb-physical = { workspace = true } rust_decimal = { workspace = true, features = ["db-postgres"] } diff --git a/nodedb/src/control/backup/restore/columnar_reissue.rs b/nodedb/src/control/backup/restore/columnar_reissue.rs index 323fb5105..f4fa03c80 100644 --- a/nodedb/src/control/backup/restore/columnar_reissue.rs +++ b/nodedb/src/control/backup/restore/columnar_reissue.rs @@ -63,7 +63,10 @@ pub fn decode_snapshot_live_rows( let mut surrogates: Vec = Vec::new(); // Memtable: non-deleted rows in schema column order → keyed object. - for (surrogate, values) in engine.scan_memtable_rows_with_surrogates() { + // A memtable row that does not read refuses the restore: dropping it + // would lose the row from the restored collection. + for scanned in engine.scan_memtable_rows_with_surrogates() { + let (surrogate, values) = scanned?; let surrogate = surrogate.ok_or_else(|| Error::Storage { engine: "columnar".into(), detail: format!( diff --git a/nodedb/src/data/executor/columnar_checkpoint/geometry_restore.rs b/nodedb/src/data/executor/columnar_checkpoint/geometry_restore.rs index 1568b629f..94912b2e4 100644 --- a/nodedb/src/data/executor/columnar_checkpoint/geometry_restore.rs +++ b/nodedb/src/data/executor/columnar_checkpoint/geometry_restore.rs @@ -46,10 +46,7 @@ //! the identity was already lost at write time, and this rebuild neither creates //! nor widens that gap. -use tracing::warn; - use super::super::core_loop::CoreLoop; -use super::super::scan_normalize::decoded_col_to_value; use crate::bridge::envelope::PhysicalPlan; use crate::types::{DatabaseId, TenantId}; use nodedb_physical::physical_plan::{ColumnarInsertIntent, ColumnarOp}; @@ -63,12 +60,17 @@ impl CoreLoop { /// /// A no-op — without decoding anything — for the overwhelmingly common case /// of a collection with no geometry column. + /// + /// `Err` when a restored row does not read: a segment that does not open + /// or decode, a corrupt cell, or a corrupt memtable cell. The rebuilt + /// R-tree would miss that row, so the restore is refused and the core + /// does not come up. pub(super) fn restore_columnar_geometry_indexes( &mut self, key: &(DatabaseId, TenantId, String), engine: &nodedb_columnar::MutationEngine, segments: &[Vec], - ) -> usize { + ) -> crate::Result { let (db_id, tenant_id, collection) = key; let schema = engine.schema().clone(); if !schema @@ -76,9 +78,18 @@ impl CoreLoop { .iter() .any(|c| c.column_type == ColumnType::Geometry) { - return 0; + return Ok(0); } + // Every fallible read runs before the spatial maps change, so a + // refused restore leaves them as the spatial checkpoint loaded them. + let mut rows = self.restored_flushed_rows(engine, segments, &schema, collection)?; + for row in engine.scan_memtable_rows() { + rows.push(row?); + } + // The checkpoint key carries the stored, database-qualified name. + let vshard = nodedb_types::CollectionKey::from_qualified_str(*db_id, collection)?.vshard(); + // The R-tree is derived from the restored rows and nothing else. An // entry a restored spatial checkpoint holds for a row this generation // no longer has (deleted or truncated after that checkpoint) must not @@ -88,12 +99,6 @@ impl CoreLoop { self.spatial_doc_map .retain(|(d, t, c, _, _), _| !(d == db_id && t == tenant_id && c == collection)); - let mut rows: Vec> = Vec::new(); - rows.extend(Self::restored_flushed_rows( - engine, segments, &schema, collection, - )); - rows.extend(engine.scan_memtable_rows()); - // The indexer takes documents, not positional rows: rebuild each row as // the `Value::Object` shape the live insert path hands it, so the two // agree by construction rather than by a second implementation of the @@ -115,19 +120,6 @@ impl CoreLoop { // rows that are already durable, and is not itself a write. Noting an // LSN here would raise the core watermark during boot from a path that // applied no record. - // The checkpoint key carries the stored, database-qualified name. - let vshard = match nodedb_types::CollectionKey::from_qualified_str(*db_id, collection) { - Ok(key) => key.vshard(), - Err(e) => { - warn!( - %collection, - error = %e, - "columnar checkpoint restore: collection name does not de-qualify; its \ - geometry rows are absent from the rebuilt R-tree" - ); - return 0; - } - }; let task = Self::replay_task( *tenant_id, *db_id, @@ -155,7 +147,7 @@ impl CoreLoop { let indexed = docs.len(); // Boot-time rebuild: nothing to roll back, so the delta is dropped. let _ = self.index_columnar_geometry_columns(&task, &schema, collection, &docs); - indexed + Ok(indexed) } /// Decode the live (non-tombstoned) rows of every restored flushed segment. @@ -164,72 +156,150 @@ impl CoreLoop { /// memtable's virtual segment, so `segments[i]` is `segment_id i + 1`, and a /// row whose delete-bitmap bit is set is not a row any more. /// - /// A segment that fails to open or decode is warned about and skipped rather - /// than aborting the restore: this rebuilds a derived index, and skipping - /// costs the same geometry entries a `scan_flushed` over the identical - /// unreadable bytes would also fail to produce. + /// `Err` on the first segment that does not open or decode, or the first + /// corrupt cell of a live row. The shared segment reader files the + /// corruption report. fn restored_flushed_rows( + &self, engine: &nodedb_columnar::MutationEngine, segments: &[Vec], schema: &nodedb_types::columnar::ColumnarSchema, collection: &str, - ) -> Vec> { + ) -> crate::Result>> { let mut out = Vec::new(); for (seg_idx, seg_bytes) in segments.iter().enumerate() { let seg_id = seg_idx as u64 + 1; - let reader = match nodedb_columnar::SegmentReader::open(seg_bytes) { - Ok(r) => r, - Err(e) => { - warn!( - %collection, - seg_id, - error = %e, - "columnar checkpoint restore: flushed segment unreadable; its \ - geometry rows are absent from the rebuilt R-tree" - ); - continue; - } - }; - - let mut decoded_cols = Vec::with_capacity(schema.columns.len()); - let mut decode_ok = true; - for col_idx in 0..schema.columns.len() { - match reader.read_column(col_idx) { - Ok(dc) => decoded_cols.push(dc), - Err(e) => { - warn!( - %collection, - seg_id, - col_idx, - error = %e, - "columnar checkpoint restore: column decode failed; the \ - segment's geometry rows are absent from the rebuilt R-tree" - ); - decode_ok = false; - break; - } - } - } - if !decode_ok { - continue; - } - + let segment = self.decode_flushed_segment( + collection, + seg_id, + seg_bytes, + schema.columns.len(), + "columnar_checkpoint_geometry_restore", + )?; let delete_bm = engine.delete_bitmap(seg_id); - for row_idx in 0..reader.row_count() as usize { + for row_idx in 0..segment.row_count() { if delete_bm.is_some_and(|bm| bm.is_deleted(row_idx as u32)) { continue; } - out.push( - decoded_cols - .iter() - .zip(&schema.columns) - .map(|(dc, col_def)| { - decoded_col_to_value(dc, row_idx, &col_def.column_type) - }) - .collect(), - ); + out.push(segment.row(schema, row_idx)?); } } - out + Ok(out) + } +} + +#[cfg(test)] +mod tests { + use nodedb_columnar::MutationEngine; + use nodedb_types::columnar::{ColumnDef, ColumnType, ColumnarSchema}; + use nodedb_types::value::Value; + + use crate::data::executor::core_loop::CoreLoop; + use crate::data::executor::core_loop::tests::make_core_with_dir; + use crate::types::{DatabaseId, TenantId}; + + type EngineKey = (DatabaseId, TenantId, String); + + fn key() -> EngineKey { + (DatabaseId::DEFAULT, TenantId::new(1), "geo".to_string()) + } + + fn schema() -> ColumnarSchema { + ColumnarSchema::new(vec![ + ColumnDef::required("id", ColumnType::String).with_primary_key(), + ColumnDef::nullable("loc", ColumnType::Geometry), + ]) + .expect("valid") + } + + /// An engine whose one row lives only in a flushed segment, encoded as + /// the flush path encodes it. Returns the engine and the segment bytes. + fn flushed_engine(core: &CoreLoop) -> (MutationEngine, Vec) { + let key = key(); + let mut engine = MutationEngine::new(key.2.clone(), schema()); + engine + .insert(&[ + Value::String("a".into()), + Value::String(r#"{"type":"Point","coordinates":[1.0,2.0]}"#.into()), + ]) + .expect("insert"); + let segment_id = engine.next_segment_id(); + let (seg_schema, columns, row_count) = engine.memtable_mut().drain_optimized(); + let memory = nodedb_mem::ScopedMemory::new( + core.governor.clone(), + key.0, + key.1, + nodedb_mem::EngineId::Columnar, + ); + let blob = + nodedb_columnar::SegmentWriter::new(nodedb_columnar::writer::PROFILE_PLAIN, memory) + .write_segment(&seg_schema, &columns, row_count, None) + .expect("write_segment"); + engine + .on_memtable_flushed(segment_id) + .expect("on_memtable_flushed"); + (engine, blob) + } + + #[test] + fn a_readable_segment_rebuilds_its_geometry_rows() { + let dir = tempfile::tempdir().expect("tempdir"); + let (mut core, _tx, _rx) = make_core_with_dir(dir.path()); + let (engine, blob) = flushed_engine(&core); + let indexed = core + .restore_columnar_geometry_indexes(&key(), &engine, &[blob]) + .expect("restore"); + assert_eq!(indexed, 1); + } + + #[test] + fn an_unopenable_segment_refuses_the_restore() { + let dir = tempfile::tempdir().expect("tempdir"); + let (mut core, _tx, _rx) = make_core_with_dir(dir.path()); + let (engine, _blob) = flushed_engine(&core); + let result = core.restore_columnar_geometry_indexes(&key(), &engine, &[Vec::new()]); + assert!( + matches!(result, Err(crate::Error::SegmentCorrupted { ref detail }) if detail.contains("segment 1 of 'geo'")), + "{result:?}" + ); + } + + /// A segment that opens and decodes, but whose geometry cell is not + /// UTF-8 text. Skipping it would leave the row out of the R-tree while + /// a full scan still sees the segment, so the restore is refused. + #[test] + fn a_corrupt_segment_cell_refuses_the_restore() { + let dir = tempfile::tempdir().expect("tempdir"); + let (mut core, _tx, _rx) = make_core_with_dir(dir.path()); + let key = key(); + let columns = vec![ + nodedb_columnar::memtable::ColumnData::String { + data: b"a".to_vec(), + offsets: vec![0, 1], + valid: None, + }, + nodedb_columnar::memtable::ColumnData::Geometry { + data: vec![0xFF, 0xFE], + offsets: vec![0, 2], + valid: Some(vec![true]), + }, + ]; + let memory = nodedb_mem::ScopedMemory::new( + core.governor.clone(), + key.0, + key.1, + nodedb_mem::EngineId::Columnar, + ); + let blob = + nodedb_columnar::SegmentWriter::new(nodedb_columnar::writer::PROFILE_PLAIN, memory) + .write_segment(&schema(), &columns, 1, None) + .expect("write_segment"); + let engine = MutationEngine::new(key.2.clone(), schema()); + + let result = core.restore_columnar_geometry_indexes(&key, &engine, &[blob]); + assert!( + matches!(result, Err(crate::Error::SegmentCorrupted { ref detail }) if detail.contains("segment 1 of 'geo'")), + "{result:?}" + ); } } diff --git a/nodedb/src/data/executor/columnar_checkpoint/load.rs b/nodedb/src/data/executor/columnar_checkpoint/load.rs index 68b2cac9a..907f57da6 100644 --- a/nodedb/src/data/executor/columnar_checkpoint/load.rs +++ b/nodedb/src/data/executor/columnar_checkpoint/load.rs @@ -90,8 +90,9 @@ impl CoreLoop { // rebuilt from the restored rows, not carried in the checkpoint — // see `geometry_restore.rs`. Done BEFORE the maps are populated so // it reads the restored engine and blobs directly and cannot see a - // half-installed state. - geometry_rows += self.restore_columnar_geometry_indexes(&key, &engine, &blobs); + // half-installed state. A restored row that does not read refuses + // the load: the rebuilt R-tree would silently miss it. + geometry_rows += self.restore_columnar_geometry_indexes(&key, &engine, &blobs)?; segments += blobs.len(); // Both halves are installed from one destructured value, in the same diff --git a/nodedb/src/data/executor/dispatch/meta_retention/columnar_plain.rs b/nodedb/src/data/executor/dispatch/meta_retention/columnar_plain.rs index 075ee38af..289ccf852 100644 --- a/nodedb/src/data/executor/dispatch/meta_retention/columnar_plain.rs +++ b/nodedb/src/data/executor/dispatch/meta_retention/columnar_plain.rs @@ -3,21 +3,29 @@ //! Plain columnar profile temporal-purge. //! //! Row-level audit purge on bitemporal plain-columnar collections. -//! Walks sealed segments (via [`nodedb_columnar::SegmentReader`]) and the -//! live memtable, groups rows by primary key, and marks every +//! Walks flushed segments (through the shared flushed-segment reader) and +//! the live memtable, groups rows by primary key, and marks every //! *superseded* row whose `_ts_system` is below the cutoff in the //! engine's per-segment delete bitmap. //! //! The single latest version per PK is always preserved — even if it is //! itself below the cutoff — so "AS OF" reads beyond the cutoff can still //! resolve each logical row's terminal state. +//! +//! A version's row index is its physical position in its segment or in the +//! memtable: the index the delete bitmap marks. use std::collections::HashMap; +use nodedb_columnar::MutationEngine; +use nodedb_types::columnar::ColumnType; +use nodedb_types::value::Value; use nodedb_types::{DatabaseId, TenantId}; use crate::data::executor::core_loop::CoreLoop; -use crate::data::executor::scan_normalize::decoded_col_to_value; + +/// The read path the corruption report names. +const SITE: &str = "columnar_temporal_purge"; pub(super) struct RowVersion { pub seg_id: u64, @@ -31,6 +39,9 @@ impl CoreLoop { /// per-segment delete bitmaps. No WAL records are appended here; the /// caller is responsible for the `RecordType::TemporalPurge` audit /// record that covers the batch. + /// + /// `Err` when a segment or memtable row does not read, or a row carries + /// no system time. Nothing is marked then. pub(super) fn plain_columnar_purge( &mut self, database_id: DatabaseId, @@ -39,14 +50,47 @@ impl CoreLoop { cutoff_system_ms: i64, ) -> crate::Result { let key = (database_id, tid, collection.to_string()); - let engine = match self.columnar_engines.get(&key) { - Some(e) => e, - None => return Ok(0), + let victims_per_seg = { + let Some(engine) = self.columnar_engines.get(&key) else { + return Ok(0); + }; + if !engine.schema().is_bitemporal() { + return Ok(0); + } + let versions = self.plain_columnar_row_versions(&key, engine)?; + purge_victims(engine, versions, cutoff_system_ms) }; - let schema = engine.schema(); - if !schema.is_bitemporal() { + if victims_per_seg.is_empty() { return Ok(0); } + + let engine_mut = + self.columnar_engines + .get_mut(&key) + .ok_or_else(|| crate::Error::Internal { + detail: format!("columnar engine for '{collection}' is gone mid-purge"), + })?; + let mut total = 0usize; + for (seg_id, mut row_indices) in victims_per_seg { + row_indices.sort_unstable(); + row_indices.dedup(); + total += row_indices.len(); + engine_mut + .delete_bitmap_mut(seg_id) + .mark_deleted_batch(&row_indices); + } + Ok(total) + } + + /// Every row version of the collection at `key`. A flushed row is listed + /// tombstoned or not. A memtable row is listed only when live. + fn plain_columnar_row_versions( + &self, + key: &(DatabaseId, TenantId, String), + engine: &MutationEngine, + ) -> crate::Result> { + let collection = key.2.as_str(); + let schema = engine.schema(); let ts_idx = schema .ts_system_idx() .ok_or_else(|| crate::Error::Storage { @@ -60,172 +104,203 @@ impl CoreLoop { detail: "bitemporal collection without primary key columns".into(), }); } + // Decode the columns up to the last one a version needs. + let column_count = pk_indices.iter().copied().fold(ts_idx, usize::max) + 1; let mut versions: Vec = Vec::new(); // Flushed segments: segment_id starts at 1 (memtable is id 0). - if let Some(segments) = self.columnar_flushed_segments.get(&key) { - for (seg_idx, seg_bytes) in segments.iter().enumerate() { - let seg_id = seg_idx as u64 + 1; - let seg_id_str = seg_id.to_string(); - let reader = if let Some(reg) = &self.quarantine_registry { - crate::storage::quarantine::engines::open_segment_with_quarantine( - reg, - seg_bytes, - collection, - &seg_id_str, - ) - .map_err(|e| crate::Error::Storage { - engine: "columnar".into(), - detail: format!("open segment for purge: {e}"), - })? - } else { - nodedb_columnar::SegmentReader::open(seg_bytes).map_err(|e| { - crate::Error::Storage { - engine: "columnar".into(), - detail: format!("open segment for purge: {e}"), - } - })? - }; - let ts_col = reader - .read_column(ts_idx) - .map_err(|e| crate::Error::Storage { - engine: "columnar".into(), - detail: format!("read _ts_system: {e}"), + let segments = self + .columnar_flushed_segments + .get(key) + .map(Vec::as_slice) + .unwrap_or_default(); + for (seg_idx, seg_bytes) in segments.iter().enumerate() { + let seg_id = seg_idx as u64 + 1; + let segment = + self.decode_flushed_segment(collection, seg_id, seg_bytes, column_count, SITE)?; + for row_idx in 0..segment.row_count() { + let ts_cell = segment.cell(ts_idx, row_idx, &ColumnType::Int64)?; + let sys_ts = + system_time(&ts_cell, collection, seg_id, row_idx).inspect_err(|e| { + crate::diag::columnar_segment_corrupt(e, collection, seg_id, "cell", SITE); })?; - let pk_cols: Vec<( - nodedb_columnar::reader::DecodedColumn, - &nodedb_types::columnar::ColumnType, - )> = pk_indices + let pk_cells = pk_indices .iter() - .map(|&i| { - reader - .read_column(i) - .map(|col| (col, &schema.columns[i].column_type)) - }) - .collect::>() - .map_err(|e| crate::Error::Storage { - engine: "columnar".into(), - detail: format!("read pk column: {e}"), - })?; - let row_count = reader.row_count() as usize; - for row_idx in 0..row_count { - let Some(sys_ts) = int64_from_decoded(&ts_col, row_idx) else { - continue; - }; - let pk_bytes = encode_pk_from_decoded_cols(&pk_cols, row_idx); - versions.push(RowVersion { - seg_id, - row_idx: row_idx as u32, - pk_bytes, - sys_ts, - }); - } + .map(|&i| segment.cell(i, row_idx, &schema.columns[i].column_type)) + .collect::>>()?; + versions.push(RowVersion { + seg_id, + row_idx: row_index_u32(collection, seg_id, row_idx)?, + pk_bytes: version_pk_bytes(&pk_cells), + sys_ts, + }); } } - // Memtable rows. + // Memtable rows, each at its physical index. let memtable_seg_id = engine.memtable_segment_id(); - let rows: Vec> = engine.scan_memtable_rows().collect(); - for (row_idx, row) in rows.iter().enumerate() { - let sys_ts = match row.get(ts_idx) { - Some(nodedb_types::value::Value::Integer(n)) => *n, - _ => continue, - }; - let pk_values: Vec<&nodedb_types::value::Value> = - pk_indices.iter().filter_map(|&i| row.get(i)).collect(); - if pk_values.len() != pk_indices.len() { + for row_idx in 0..engine.memtable().row_count() { + // `None` is a row the memtable delete bitmap marks. + let Some(row) = engine.get_memtable_row(row_idx)? else { continue; - } - let pk_bytes = if pk_values.len() == 1 { - nodedb_columnar::pk_index::encode_pk(pk_values[0]) - } else { - nodedb_columnar::pk_index::encode_composite_pk(&pk_values) }; + let sys_ts = system_time( + row.get(ts_idx).unwrap_or(&Value::Null), + collection, + memtable_seg_id, + row_idx, + )?; + let pk_cells = pk_indices + .iter() + .map(|&i| { + row.get(i).cloned().ok_or_else(|| crate::Error::Internal { + detail: format!( + "columnar '{collection}': memtable row {row_idx} holds no \ + primary-key cell {i}" + ), + }) + }) + .collect::>>()?; versions.push(RowVersion { seg_id: memtable_seg_id, - row_idx: row_idx as u32, - pk_bytes, + row_idx: row_index_u32(collection, memtable_seg_id, row_idx)?, + pk_bytes: version_pk_bytes(&pk_cells), sys_ts, }); } + Ok(versions) + } +} - // Find latest system_ts per PK. - let mut latest: HashMap, i64> = HashMap::new(); - for v in &versions { - latest - .entry(v.pk_bytes.clone()) - .and_modify(|cur| { - if v.sys_ts > *cur { - *cur = v.sys_ts; - } - }) - .or_insert(v.sys_ts); - } +/// The row indices to tombstone, per segment: every version superseded by a +/// later one of its PK, below `cutoff_system_ms`, and not tombstoned yet. +fn purge_victims( + engine: &MutationEngine, + versions: Vec, + cutoff_system_ms: i64, +) -> HashMap> { + let mut latest: HashMap<&[u8], i64> = HashMap::new(); + for v in &versions { + latest + .entry(v.pk_bytes.as_slice()) + .and_modify(|cur| *cur = (*cur).max(v.sys_ts)) + .or_insert(v.sys_ts); + } - // Victims: superseded AND below cutoff AND not already tombstoned. - let mut victims_per_seg: HashMap> = HashMap::new(); - for v in versions { - let lat = latest.get(&v.pk_bytes).copied().unwrap_or(v.sys_ts); - if v.sys_ts < cutoff_system_ms && v.sys_ts < lat { - let already = engine - .delete_bitmap(v.seg_id) - .is_some_and(|bm| bm.is_deleted(v.row_idx)); - if !already { - victims_per_seg.entry(v.seg_id).or_default().push(v.row_idx); - } + let mut victims_per_seg: HashMap> = HashMap::new(); + for v in &versions { + let lat = latest + .get(v.pk_bytes.as_slice()) + .copied() + .unwrap_or(v.sys_ts); + if v.sys_ts < cutoff_system_ms && v.sys_ts < lat { + let already = engine + .delete_bitmap(v.seg_id) + .is_some_and(|bm| bm.is_deleted(v.row_idx)); + if !already { + victims_per_seg.entry(v.seg_id).or_default().push(v.row_idx); } } + } + victims_per_seg +} - if victims_per_seg.is_empty() { - return Ok(0); - } - - let engine_mut = self - .columnar_engines - .get_mut(&key) - .expect("engine vanished mid-purge"); - let mut total = 0usize; - for (seg_id, mut row_indices) in victims_per_seg { - row_indices.sort_unstable(); - row_indices.dedup(); - total += row_indices.len(); - engine_mut - .delete_bitmap_mut(seg_id) - .mark_deleted_batch(&row_indices); - } - Ok(total) +/// The system time `_ts_system` holds. The column is a required `Int64`, so +/// any other cell is an error. +fn system_time(cell: &Value, collection: &str, seg_id: u64, row_idx: usize) -> crate::Result { + match cell { + Value::Integer(n) => Ok(*n), + other => Err(crate::Error::SegmentCorrupted { + detail: format!( + "columnar '{collection}': row {row_idx} of segment {seg_id} holds {other:?} \ + in _ts_system, not an integer system time" + ), + }), } } -fn int64_from_decoded(col: &nodedb_columnar::reader::DecodedColumn, row_idx: usize) -> Option { - use nodedb_columnar::reader::DecodedColumn; - match col { - DecodedColumn::Int64 { values, valid } | DecodedColumn::Timestamp { values, valid } => { - (row_idx < valid.len() && valid[row_idx]).then(|| values[row_idx]) +/// Encode a version's primary key from its PK cells, in PK column order. +/// The bytes match what the engine's PK index holds for the same row. +fn version_pk_bytes(cells: &[Value]) -> Vec { + match cells { + [single] => nodedb_columnar::pk_index::encode_pk(single), + _ => { + let refs: Vec<&Value> = cells.iter().collect(); + nodedb_columnar::pk_index::encode_composite_pk(&refs) } - _ => None, } } -/// Encode the primary key of one flushed row. Each cell is typed by its -/// declared column type, so the key bytes match what the memtable path -/// encodes for the same row: a `TIMESTAMP` key cell is an instant on both. -fn encode_pk_from_decoded_cols( - cols: &[( - nodedb_columnar::reader::DecodedColumn, - &nodedb_types::columnar::ColumnType, - )], - row_idx: usize, -) -> Vec { - let values: Vec = cols - .iter() - .map(|(c, declared)| decoded_col_to_value(c, row_idx, declared)) - .collect(); - if values.len() == 1 { - nodedb_columnar::pk_index::encode_pk(&values[0]) - } else { - let refs: Vec<&nodedb_types::value::Value> = values.iter().collect(); - nodedb_columnar::pk_index::encode_composite_pk(&refs) +/// `row_idx` as the `u32` row index a delete bitmap marks. +fn row_index_u32(collection: &str, seg_id: u64, row_idx: usize) -> crate::Result { + u32::try_from(row_idx).map_err(|_| crate::Error::Internal { + detail: format!( + "columnar '{collection}': row {row_idx} of segment {seg_id} is past the u32 row \ + index range" + ), + }) +} + +#[cfg(test)] +mod tests { + use nodedb_columnar::pk_index::encode_pk; + use nodedb_types::columnar::{ColumnDef, ColumnarSchema}; + + use super::*; + use crate::data::executor::core_loop::tests::make_core_with_dir; + + fn schema() -> ColumnarSchema { + ColumnarSchema::new(vec![ + ColumnDef::required("_ts_system", ColumnType::Int64), + ColumnDef::required("_ts_valid_from", ColumnType::Int64), + ColumnDef::required("_ts_valid_until", ColumnType::Int64), + ColumnDef::required("id", ColumnType::Int64).with_primary_key(), + ]) + .expect("valid") + } + + fn version(sys_ts: i64, id: i64) -> Vec { + vec![ + Value::Integer(sys_ts), + Value::Integer(i64::MIN), + Value::Integer(i64::MAX), + Value::Integer(id), + ] + } + + /// A tombstoned memtable row ahead of the superseded version does not + /// shift the index the purge marks. + #[test] + fn the_purge_marks_the_physical_row_behind_a_deleted_row() { + let dir = tempfile::tempdir().expect("tempdir"); + let (mut core, _tx, _rx) = make_core_with_dir(dir.path()); + let mut engine = MutationEngine::new("bt".to_string(), schema()); + // Physical rows: 0 = id 9 (deleted below), 1 = id 1 at 10, 2 = id 1 at 50. + engine.insert(&version(5, 9)).expect("insert"); + engine.insert(&version(10, 1)).expect("insert"); + engine.insert(&version(50, 1)).expect("insert"); + engine.delete(&Value::Integer(9)).expect("delete"); + let memtable_seg = engine.memtable_segment_id(); + let key = (DatabaseId::DEFAULT, TenantId::new(1), "bt".to_string()); + core.columnar_engines.insert(key.clone(), engine); + + let purged = core + .plain_columnar_purge(key.0, key.1, "bt", 100) + .expect("purge"); + + assert_eq!(purged, 1, "the version of id 1 at 10 is superseded"); + let engine = core.columnar_engines.get(&key).expect("engine"); + let bitmap = engine.delete_bitmap(memtable_seg).expect("bitmap"); + assert!(bitmap.is_deleted(0), "the deleted row stays deleted"); + assert!(bitmap.is_deleted(1), "the superseded version is marked"); + assert!(!bitmap.is_deleted(2), "the latest version stays live"); + assert_eq!( + engine + .pk_index() + .get(&encode_pk(&Value::Integer(1))) + .map(|loc| loc.row_index), + Some(2) + ); } } diff --git a/nodedb/src/data/executor/handlers/columnar_read/convert.rs b/nodedb/src/data/executor/handlers/columnar_read/convert.rs index 6300f0bc5..1f348e55c 100644 --- a/nodedb/src/data/executor/handlers/columnar_read/convert.rs +++ b/nodedb/src/data/executor/handlers/columnar_read/convert.rs @@ -76,7 +76,8 @@ pub(in crate::data::executor) fn row_to_projected_value( /// A time cell is written as its column's kind says: an instant column /// yields a typed instant ext, a `Millis` column the integer stored. `Err` /// when a stored millisecond count overflows the microsecond range an -/// instant carries; the cell is never written wrapped. +/// instant carries; the cell is never written wrapped. `Err` when the cell +/// does not read as its declared type, never a written NULL. pub(in crate::data::executor) fn emit_column_value( buf: &mut Vec, mt: &crate::engine::timeseries::columnar_memtable::ColumnarMemtable, @@ -85,41 +86,24 @@ pub(in crate::data::executor) fn emit_column_value( col_data: &crate::engine::timeseries::columnar_memtable::ColumnData, row_idx: usize, ) -> crate::Result<()> { - use crate::engine::timeseries::columnar_memtable::{ - ColumnData as TsColumnData, ColumnType as TsColumnType, - }; - match col_type { - TsColumnType::Timestamp(kind) => { - let millis = col_data.as_timestamps()[row_idx]; - write_time_cell(buf, *kind, millis)?; - } - TsColumnType::Float64 => { - let v = col_data.as_f64()[row_idx]; - if v.is_finite() { - nodedb_query::msgpack_scan::write_f64(buf, v); - } else { - nodedb_query::msgpack_scan::write_null(buf); - } - } - TsColumnType::Symbol => { - if let TsColumnData::Symbol(ids) = col_data { - let sym_id = ids[row_idx]; - if let Some(s) = mt.symbol_dict(col_idx).and_then(|dict| dict.get(sym_id)) { - nodedb_query::msgpack_scan::write_str(buf, s); - } else { - nodedb_query::msgpack_scan::write_null(buf); - } - } else { - nodedb_query::msgpack_scan::write_null(buf); - } - } - TsColumnType::Int64 => { - if let TsColumnData::Int64(vals) = col_data { - nodedb_query::msgpack_scan::write_i64(buf, vals[row_idx]); - } else { - nodedb_query::msgpack_scan::write_null(buf); - } - } + use crate::data::executor::handlers::timeseries::cell_read::{TsCell, read_ts_cell}; + let column = mt + .schema() + .columns + .get(col_idx) + .map_or("", |(name, _)| name.as_str()); + match read_ts_cell( + col_data, + *col_type, + column, + mt.symbol_dict(col_idx), + row_idx, + )? { + TsCell::Time(kind, millis) => write_time_cell(buf, kind, millis)?, + TsCell::Float(v) if v.is_finite() => nodedb_query::msgpack_scan::write_f64(buf, v), + TsCell::Float(_) | TsCell::Null => nodedb_query::msgpack_scan::write_null(buf), + TsCell::Symbol(s) => nodedb_query::msgpack_scan::write_str(buf, s), + TsCell::Int(n) => nodedb_query::msgpack_scan::write_i64(buf, n), } Ok(()) } diff --git a/nodedb/src/data/executor/handlers/columnar_read/flushed_segment.rs b/nodedb/src/data/executor/handlers/columnar_read/flushed_segment.rs new file mode 100644 index 000000000..409650a65 --- /dev/null +++ b/nodedb/src/data/executor/handlers/columnar_read/flushed_segment.rs @@ -0,0 +1,206 @@ +// SPDX-License-Identifier: BUSL-1.1 + +//! The one reader every flushed columnar segment read goes through. +//! +//! A segment that does not open, a column that does not decode, and a cell +//! that does not hold its declared type are all corruption. Each one returns +//! a typed error and files one corruption report here, where it is detected. +//! No caller skips a segment or a row it cannot read: a skipped row reads as +//! an absent row, which is a wrong answer, not a degraded one. + +use nodedb_columnar::reader::DecodedColumn; +use nodedb_types::columnar::{ColumnType, ColumnarSchema}; +use nodedb_types::value::Value; + +use crate::data::executor::core_loop::CoreLoop; +use crate::data::executor::scan_normalize::decoded_col_to_value; + +/// One flushed segment with every schema column decoded. +pub(in crate::data::executor) struct DecodedFlushedSegment<'a> { + collection: &'a str, + segment_id: u64, + site: &'static str, + row_count: usize, + columns: Vec, +} + +impl CoreLoop { + /// Open flushed segment `segment_id` of `collection` and decode its first + /// `column_count` columns. + /// + /// The open goes through the quarantine registry when one is wired, so a + /// repeated CRC failure quarantines the segment. `site` names the read + /// path in the corruption report. + pub(in crate::data::executor) fn decode_flushed_segment<'a>( + &self, + collection: &'a str, + segment_id: u64, + seg_bytes: &[u8], + column_count: usize, + site: &'static str, + ) -> crate::Result> { + let report = |err: crate::Error, stage: &'static str| { + crate::diag::columnar_segment_corrupt(&err, collection, segment_id, stage, site); + err + }; + let opened = match &self.quarantine_registry { + Some(reg) => crate::storage::quarantine::engines::open_segment_with_quarantine( + reg, + seg_bytes, + collection, + &segment_id.to_string(), + ) + .map_err(crate::Error::from), + None => nodedb_columnar::SegmentReader::open(seg_bytes).map_err(crate::Error::from), + }; + let reader = + opened.map_err(|e| report(segment_error(collection, segment_id, e), "open"))?; + let columns = (0..column_count) + .map(|col_idx| reader.read_column(col_idx)) + .collect::, _>>() + .map_err(|e| { + report( + segment_error(collection, segment_id, crate::Error::from(e)), + "column", + ) + })?; + let row_count = usize::try_from(reader.row_count()).map_err(|_| { + report( + crate::Error::SegmentCorrupted { + detail: format!( + "columnar segment {segment_id} of '{collection}' reports {} rows, \ + more than this machine can address", + reader.row_count() + ), + }, + "open", + ) + })?; + Ok(DecodedFlushedSegment { + collection, + segment_id, + site, + row_count, + columns, + }) + } +} + +impl DecodedFlushedSegment<'_> { + /// Rows the segment holds, tombstoned ones included. + pub(in crate::data::executor) fn row_count(&self) -> usize { + self.row_count + } + + /// Row `row_idx`, each cell typed by its declared column. + pub(in crate::data::executor) fn row( + &self, + schema: &ColumnarSchema, + row_idx: usize, + ) -> crate::Result> { + self.columns + .iter() + .zip(&schema.columns) + .map(|(col, def)| self.decode_cell(col, row_idx, &def.column_type)) + .collect() + } + + /// Cell `row_idx` of column `col_idx`, typed by `declared`. + pub(in crate::data::executor) fn cell( + &self, + col_idx: usize, + row_idx: usize, + declared: &ColumnType, + ) -> crate::Result { + let Some(col) = self.columns.get(col_idx) else { + return Err(crate::Error::Internal { + detail: format!( + "columnar segment {} of '{}': column {col_idx} was not decoded", + self.segment_id, self.collection + ), + }); + }; + self.decode_cell(col, row_idx, declared) + } + + fn decode_cell( + &self, + col: &DecodedColumn, + row_idx: usize, + declared: &ColumnType, + ) -> crate::Result { + decoded_col_to_value(col, row_idx, declared).map_err(|e| { + let err = segment_error(self.collection, self.segment_id, e); + crate::diag::columnar_segment_corrupt( + &err, + self.collection, + self.segment_id, + "cell", + self.site, + ); + err + }) + } +} + +/// `err` with the segment and collection named in its detail when it is a +/// corruption error, so the refusal says which segment is damaged. +fn segment_error(collection: &str, segment_id: u64, err: crate::Error) -> crate::Error { + match err { + crate::Error::SegmentCorrupted { detail } => crate::Error::SegmentCorrupted { + detail: format!("columnar segment {segment_id} of '{collection}': {detail}"), + }, + other => other, + } +} + +#[cfg(test)] +mod tests { + use nodedb_types::columnar::ColumnDef; + + use super::*; + use crate::data::executor::core_loop::tests::make_core_with_dir; + + fn schema() -> ColumnarSchema { + ColumnarSchema::new(vec![ + ColumnDef::required("id", ColumnType::Int64).with_primary_key(), + ColumnDef::nullable("doc", ColumnType::Json), + ]) + .expect("valid") + } + + #[test] + fn unopenable_segment_bytes_are_refused_as_corruption() { + let dir = tempfile::tempdir().expect("tempdir"); + let (core, _tx, _rx) = make_core_with_dir(dir.path()); + let result = core.decode_flushed_segment("m", 1, &[], schema().columns.len(), "test"); + assert!( + matches!(result, Err(crate::Error::SegmentCorrupted { ref detail }) if detail.contains("segment 1 of 'm'")) + ); + } + + #[test] + fn a_cell_that_does_not_fit_its_type_is_refused_as_corruption() { + let segment = DecodedFlushedSegment { + collection: "m", + segment_id: 2, + site: "test", + row_count: 1, + columns: vec![ + DecodedColumn::Int64 { + values: vec![1], + valid: vec![true], + }, + DecodedColumn::Binary { + data: vec![0xC1], + offsets: vec![0, 1], + valid: vec![true], + }, + ], + }; + let result = segment.row(&schema(), 0); + assert!( + matches!(result, Err(crate::Error::SegmentCorrupted { ref detail }) if detail.contains("segment 2 of 'm'")) + ); + } +} diff --git a/nodedb/src/data/executor/handlers/columnar_read/materialize_scan.rs b/nodedb/src/data/executor/handlers/columnar_read/materialize_scan.rs index 713b1a410..84f1552e0 100644 --- a/nodedb/src/data/executor/handlers/columnar_read/materialize_scan.rs +++ b/nodedb/src/data/executor/handlers/columnar_read/materialize_scan.rs @@ -26,7 +26,6 @@ use nodedb_types::value::Value; use crate::bridge::envelope::Response; use crate::data::executor::core_loop::CoreLoop; -use crate::data::executor::scan_normalize::decoded_col_to_value; use crate::data::executor::task::ExecutionTask; impl CoreLoop { @@ -101,44 +100,19 @@ impl CoreLoop { continue; } - let reader = match nodedb_columnar::SegmentReader::open(seg_bytes) { - Ok(r) => r, - Err(e) => { - tracing::warn!( - collection, - seg_id, - error = %e, - "materialize_scan: failed to open flushed segment; skipping" - ); - continue; - } + // A segment that does not read refuses the scan. Skipping it + // would hand the materializer a clone without its rows. + let segment = match self.decode_flushed_segment( + collection, + u64::from(seg_id), + seg_bytes, + schema.columns.len(), + "columnar_materialize_scan", + ) { + Ok(segment) => segment, + Err(e) => return self.response_error(task, e), }; - - let row_count = reader.row_count() as usize; - let col_count = schema.columns.len(); - - // Decode all columns once per segment for efficiency. - let mut decoded_cols = Vec::with_capacity(col_count); - let mut decode_ok = true; - for col_idx in 0..col_count { - match reader.read_column(col_idx) { - Ok(dc) => decoded_cols.push(dc), - Err(e) => { - tracing::warn!( - collection, - seg_id, - col_idx, - error = %e, - "materialize_scan: column decode failed; skipping segment" - ); - decode_ok = false; - break; - } - } - } - if !decode_ok { - continue; - } + let row_count = segment.row_count(); // Starting row within this segment. let first_row_in_seg = if seg_id == start_segment { @@ -166,11 +140,11 @@ impl CoreLoop { // Bitemporal system-time filter. if let (Some(ts_idx), Some(cutoff)) = (ts_system_idx, system_as_of_ms) { - let ts_val = decoded_col_to_value( - &decoded_cols[ts_idx], - row_idx, - &schema.columns[ts_idx].column_type, - ); + let ts_val = + match segment.cell(ts_idx, row_idx, &schema.columns[ts_idx].column_type) { + Ok(v) => v, + Err(e) => return self.response_error(task, e), + }; if let Value::Integer(ts) = ts_val && ts > cutoff { @@ -179,27 +153,23 @@ impl CoreLoop { } // Build a Value::Object for this row. - let mut map = std::collections::HashMap::new(); - for (col_idx, col_def) in schema.columns.iter().enumerate() { - let val = - decoded_col_to_value(&decoded_cols[col_idx], row_idx, &col_def.column_type); - map.insert(col_def.name.clone(), val); - } - - // Encode as msgpack value bytes (the Insert handler reads this format). - let ndb_val = Value::Object(map); - let value_bytes = match nodedb_types::value_to_msgpack(&ndb_val) { + let row = match segment.row(&schema, row_idx) { + Ok(row) => row, + Err(e) => return self.response_error(task, e), + }; + let map: std::collections::HashMap = schema + .columns + .iter() + .map(|col_def| col_def.name.clone()) + .zip(row) + .collect(); + + // Encode as msgpack value bytes (the Insert handler reads this + // format). A row that does not encode refuses the scan rather + // than leaving it out of the clone. + let value_bytes = match encode_row(&Value::Object(map)) { Ok(b) => b, - Err(e) => { - tracing::warn!( - collection, - seg_id, - row_idx, - error = %e, - "materialize_scan: row msgpack encode failed; skipping" - ); - continue; - } + Err(e) => return self.response_error(task, e), }; // Emit the real per-row surrogate when available so the @@ -255,10 +225,14 @@ impl CoreLoop { let ts_system_idx = schema.columns.iter().position(|c| c.name == TS_SYSTEM); let rows_with_surrogates: Vec<(Option, Vec)> = - engine + match engine .scan_memtable_rows_with_surrogates() .skip(memtable_start_row) - .collect(); + .collect::>() + { + Ok(rows) => rows, + Err(e) => return self.response_error(task, crate::Error::from(e)), + }; for (mt_idx, (row_surrogate, row)) in rows_with_surrogates.iter().enumerate() { // Bitemporal system-time filter. @@ -276,18 +250,9 @@ impl CoreLoop { map.insert(col_def.name.clone(), row[col_idx].clone()); } } - let ndb_val = Value::Object(map); - let value_bytes = match nodedb_types::value_to_msgpack(&ndb_val) { + let value_bytes = match encode_row(&Value::Object(map)) { Ok(b) => b, - Err(e) => { - tracing::warn!( - collection, - mt_idx, - error = %e, - "materialize_scan: memtable row encode failed; skipping" - ); - continue; - } + Err(e) => return self.response_error(task, e), }; let abs_row = memtable_start_row + mt_idx; @@ -320,6 +285,15 @@ impl CoreLoop { } } +/// Encode one materialized row as MessagePack value bytes, the format the +/// Insert handler reads. +fn encode_row(row: &Value) -> crate::Result> { + nodedb_types::value_to_msgpack(row).map_err(|e| crate::Error::Serialization { + format: "msgpack".into(), + detail: format!("encode materialized columnar row: {e}"), + }) +} + /// Encode `(seg_id, row_idx)` as a compact 32-bit tag. /// /// This is the **fallback** used only for flushed-segment rows that have no diff --git a/nodedb/src/data/executor/handlers/columnar_read/materialize_scan_ts.rs b/nodedb/src/data/executor/handlers/columnar_read/materialize_scan_ts.rs index eca7e34a9..05909928e 100644 --- a/nodedb/src/data/executor/handlers/columnar_read/materialize_scan_ts.rs +++ b/nodedb/src/data/executor/handlers/columnar_read/materialize_scan_ts.rs @@ -50,9 +50,15 @@ use nodedb_types::value::Value; use super::materialize_scan::{build_response, encode_cursor}; use crate::data::executor::core_loop::CoreLoop; +use crate::data::executor::handlers::timeseries::cell_read::{TsCell, read_ts_cell}; +use crate::data::executor::handlers::timeseries::partition_read::{ + TsPartitionColumns, partition_corrupt, read_ts_partition, +}; use crate::data::executor::task::ExecutionTask; use crate::engine::timeseries::columnar_memtable::{ColumnData, ColumnType}; -use crate::engine::timeseries::columnar_segment::ColumnarSegmentReader; + +/// The read path the corruption report names. +const SITE: &str = "timeseries_materialize_scan"; impl CoreLoop { /// Execute a cursor-paginated materialize scan for a timeseries collection. @@ -92,7 +98,10 @@ impl CoreLoop { for row_idx in first_row..row_count { // `system_as_of_ms` filter: skip rows newer than the cutoff. if let (Some(sys_idx), Some(cutoff)) = (ts_system_idx, system_as_of_ms) { - let ts_val = ts_memtable_value(mt.column(sys_idx), row_idx); + let ts_val = match memtable_system_time(mt, sys_idx, row_idx) { + Ok(ts) => ts, + Err(e) => return self.response_error(task, e), + }; if ts_val > cutoff { continue; } @@ -155,52 +164,38 @@ impl CoreLoop { if *part_id < first_part_id { continue; } - if !part_dir.exists() { - continue; - } - let schema = match ColumnarSegmentReader::read_schema(part_dir, None) { - Ok(s) => s, - Err(e) => { - tracing::warn!( - collection, - part_id, - error = %e, - "ts_materialize_scan: failed to read schema; skipping partition" - ); - continue; - } + // A partition that does not read refuses the scan, a missing + // directory included. Skipping it would hand the materializer + // a clone without its rows. + let TsPartitionColumns { + schema, + columns: col_data, + sym_dicts, + } = match read_ts_partition(part_dir, SITE) { + Ok(partition) => partition, + Err(e) => return self.response_error(task, e), }; let ts_system_idx = schema.ts_system_idx(); - // Read all columns. - let col_data: Vec> = schema - .columns - .iter() - .map(|(name, ty)| { - ColumnarSegmentReader::read_column(part_dir, name, *ty, None).ok() - }) - .collect(); - - // Read symbol dictionaries. - let sym_dicts: HashMap = schema - .columns - .iter() - .enumerate() - .filter(|(_, (_, ty))| *ty == ColumnType::Symbol) - .filter_map(|(i, (name, _))| { - ColumnarSegmentReader::read_symbol_dict(part_dir, name, None) - .ok() - .map(|dict| (i, dict)) - }) - .collect(); - // Determine row count from the timestamp column. - let ts_col = col_data.get(schema.timestamp_idx).and_then(|d| d.as_ref()); - let row_count = match ts_col { + let row_count = match col_data.get(schema.timestamp_idx).and_then(|d| d.as_ref()) { Some(col) => col.len(), - None => continue, + None => { + return self.response_error( + task, + partition_corrupt( + part_dir, + "schema", + SITE, + format!( + "time column index {} is outside its schema", + schema.timestamp_idx + ), + ), + ); + } }; let first_row_in_part = if *part_id == start_segment as usize { @@ -209,21 +204,25 @@ impl CoreLoop { 0 }; + let partition = PartitionCells { + part_dir, + schema_columns: &schema.columns, + col_data: &col_data, + sym_dicts: &sym_dicts, + }; for row_idx in first_row_in_part..row_count { // `system_as_of_ms` filter. - if let (Some(sys_idx), Some(cutoff)) = (ts_system_idx, system_as_of_ms) - && let Some(sys_col) = &col_data[sys_idx] - && ts_partition_value(sys_col, row_idx) > cutoff - { - continue; + if let (Some(sys_idx), Some(cutoff)) = (ts_system_idx, system_as_of_ms) { + let ts_val = match partition.system_time(sys_idx, row_idx) { + Ok(ts) => ts, + Err(e) => return self.response_error(task, e), + }; + if ts_val > cutoff { + continue; + } } - let value_bytes = match encode_ts_partition_row( - &schema.columns, - &col_data, - &sym_dicts, - row_idx, - ) { + let value_bytes = match partition.encode_row(row_idx) { Ok(b) => b, Err(e) => return self.response_error(task, e), }; @@ -302,35 +301,104 @@ fn encode_ts_memtable_row( let mut map: HashMap = HashMap::with_capacity(col_count); let schema = mt.schema(); - for (col_idx, (col_name, col_type)) in schema.columns.iter().enumerate() { - let col_data = mt.column(col_idx); - let val = memtable_col_to_value(col_data, col_type, col_idx, mt, row_idx) - .map_err(|e| instant_read_error(col_name, e))?; - map.insert(col_name.clone(), val); + for (col_idx, (col_name, _)) in schema.columns.iter().enumerate() { + let cell = memtable_cell(mt, col_idx, row_idx)?; + map.insert(col_name.clone(), cell_to_value(cell, col_name)?); } encode_row_map(map, row_idx) } -/// Encode a partition row as msgpack `Value::Object` bytes. -fn encode_ts_partition_row( - schema_columns: &[(String, ColumnType)], - col_data: &[Option], - sym_dicts: &HashMap, +/// Cell `row_idx` of memtable column `col_idx`. A cell that does not read +/// as its declared type is an error. +fn memtable_cell( + mt: &crate::engine::timeseries::columnar_memtable::ColumnarMemtable, + col_idx: usize, row_idx: usize, -) -> crate::Result> { - let mut map: HashMap = HashMap::with_capacity(schema_columns.len()); +) -> crate::Result> { + let (col_name, col_type) = &mt.schema().columns[col_idx]; + read_ts_cell( + mt.column(col_idx), + *col_type, + col_name, + mt.symbol_dict(col_idx), + row_idx, + ) + .map_err(crate::Error::from) +} + +/// The system time of memtable row `row_idx`, read from column `sys_idx`. +fn memtable_system_time( + mt: &crate::engine::timeseries::columnar_memtable::ColumnarMemtable, + sys_idx: usize, + row_idx: usize, +) -> crate::Result { + let cell = memtable_cell(mt, sys_idx, row_idx)?; + system_time(cell).ok_or_else(|| crate::Error::SegmentCorrupted { + detail: format!( + "timeseries memtable column {} row {row_idx} holds {cell:?}, not a system time", + mt.schema().columns[sys_idx].0 + ), + }) +} - for (col_i, (col_name, col_type)) in schema_columns.iter().enumerate() { - let Some(data) = &col_data[col_i] else { - continue; +/// The decoded columns of one partition, read cell by cell. +struct PartitionCells<'a> { + part_dir: &'a std::path::Path, + schema_columns: &'a [(String, ColumnType)], + col_data: &'a [Option], + sym_dicts: &'a HashMap, +} + +impl PartitionCells<'_> { + /// Cell `row_idx` of column `col_i`. A column the partition read did not + /// decode, and a cell that does not read as its declared type, are + /// corruption: each files one report. + fn cell(&self, col_i: usize, row_idx: usize) -> crate::Result> { + let (col_name, col_type) = &self.schema_columns[col_i]; + let Some(data) = self.col_data.get(col_i).and_then(Option::as_ref) else { + return Err(partition_corrupt( + self.part_dir, + "column", + SITE, + format!("column '{col_name}' was not decoded"), + )); }; - let val = partition_col_to_value(data, col_type, col_i, sym_dicts, row_idx) - .map_err(|e| instant_read_error(col_name, e))?; - map.insert(col_name.clone(), val); + read_ts_cell( + data, + *col_type, + col_name, + self.sym_dicts.get(&col_i), + row_idx, + ) + .map_err(|e| partition_corrupt(self.part_dir, "cell", SITE, e)) } - encode_row_map(map, row_idx) + /// The system time of row `row_idx`, read from column `sys_idx`. + fn system_time(&self, sys_idx: usize, row_idx: usize) -> crate::Result { + let cell = self.cell(sys_idx, row_idx)?; + system_time(cell).ok_or_else(|| { + partition_corrupt( + self.part_dir, + "cell", + SITE, + format!( + "column {} row {row_idx} holds {cell:?}, not a system time", + self.schema_columns[sys_idx].0 + ), + ) + }) + } + + /// Encode row `row_idx` as msgpack `Value::Object` bytes. + fn encode_row(&self, row_idx: usize) -> crate::Result> { + let mut map: HashMap = HashMap::with_capacity(self.schema_columns.len()); + for (col_i, (col_name, _)) in self.schema_columns.iter().enumerate() { + let cell = self.cell(col_i, row_idx)?; + map.insert(col_name.clone(), cell_to_value(cell, col_name)?); + } + encode_row_map(map, row_idx) + } } /// Serialize one row map to msgpack bytes. @@ -352,91 +420,24 @@ fn instant_read_error(column: &str, e: nodedb_types::NdbDateTimeError) -> crate: // Column-to-Value converters // --------------------------------------------------------------------------- -/// Convert a memtable column entry to `nodedb_types::Value`. -fn memtable_col_to_value( - col_data: &ColumnData, - col_type: &ColumnType, - col_idx: usize, - mt: &crate::engine::timeseries::columnar_memtable::ColumnarMemtable, - row_idx: usize, -) -> Result { - let value = match col_type { - ColumnType::Timestamp(kind) => { - return kind.cell_value(col_data.as_timestamps()[row_idx]); - } - ColumnType::Float64 => { - let v = col_data.as_f64()[row_idx]; - if v.is_nan() { - Value::Null - } else { - Value::Float(v) - } - } - ColumnType::Int64 => Value::Integer(col_data.as_i64()[row_idx]), - ColumnType::Symbol => { - let sym_id = col_data.as_symbols()[row_idx]; - mt.symbol_dict(col_idx) - .and_then(|d| d.get(sym_id)) - .map(|s| Value::String(s.to_string())) - .unwrap_or(Value::Null) - } - }; - Ok(value) -} - -/// Convert a partition column entry to `nodedb_types::Value`. -fn partition_col_to_value( - data: &ColumnData, - col_type: &ColumnType, - col_i: usize, - sym_dicts: &HashMap, - row_idx: usize, -) -> Result { - let value = match col_type { - ColumnType::Timestamp(kind) => { - return kind.cell_value(data.as_timestamps()[row_idx]); - } - ColumnType::Float64 => { - let v = data.as_f64()[row_idx]; - if v.is_nan() { - Value::Null - } else { - Value::Float(v) - } - } - ColumnType::Int64 => { - if let ColumnData::Int64(vals) = data { - Value::Integer(vals[row_idx]) - } else { - Value::Null - } - } - ColumnType::Symbol => { - if let ColumnData::Symbol(ids) = data { - sym_dicts - .get(&col_i) - .and_then(|dict| dict.get(ids[row_idx])) - .map(|s| Value::String(s.to_string())) - .unwrap_or(Value::Null) - } else { - Value::Null - } - } - }; - Ok(value) -} - -/// Extract a timestamp value from a column (for `_ts_system` filtering). -fn ts_memtable_value(col_data: &ColumnData, row_idx: usize) -> i64 { - match col_data { - ColumnData::Timestamp(v) | ColumnData::Int64(v) => v.get(row_idx).copied().unwrap_or(0), - _ => 0, - } +/// The `Value` a timeseries cell of column `column` emits. +fn cell_to_value(cell: TsCell<'_>, column: &str) -> crate::Result { + Ok(match cell { + TsCell::Null => Value::Null, + TsCell::Time(kind, millis) => kind + .cell_value(millis) + .map_err(|e| instant_read_error(column, e))?, + TsCell::Float(f) => Value::Float(f), + TsCell::Int(n) => Value::Integer(n), + TsCell::Symbol(s) => Value::String(s.to_string()), + }) } -fn ts_partition_value(col_data: &ColumnData, row_idx: usize) -> i64 { - match col_data { - ColumnData::Timestamp(v) | ColumnData::Int64(v) => v.get(row_idx).copied().unwrap_or(0), - _ => 0, +/// The millisecond system time a `_ts_system` cell holds. `None` for any +/// cell other than a time or an integer. +fn system_time(cell: TsCell<'_>) -> Option { + match cell { + TsCell::Time(_, millis) | TsCell::Int(millis) => Some(millis), + TsCell::Null | TsCell::Float(_) | TsCell::Symbol(_) => None, } } diff --git a/nodedb/src/data/executor/handlers/columnar_read/mod.rs b/nodedb/src/data/executor/handlers/columnar_read/mod.rs index 5fde1a0de..25a99ccdd 100644 --- a/nodedb/src/data/executor/handlers/columnar_read/mod.rs +++ b/nodedb/src/data/executor/handlers/columnar_read/mod.rs @@ -9,6 +9,7 @@ pub mod bitemporal; pub mod convert; pub mod filter; +pub mod flushed_segment; pub mod materialize_scan; pub mod materialize_scan_ts; pub mod scan; diff --git a/nodedb/src/data/executor/handlers/columnar_read/scan_flushed.rs b/nodedb/src/data/executor/handlers/columnar_read/scan_flushed.rs index eb7df962a..e4a45869a 100644 --- a/nodedb/src/data/executor/handlers/columnar_read/scan_flushed.rs +++ b/nodedb/src/data/executor/handlers/columnar_read/scan_flushed.rs @@ -6,7 +6,6 @@ use crate::bridge::expr_eval::ComputedColumn; use crate::bridge::scan_filter::ScanFilter; use crate::data::executor::core_loop::CoreLoop; use crate::data::executor::handlers::transaction::overlay::ColumnarMatchedRow; -use crate::data::executor::scan_normalize::decoded_col_to_value; use super::bitemporal::bitemporal_row_visible; use super::convert::row_to_projected_value; @@ -43,6 +42,9 @@ impl CoreLoop { /// type, and each row is projected by `row_to_projected_value` — the same /// two steps the live-memtable phase applies — so a row reads identically /// from a segment and from the memtable. + /// + /// `Err` when a segment does not open or decode, or a live row holds a + /// corrupt cell. No segment is skipped. pub(in crate::data::executor) fn scan_flushed_columnar_segments( &self, ctx: FlushedScanCtx<'_>, @@ -76,64 +78,16 @@ impl CoreLoop { // active memtable virtual segment). Mirror: materialize_scan.rs. let seg_id = seg_idx as u64 + 1; - let reader = if let Some(ref reg) = self.quarantine_registry { - match crate::storage::quarantine::engines::open_segment_with_quarantine( - reg, - seg_bytes, - collection, - &seg_id.to_string(), - ) { - Ok(r) => r, - Err(e) => { - tracing::warn!( - collection, - seg_id, - error = %e, - "execute_columnar_scan: failed to open flushed segment (quarantine); skipping" - ); - continue; - } - } - } else { - match nodedb_columnar::SegmentReader::open(seg_bytes) { - Ok(r) => r, - Err(e) => { - tracing::warn!( - collection, - seg_id, - error = %e, - "execute_columnar_scan: failed to open flushed segment; skipping" - ); - continue; - } - } - }; - - let row_count = reader.row_count() as usize; - let col_count = schema.columns.len(); - - // Decode all columns for this segment up front. - let mut decoded_cols = Vec::with_capacity(col_count); - let mut decode_ok = true; - for col_idx in 0..col_count { - match reader.read_column(col_idx) { - Ok(dc) => decoded_cols.push(dc), - Err(e) => { - tracing::warn!( - collection, - seg_id, - col_idx, - error = %e, - "execute_columnar_scan: column decode failed; skipping segment" - ); - decode_ok = false; - break; - } - } - } - if !decode_ok { - continue; - } + // A segment that does not read refuses the scan. Skipping it + // would answer the query without the segment's rows. + let segment = self.decode_flushed_segment( + collection, + seg_id, + seg_bytes, + schema.columns.len(), + "columnar_scan", + )?; + let row_count = segment.row_count(); // Fetch the delete bitmap for this segment once per segment. let delete_bm = self @@ -177,13 +131,7 @@ impl CoreLoop { // Build the row as Vec using the shared decoder, // typing each cell by its declared column type. - let row: Vec = decoded_cols - .iter() - .zip(&schema.columns) - .map(|(dc, col_def)| { - decoded_col_to_value(dc, row_idx, &col_def.column_type) - }) - .collect(); + let row = segment.row(schema, row_idx)?; if !bitemporal_row_visible( &row, diff --git a/nodedb/src/data/executor/handlers/columnar_write/flush.rs b/nodedb/src/data/executor/handlers/columnar_write/flush.rs index 64314f876..19c0cc485 100644 --- a/nodedb/src/data/executor/handlers/columnar_write/flush.rs +++ b/nodedb/src/data/executor/handlers/columnar_write/flush.rs @@ -1,119 +1,255 @@ // SPDX-License-Identifier: BUSL-1.1 -//! Post-insert memtable flush: drains the columnar memtable to a segment -//! once the flush threshold is reached, retaining encoded segment bytes -//! and their surrogate sidecar in memory. +//! Post-insert memtable flush: encodes the columnar memtable to a segment +//! once the flush threshold is reached, retains the segment bytes and their +//! surrogate sidecar in memory, then drains the memtable. +//! +//! The segment encodes from the memtable in place. The memtable drains only +//! after the segment bytes are retained. An encode error leaves every row in +//! the memtable, so the rows stay readable and the next flush retries them. + +use nodedb_columnar::memtable::DICT_ENCODE_MAX_CARDINALITY; +use nodedb_columnar::{ColumnarError, MutationEngine, SegmentWriter}; +use nodedb_types::Surrogate; +use nodedb_wal::crypto::WalEncryptionKey; use crate::bridge::envelope::{ErrorCode, Response}; use crate::data::executor::core_loop::CoreLoop; use crate::data::executor::task::ExecutionTask; +/// What one memtable flush moved to a segment. +#[derive(Debug, PartialEq, Eq)] +pub(in crate::data::executor) struct FlushedMemtable { + /// The segment id the memtable rows now live under. + pub segment_id: u64, + /// The number of rows the segment holds. `0` when the memtable was empty + /// and no segment was written. + pub row_count: usize, +} + +/// Encode `engine`'s memtable to a segment, hand the segment bytes and the +/// index-aligned row surrogates to `retain`, then drain the memtable and +/// remap its rows to the new segment id. +/// +/// An encode error returns before `retain` runs and before the drain, so the +/// memtable keeps every row. An empty memtable writes no segment. +pub(in crate::data::executor) fn flush_memtable_to_segment( + engine: &mut MutationEngine, + writer: &SegmentWriter, + kek: Option<&WalEncryptionKey>, + retain: impl FnOnce(Vec, Vec>), +) -> Result { + let segment_id = engine.next_segment_id(); + let row_count = engine.memtable().row_count(); + if row_count > 0 { + // Dictionary encoding rewrites low-cardinality string columns in + // place. The memtable reads a dictionary column like a plain one, so + // the rows stay readable if the encode below fails. + engine + .memtable_mut() + .try_dict_encode_columns(DICT_ENCODE_MAX_CARDINALITY); + let memtable = engine.memtable(); + let bytes = writer.write_segment(memtable.schema(), memtable.columns(), row_count, kek)?; + // The surrogates are index-aligned with the memtable rows the segment + // holds. `on_memtable_flushed` below clears them. + retain(bytes, engine.memtable_surrogates().to_vec()); + engine.memtable_mut().drain(); + } + engine.on_memtable_flushed(segment_id)?; + Ok(FlushedMemtable { + segment_id, + row_count, + }) +} + impl CoreLoop { /// Flush the columnar memtable at `engine_key` to a segment if the /// flush threshold has been reached. No-op otherwise. + /// + /// An encode error fails the write with the memtable rows intact. The + /// write's rows are applied and logged, so the error is `Internal`, which + /// the write funnel treats as a write that can have landed. pub(in crate::data::executor) fn flush_columnar_memtable_if_needed( &mut self, task: &ExecutionTask, engine_key: &(nodedb_types::DatabaseId, crate::types::TenantId, String), collection: &str, ) -> Result<(), Response> { - let engine = match self.columnar_engines.get_mut(engine_key) { - Some(e) => e, - None => { - return Err(self.response_error( - task, - ErrorCode::Internal { - detail: "columnar engine missing after insert loop".into(), - }, - )); - } + let Some(engine) = self.columnar_engines.get_mut(engine_key) else { + return Err(self.response_error( + task, + ErrorCode::Internal { + detail: "columnar engine missing after insert loop".into(), + }, + )); }; + if !engine.should_flush() { + return Ok(()); + } + + let writer = SegmentWriter::new( + nodedb_columnar::writer::PROFILE_PLAIN, + nodedb_mem::ScopedMemory::new( + self.governor.clone(), + engine_key.0, + engine_key.1, + nodedb_mem::EngineId::Columnar, + ), + ); + let kek = self.segment_keks.columnar_segment_kek.as_ref(); + let flushed_segments = &mut self.columnar_flushed_segments; + let flushed_surrogates = &mut self.columnar_flushed_surrogates; + // Lockstep invariant: `retain` pushes to BOTH maps for the same key in + // the same order, so the segment-bytes Vec and the surrogate sidecar + // stay equal-length and index-aligned (outer index == segment index, + // segment_id == index + 1). An encode error pushes to neither. + let outcome = flush_memtable_to_segment(engine, &writer, kek, |bytes, surrogates| { + flushed_segments + .entry(engine_key.clone()) + .or_default() + .push(bytes); + flushed_surrogates + .entry(engine_key.clone()) + .or_default() + .push(surrogates); + }); - // Flush memtable to a segment if the threshold has been reached. - if engine.should_flush() { - let new_segment_id = engine.next_segment_id(); - let (schema, columns, row_count) = engine.memtable_mut().drain_optimized(); - // Capture the memtable's per-row surrogates BEFORE `on_memtable_flushed` - // clears them. `drain_optimized` drains the row data but leaves - // `memtable_surrogates` intact; only `on_memtable_flushed` (below) - // clears it. This snapshot is the pre-clear, index-aligned identity - // table for the rows we are about to encode into the segment. - let flushed_surrogates: Vec> = - engine.memtable_surrogates().to_vec(); - if row_count > 0 { - let kek = self.segment_keks.columnar_segment_kek.as_ref(); - let memory = nodedb_mem::ScopedMemory::new( - self.governor.clone(), - engine_key.0, - engine_key.1, - nodedb_mem::EngineId::Columnar, + match outcome { + Ok(flushed) => { + tracing::debug!( + core = self.core_id, + %collection, + new_segment_id = flushed.segment_id, + row_count = flushed.row_count, + "columnar memtable flushed and segment bytes retained in memory" ); - match nodedb_columnar::SegmentWriter::new( - nodedb_columnar::writer::PROFILE_PLAIN, - memory, - ) - .write_segment(&schema, &columns, row_count, kek) - { - Ok(bytes) => { - // Lockstep invariant: push to BOTH maps for the same key in - // the SAME order so the segment-bytes Vec and the surrogate - // sidecar stay equal-length and index-aligned (outer index - // == segment index; segment_id == index + 1). On the Err - // branch below we push to NEITHER, preserving lockstep. - self.columnar_flushed_segments - .entry(engine_key.clone()) - .or_default() - .push(bytes); - self.columnar_flushed_surrogates - .entry(engine_key.clone()) - .or_default() - .push(flushed_surrogates); - tracing::debug!( - core = self.core_id, - %collection, - new_segment_id, - row_count, - "columnar memtable flushed and segment bytes retained in memory" - ); - } - Err(e) => { - // The memtable was already drained above, so these rows - // are no longer in memory and were NOT encoded to a - // segment. We must not continue and report success: on - // the sync path that would call `sync_commit` + return - // `AckStatus::Applied`, telling the client durably-lost - // data was applied and advancing the HWM so the retry is - // never re-admitted. Fail hard instead — the HWM stays - // put and the client (or SQL caller) retries. - tracing::error!( - core = self.core_id, - %collection, - new_segment_id, - row_count, - error = %e, - "columnar segment encode failed; drained rows not durable, failing the write" - ); - return Err(self.response_error( - task, - ErrorCode::Internal { - detail: format!( - "columnar segment encode failed, {row_count} rows not durable: {e}" - ), - }, - )); - } - } + Ok(()) } - if let Err(e) = engine.on_memtable_flushed(new_segment_id) { - return Err(self.response_error( + Err(e) => { + tracing::error!( + core = self.core_id, + %collection, + error = %e, + "columnar memtable flush failed; the rows stay in the memtable" + ); + Err(self.response_error( task, ErrorCode::Internal { - detail: format!("columnar flush: segment ID counter exhausted: {e}"), + detail: format!( + "columnar memtable flush of '{collection}' failed, \ + the rows stay in the memtable: {e}" + ), }, - )); + )) } } + } +} + +#[cfg(test)] +mod tests { + use std::sync::Arc; + + use nodedb_columnar::writer::PROFILE_PLAIN; + use nodedb_mem::{EngineId, EngineLimits, GovernorConfig, MemoryGovernor, ScopedMemory}; + use nodedb_types::columnar::{ColumnDef, ColumnType, ColumnarSchema}; + use nodedb_types::{DatabaseId, TenantId, Value}; + + use super::*; + + /// A writer whose columnar budget is `per_engine` bytes. + fn writer(per_engine: usize) -> SegmentWriter { + let governor = Arc::new( + MemoryGovernor::new(GovernorConfig { + global_ceiling: per_engine * EngineId::ALL.len(), + engine_limits: EngineLimits::uniform(per_engine), + }) + .expect("test governor"), + ); + SegmentWriter::new( + PROFILE_PLAIN, + ScopedMemory::new( + governor, + DatabaseId::DEFAULT, + TenantId::new(0), + EngineId::Columnar, + ), + ) + } + + /// An engine whose memtable holds rows `(1, "a")`, `(2, "b")`, `(3, "a")`. + fn engine_with_rows() -> MutationEngine { + let schema = ColumnarSchema::new(vec![ + ColumnDef::required("id", ColumnType::Int64).with_primary_key(), + ColumnDef::required("tag", ColumnType::String), + ]) + .expect("schema"); + let mut engine = MutationEngine::new("flush_test".to_string(), schema); + for (id, tag) in [(1, "a"), (2, "b"), (3, "a")] { + engine + .insert(&[Value::Integer(id), Value::String(tag.to_string())]) + .expect("insert"); + } + engine + } + + fn memtable_rows(engine: &MutationEngine) -> Vec> { + engine + .scan_memtable_rows() + .map(|row| row.expect("read memtable row")) + .collect() + } + + #[test] + fn an_encode_error_keeps_every_memtable_row() { + let mut engine = engine_with_rows(); + let before = memtable_rows(&engine); + let segment_id = engine.next_segment_id(); + + let mut retained = false; + let err = flush_memtable_to_segment(&mut engine, &writer(1), None, |_, _| { + retained = true; + }) + .expect_err("a one-byte columnar budget refuses the segment encode"); + + assert!( + matches!(err, ColumnarError::BudgetExhausted(_)), + "unexpected error: {err:?}" + ); + assert!(!retained, "a failed encode retains no segment"); + assert_eq!(engine.memtable().row_count(), 3); + assert_eq!(memtable_rows(&engine), before); + assert_eq!( + engine.next_segment_id(), + segment_id, + "a failed flush allocates no segment id" + ); + } + + #[test] + fn a_flush_retains_the_segment_then_drains() { + let mut engine = engine_with_rows(); + let segment_id = engine.next_segment_id(); - Ok(()) + let mut retained: Option<(Vec, usize)> = None; + let flushed = flush_memtable_to_segment( + &mut engine, + &writer(usize::MAX / EngineId::COUNT), + None, + |bytes, surrogates| retained = Some((bytes, surrogates.len())), + ) + .expect("flush"); + + assert_eq!( + flushed, + FlushedMemtable { + segment_id, + row_count: 3, + } + ); + let (bytes, surrogate_count) = retained.expect("the segment is retained"); + assert!(!bytes.is_empty()); + assert_eq!(surrogate_count, 3); + assert_eq!(engine.memtable().row_count(), 0); } } diff --git a/nodedb/src/data/executor/handlers/columnar_write/read_prior.rs b/nodedb/src/data/executor/handlers/columnar_write/read_prior.rs index 5981d636c..bc34a5d80 100644 --- a/nodedb/src/data/executor/handlers/columnar_write/read_prior.rs +++ b/nodedb/src/data/executor/handlers/columnar_write/read_prior.rs @@ -2,87 +2,88 @@ //! Flushed-segment PK lookup for ON CONFLICT DO UPDATE prior-row reads. -use std::mem::size_of; - -use nodedb_types::decode_bounds::checked_decode_capacity; use nodedb_types::value::Value; use crate::data::executor::core_loop::CoreLoop; impl CoreLoop { /// Read the live row bound to `pk_bytes`, wherever it lives: the memtable - /// first, then a flushed segment. `None` when the PK is unbound. + /// first, then a flushed segment. `Ok(None)` when the PK is unbound. + /// + /// `Err` when the bound row does not read. A corrupt row is never + /// reported as an absent one: an upsert would then insert a duplicate + /// key, and a write policy would decide against no prior row. pub(in crate::data::executor) fn read_columnar_row_by_pk( &self, engine_key: &(nodedb_types::DatabaseId, crate::types::TenantId, String), pk_bytes: &[u8], - ) -> Option> { - self.columnar_engines - .get(engine_key) - .and_then(|e| e.lookup_memtable_row_by_pk(pk_bytes)) - .or_else(|| self.read_flushed_row_by_pk(engine_key, pk_bytes)) + ) -> crate::Result>> { + let Some(engine) = self.columnar_engines.get(engine_key) else { + return Ok(None); + }; + if let Some(row) = engine.lookup_memtable_row_by_pk(pk_bytes)? { + return Ok(Some(row)); + } + self.read_flushed_row_by_pk(engine_key, pk_bytes) } /// Read a single row from a flushed columnar segment by PK, if the PK - /// index points to one. Returns `None` when the PK lives in the - /// memtable, when the segment is not in memory, or when the row was - /// tombstoned. Used by the `ON CONFLICT DO UPDATE` path to locate a - /// prior row that has already been flushed out of the memtable. + /// index points to one. `Ok(None)` when the PK is unbound, lives in the + /// memtable, or was tombstoned. Used by the `ON CONFLICT DO UPDATE` and + /// write-policy paths to read a prior row already flushed out of the + /// memtable. + /// + /// `Err` when the PK index points at a segment or row this core does not + /// hold, or when the segment does not decode. The shared segment reader + /// files the corruption report. pub(in crate::data::executor) fn read_flushed_row_by_pk( &self, engine_key: &(nodedb_types::DatabaseId, crate::types::TenantId, String), pk_bytes: &[u8], - ) -> Option> { - let engine = self.columnar_engines.get(engine_key)?; - let loc = engine.pk_index().get(pk_bytes).copied()?; + ) -> crate::Result>> { + let Some(engine) = self.columnar_engines.get(engine_key) else { + return Ok(None); + }; + let Some(loc) = engine.pk_index().get(pk_bytes).copied() else { + return Ok(None); + }; // Memtable case is already covered by the engine-side lookup. if loc.segment_id == engine.memtable_segment_id() { - return None; + return Ok(None); } // Tombstoned — prior row no longer logically present. if engine .delete_bitmap(loc.segment_id) .is_some_and(|bm| bm.is_deleted(loc.row_index)) { - return None; + return Ok(None); } - let segs = self.columnar_flushed_segments.get(engine_key)?; - // Segments are pushed in order starting at segment_id=1. - let seg_idx = (loc.segment_id as usize).checked_sub(1)?; - let seg_bytes = segs.get(seg_idx)?; - let seg_id_str = loc.segment_id.to_string(); - let reader = if let Some(reg) = &self.quarantine_registry { - crate::storage::quarantine::engines::open_segment_with_quarantine( - reg, - seg_bytes, - &engine_key.2, - &seg_id_str, - ) - .ok()? - } else { - nodedb_columnar::SegmentReader::open(seg_bytes).ok()? + let collection = engine_key.2.as_str(); + let unheld = || crate::Error::Internal { + detail: format!( + "columnar '{collection}': the PK index binds a row to flushed segment {} row {}, \ + which this core does not hold", + loc.segment_id, loc.row_index + ), }; + // Segments are pushed in order starting at segment_id=1. + let seg_bytes = usize::try_from(loc.segment_id) + .ok() + .and_then(|id| id.checked_sub(1)) + .and_then(|seg_idx| self.columnar_flushed_segments.get(engine_key)?.get(seg_idx)) + .ok_or_else(unheld)?; let schema = engine.schema(); - // Schema cardinality is catalog-owned rather than byte-decoded; bind - // it once so the allocation's direct trusted bound is explicit. - let column_count = schema.columns.len(); - let row_capacity = checked_decode_capacity( - column_count, - size_of::(), - column_count, - 1, - column_count, - usize::MAX, + let segment = self.decode_flushed_segment( + collection, + loc.segment_id, + seg_bytes, + schema.columns.len(), + "columnar_prior_row", )?; - let mut row = Vec::with_capacity(row_capacity); - for (col_idx, col_def) in schema.columns.iter().enumerate() { - let decoded = reader.read_column(col_idx).ok()?; - row.push(crate::data::executor::scan_normalize::decoded_col_to_value( - &decoded, - loc.row_index as usize, - &col_def.column_type, - )); + let row_idx = loc.row_index as usize; + if row_idx >= segment.row_count() { + return Err(unheld()); } - Some(row) + segment.row(schema, row_idx).map(Some) } } diff --git a/nodedb/src/data/executor/handlers/columnar_write/row_ingest.rs b/nodedb/src/data/executor/handlers/columnar_write/row_ingest.rs index 060654b3c..97c6841a8 100644 --- a/nodedb/src/data/executor/handlers/columnar_write/row_ingest.rs +++ b/nodedb/src/data/executor/handlers/columnar_write/row_ingest.rs @@ -3,6 +3,7 @@ //! Core row-ingest path: per-row value coercion, ON CONFLICT DO UPDATE merge //! resolution, and the row-level `MutationEngine` insert call. +use nodedb_columnar::{BatchConflict, BatchRow}; use nodedb_types::columnar::ColumnarSchema; use nodedb_types::columnar::schema::{TS_SYSTEM, TS_VALID_FROM, TS_VALID_UNTIL}; use nodedb_types::surrogate::Surrogate; @@ -59,8 +60,9 @@ impl CoreLoop { /// `InsertUnique` on a PK the index or an earlier row of the batch /// already carries). /// - /// Every row is resolved and checked before any row is written, so a - /// refusal applies nothing. Returns the accepted row count (and, on + /// Every row is resolved and checked before any row is written. The + /// engine then writes the whole batch or none of it, so a refusal at any + /// stage applies nothing. Returns the accepted row count (and, on /// request, the stored post-images), or `Err(Response)` on the first /// error. pub(in crate::data::executor) fn insert_columnar_rows( @@ -75,52 +77,44 @@ impl CoreLoop { collect_stored_rows, .. } = params; - let mut accepted = 0u64; - let mut stored_rows: Vec> = Vec::new(); + let conflict = match intent { + ColumnarInsertIntent::InsertIfAbsent => BatchConflict::Skip, + ColumnarInsertIntent::InsertUnique + | ColumnarInsertIntent::Insert + | ColumnarInsertIntent::Put => BatchConflict::Upsert, + }; - for row in resolved { - let engine = match self.columnar_engines.get_mut(engine_key) { - Some(e) => e, - None => { - return Err(self.response_error( - task, - ErrorCode::Internal { - detail: "columnar engine vanished during insert".into(), - }, - )); - } - }; - let result = match intent { - ColumnarInsertIntent::InsertIfAbsent => engine.insert_if_absent(&row.values), - ColumnarInsertIntent::InsertUnique - | ColumnarInsertIntent::Insert - | ColumnarInsertIntent::Put => match row.surrogate { - Some(s) => engine.insert_with_surrogate(&row.values, s), - None => engine.insert(&row.values), + let Some(engine) = self.columnar_engines.get_mut(engine_key) else { + return Err(self.response_error( + task, + ErrorCode::Internal { + detail: "columnar engine vanished during insert".into(), }, - }; + )); + }; + let batch = resolved.iter().map(|row| BatchRow { + values: &row.values, + surrogate: row.surrogate, + }); + let results = match engine.insert_batch(batch, conflict) { + Ok(results) => results, + Err(e) => { + return Err(self.response_error(task, ErrorCode::from(crate::Error::from(e)))); + } + }; - match result { - // An `insert_if_absent` that hit an existing key returns an - // EMPTY `wal_records` — that is the engine's documented no-op - // signal, and the only way to tell a skip from a write. Counting - // it reported an `INSERT 1` for a row that was never stored, and - // returning it would hand back a row that does not exist. - Ok(mutation) if mutation.wal_records.is_empty() => {} - Ok(_) => { - accepted += 1; - if collect_stored_rows { - stored_rows.push(row.values); - } - } - Err(e) => { - return Err(self.response_error( - task, - ErrorCode::Internal { - detail: format!("columnar insert failed: {e}"), - }, - )); - } + // A skipped row returns an EMPTY `wal_records`: the engine's no-op + // signal, and the only way to tell a skip from a write. A skipped + // row is neither counted nor returned, as it was never stored. + let mut accepted = 0u64; + let mut stored_rows: Vec> = Vec::new(); + for (row, result) in resolved.into_iter().zip(results) { + if result.wal_records.is_empty() { + continue; + } + accepted += 1; + if collect_stored_rows { + stored_rows.push(row.values); } } @@ -162,9 +156,18 @@ impl CoreLoop { std::collections::HashMap::new(); for (row_idx, row) in ndb_rows.iter().enumerate() { - let obj = match row { - nodedb_types::Value::Object(m) => m, - _ => continue, + // A row that is not an object has no fields to write. Skipping it + // would report a statement that stored fewer rows than it sent. + let nodedb_types::Value::Object(obj) = row else { + return Err(self.response_error( + task, + crate::Error::BadRequest { + detail: format!( + "columnar insert into '{}': row {row_idx} is not an object", + engine_key.2 + ), + }, + )); }; // Build Value slice in schema order. For bitemporal @@ -192,19 +195,12 @@ impl CoreLoop { Some(Value::Integer(i)) => Value::Integer(*i), _ => Value::Integer(i64::MAX), }), - _ => ndb_field_to_value(obj.get(&col.name), &col.column_type), + _ => ndb_field_to_value(obj.get(&col.name), col), }) .collect::, crate::Error>>() { Ok(v) => v, - Err(e) => { - return Err(self.response_error( - task, - ErrorCode::Internal { - detail: format!("columnar insert coercion: {e}"), - }, - )); - } + Err(e) => return Err(self.response_error(task, ErrorCode::from(e))), }; let pk_bytes = if merging || unique { @@ -222,12 +218,9 @@ impl CoreLoop { match engine.encode_pk_from_row(&values) { Ok(b) => b, Err(e) => { - return Err(self.response_error( - task, - ErrorCode::Internal { - detail: format!("columnar insert: pk encode failed: {e}"), - }, - )); + return Err( + self.response_error(task, ErrorCode::from(crate::Error::from(e))) + ); } } } else { @@ -238,12 +231,16 @@ impl CoreLoop { // UPDATE, plain otherwise. let final_values: Vec = match intent { ColumnarInsertIntent::Put if merging => { - let prior_row = batch_rows.get(&pk_bytes).cloned().or_else(|| { - self.columnar_engines - .get(engine_key) - .and_then(|e| e.lookup_memtable_row_by_pk(&pk_bytes)) - .or_else(|| self.read_flushed_row_by_pk(engine_key, &pk_bytes)) - }); + // A prior row that does not read refuses the statement: + // merging against "no prior row" would insert a + // duplicate key. + let prior_row = match batch_rows.get(&pk_bytes) { + Some(row) => Some(row.clone()), + None => match self.read_columnar_row_by_pk(engine_key, &pk_bytes) { + Ok(row) => row, + Err(e) => return Err(self.response_error(task, e)), + }, + }; match prior_row { None => values, Some(prior) => self.merge_on_conflict( @@ -350,16 +347,9 @@ impl CoreLoop { schema .columns .iter() - .map(|col| ndb_field_to_value(merged_obj.get(&col.name), &col.column_type)) + .map(|col| ndb_field_to_value(merged_obj.get(&col.name), col)) .collect::, crate::Error>>() - .map_err(|e| { - self.response_error( - task, - ErrorCode::Internal { - detail: format!("columnar ON CONFLICT coercion: {e}"), - }, - ) - }) + .map_err(|e| self.response_error(task, ErrorCode::from(e))) } } @@ -476,6 +466,82 @@ mod tests { ); } + /// An upsert whose prior row lives in a flushed segment that does not + /// read is refused. Reading the corrupt row as absent would insert a + /// second row under the same key. + #[test] + fn an_upsert_over_a_corrupt_prior_row_is_refused() { + let dir = tempfile::tempdir().expect("tempdir"); + let (mut core, _tx, _rx) = make_core_with_dir(dir.path()); + let seeded = insert(&mut core, ColumnarInsertIntent::Insert, vec![row("a")]); + assert_eq!(seeded.status, Status::Ok, "{:?}", seeded.error_code); + + // Flush row "a" out of the memtable, then hold segment bytes that do + // not open in its place. + let key = ( + DatabaseId::DEFAULT, + TenantId::new(TID), + COLLECTION.to_string(), + ); + let engine = core.columnar_engines.get_mut(&key).expect("engine"); + let segment_id = engine.next_segment_id(); + let _drained = engine.memtable_mut().drain_optimized(); + let sidecar = engine.memtable_surrogates().to_vec(); + engine.on_memtable_flushed(segment_id).expect("flush"); + core.columnar_flushed_segments + .insert(key.clone(), vec![Vec::new()]); + core.columnar_flushed_surrogates.insert(key, vec![sidecar]); + let before = live_rows(&core); + + let payload = + nodedb_types::value_to_msgpack(&Value::Array(vec![row("a")])).expect("encode rows"); + let schema = schema_bytes(); + let updates = vec![( + "v".to_string(), + nodedb_physical::physical_plan::UpdateValue::Literal( + nodedb_types::value_to_msgpack(&Value::Integer(2)).expect("encode"), + ), + )]; + let refused = core.execute_columnar_insert( + &task(), + ColumnarInsertParams { + collection: COLLECTION, + payload: &payload, + format: "msgpack", + intent: ColumnarInsertIntent::Put, + on_conflict_updates: &updates, + surrogates: &[], + schema_bytes: &schema, + provenance: None, + rls_write_check: &RlsWriteCheck::already_decided_elsewhere(), + returning: None, + rls_filters: &[], + spatial_undo: None, + }, + ); + + assert_ne!(refused.status, Status::Ok, "the upsert is refused"); + assert_eq!(live_rows(&core), before); + let bound = core + .columnar_engines + .get(&( + DatabaseId::DEFAULT, + TenantId::new(TID), + COLLECTION.to_string(), + )) + .expect("engine") + .pk_index() + .get(&nodedb_columnar::pk_index::encode_pk(&Value::String( + "a".into(), + ))) + .copied() + .expect("key stays bound"); + assert_eq!( + bound.segment_id, segment_id, + "the key still names the flushed row: no replacement row was written" + ); + } + #[test] fn a_key_repeated_inside_a_unique_batch_writes_no_row() { let dir = tempfile::tempdir().expect("tempdir"); diff --git a/nodedb/src/data/executor/handlers/timeseries/aggregate.rs b/nodedb/src/data/executor/handlers/timeseries/aggregate.rs index 023bb735c..ef347768b 100644 --- a/nodedb/src/data/executor/handlers/timeseries/aggregate.rs +++ b/nodedb/src/data/executor/handlers/timeseries/aggregate.rs @@ -15,7 +15,6 @@ use crate::bridge::envelope::Response; use crate::data::executor::core_loop::CoreLoop; use crate::data::executor::task::ExecutionTask; -use crate::engine::timeseries::grouped_filter::UnsupportedPredicate; use crate::engine::timeseries::grouped_scan::{ GroupedAggResult, PartitionAggParams, aggregate_memtable, aggregate_partition, }; @@ -95,13 +94,14 @@ impl CoreLoop { let data_dir = &self.data_dir; let db_id = task.request.database_id.as_u64(); let tenant = tid.as_u64(); + // Every listed partition is read. A missing directory refuses + // the aggregate in `aggregate_partition`. let partition_dirs: Vec = entries .iter() .map(|e| { super::paths::ts_collection_dir(data_dir, db_id, tenant, collection) .join(&e.dir_name) }) - .filter(|p| p.exists()) .collect(); if partition_dirs.len() <= 1 { @@ -123,7 +123,7 @@ impl CoreLoop { }) { Ok(Some(part_result)) => merged.merge(&part_result), Ok(None) => {} - Err(e) => return self.response_error(task, crate::Error::from(e)), + Err(e) => return self.response_error(task, e), } } } else { @@ -140,7 +140,7 @@ impl CoreLoop { let thread_count = available.min(partition_dirs.len()).min(8); let chunk_size = partition_dirs.len().div_ceil(thread_count); - let partition_results: Result, UnsupportedPredicate> = + let partition_results: crate::Result> = std::thread::scope(|s| { let handles: Vec<_> = partition_dirs .chunks(chunk_size) @@ -149,7 +149,7 @@ impl CoreLoop { let ag = &agg_owned; let fl = &filters_owned; let nc = &needed_owned; - s.spawn(move || -> Result<_, UnsupportedPredicate> { + s.spawn(move || -> crate::Result<_> { let mut local = GroupedAggResult::new(ag.len()); for dir in chunk { // Parallel threads: no io_uring (fadvise fallback). @@ -175,12 +175,25 @@ impl CoreLoop { }) .collect(); - handles.into_iter().filter_map(|h| h.join().ok()).collect() + // A worker that panicked aggregated none of its + // partitions: the aggregate is refused, never + // answered without them. + handles + .into_iter() + .map(|h| { + h.join().map_err(|_| crate::Error::Internal { + detail: format!( + "timeseries aggregate on '{collection}': a partition \ + worker panicked" + ), + })? + }) + .collect() }); let partition_results = match partition_results { Ok(results) => results, - Err(e) => return self.response_error(task, crate::Error::from(e)), + Err(e) => return self.response_error(task, e), }; for r in &partition_results { merged.merge(r); diff --git a/nodedb/src/data/executor/handlers/timeseries/cell_read.rs b/nodedb/src/data/executor/handlers/timeseries/cell_read.rs new file mode 100644 index 000000000..197356720 --- /dev/null +++ b/nodedb/src/data/executor/handlers/timeseries/cell_read.rs @@ -0,0 +1,159 @@ +// SPDX-License-Identifier: BUSL-1.1 + +//! One typed read of a timeseries cell, for the timeseries row emitters. +//! +//! A column whose stored data does not match its declared type, a row past +//! the end of its column, and a symbol id its dictionary does not hold are +//! all errors. None of them reads as NULL. A NULL float is stored as NaN and +//! a NULL symbol as [`NULL_SYMBOL_ID`]. + +use nodedb_types::timeseries::SymbolDictionary; + +use crate::engine::timeseries::columnar_memtable::{ColumnData, ColumnType, TimeKind}; + +/// The symbol id a NULL symbol cell stores. +pub(in crate::data::executor) const NULL_SYMBOL_ID: u32 = u32::MAX; + +/// One timeseries cell, read as its declared column type. +#[derive(Debug, Clone, Copy, PartialEq)] +pub(in crate::data::executor) enum TsCell<'a> { + Null, + /// A time cell: its kind and the stored millisecond count. + Time(TimeKind, i64), + Float(f64), + Int(i64), + Symbol(&'a str), +} + +/// A timeseries cell that does not read as its declared type. +#[derive(Debug, thiserror::Error)] +pub(in crate::data::executor) enum TsCellError { + #[error("column '{column}' holds {stored} data, declared {declared:?}")] + TypeMismatch { + column: String, + stored: &'static str, + declared: ColumnType, + }, + #[error("column '{column}' has no row {row}")] + RowOutOfRange { column: String, row: usize }, + #[error("column '{column}' row {row} holds symbol {id}, which its dictionary does not hold")] + UnknownSymbol { column: String, row: usize, id: u32 }, +} + +impl From for crate::Error { + fn from(e: TsCellError) -> Self { + crate::Error::SegmentCorrupted { + detail: e.to_string(), + } + } +} + +/// Read row `row` of `data`, a column named `column` declared `declared`. +/// `dict` is the column's symbol dictionary, when it has one. +pub(in crate::data::executor) fn read_ts_cell<'a>( + data: &'a ColumnData, + declared: ColumnType, + column: &str, + dict: Option<&'a SymbolDictionary>, + row: usize, +) -> Result, TsCellError> { + let out_of_range = || TsCellError::RowOutOfRange { + column: column.to_string(), + row, + }; + match (declared, data) { + (ColumnType::Timestamp(kind), ColumnData::Timestamp(v)) => { + let millis = *v.get(row).ok_or_else(out_of_range)?; + Ok(TsCell::Time(kind, millis)) + } + (ColumnType::Float64, ColumnData::Float64(v)) => { + let f = *v.get(row).ok_or_else(out_of_range)?; + Ok(if f.is_nan() { + TsCell::Null + } else { + TsCell::Float(f) + }) + } + (ColumnType::Int64, ColumnData::Int64(v)) => { + Ok(TsCell::Int(*v.get(row).ok_or_else(out_of_range)?)) + } + (ColumnType::Symbol, ColumnData::Symbol(ids)) => { + let id = *ids.get(row).ok_or_else(out_of_range)?; + if id == NULL_SYMBOL_ID { + return Ok(TsCell::Null); + } + dict.and_then(|d| d.get(id)) + .map(TsCell::Symbol) + .ok_or_else(|| TsCellError::UnknownSymbol { + column: column.to_string(), + row, + id, + }) + } + (declared, stored) => Err(TsCellError::TypeMismatch { + column: column.to_string(), + stored: stored_kind(stored), + declared, + }), + } +} + +/// The name of the data shape a column holds. +fn stored_kind(data: &ColumnData) -> &'static str { + match data { + ColumnData::Timestamp(_) => "timestamp", + ColumnData::Float64(_) => "float64", + ColumnData::Int64(_) => "int64", + ColumnData::Symbol(_) => "symbol", + ColumnData::DictEncoded { .. } => "dictionary-encoded", + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn a_column_of_the_wrong_shape_is_an_error() { + let data = ColumnData::Float64(vec![1.0]); + assert!(matches!( + read_ts_cell(&data, ColumnType::Int64, "n", None, 0), + Err(TsCellError::TypeMismatch { .. }) + )); + } + + #[test] + fn a_row_past_the_column_is_an_error() { + let data = ColumnData::Int64(vec![1]); + assert!(matches!( + read_ts_cell(&data, ColumnType::Int64, "n", None, 1), + Err(TsCellError::RowOutOfRange { row: 1, .. }) + )); + } + + #[test] + fn a_symbol_the_dictionary_lacks_is_an_error_and_the_null_id_is_null() { + let data = ColumnData::Symbol(vec![7, NULL_SYMBOL_ID]); + assert!(matches!( + read_ts_cell(&data, ColumnType::Symbol, "tag", None, 0), + Err(TsCellError::UnknownSymbol { id: 7, .. }) + )); + assert_eq!( + read_ts_cell(&data, ColumnType::Symbol, "tag", None, 1).expect("null"), + TsCell::Null + ); + } + + #[test] + fn a_nan_float_is_null() { + let data = ColumnData::Float64(vec![f64::NAN, 2.5]); + assert_eq!( + read_ts_cell(&data, ColumnType::Float64, "v", None, 0).expect("null"), + TsCell::Null + ); + assert_eq!( + read_ts_cell(&data, ColumnType::Float64, "v", None, 1).expect("float"), + TsCell::Float(2.5) + ); + } +} diff --git a/nodedb/src/data/executor/handlers/timeseries/ingest_resolved_returning.rs b/nodedb/src/data/executor/handlers/timeseries/ingest_resolved_returning.rs new file mode 100644 index 000000000..89b0e3cc2 --- /dev/null +++ b/nodedb/src/data/executor/handlers/timeseries/ingest_resolved_returning.rs @@ -0,0 +1,88 @@ +// SPDX-License-Identifier: BUSL-1.1 + +//! The `RETURNING` rows of a resolved timeseries ingest, decoded from the +//! landed rows' images. +//! +//! Each landed row's image is the row as a scan reads it. An image that +//! does not decode is corruption of the resolved record. The statement is +//! refused with the corruption error and one corruption report. It never +//! returns fewer rows than it stored. + +/// The site name the corruption report carries. +const SITE: &str = "timeseries_resolved_returning"; + +/// Decode every landed row's image, in row order. +/// +/// `Err(SegmentCorrupted)` for the first image that does not decode, after +/// filing its corruption report. +pub(super) fn decode_returning_images( + collection: &str, + images: &[&[u8]], +) -> crate::Result> { + images + .iter() + .enumerate() + .map(|(row, image)| { + crate::util::bounded_msgpack::read_value(image).map_err(|e| { + let err = crate::Error::SegmentCorrupted { + detail: format!( + "timeseries '{collection}': the image of landed row {row} does not \ + decode: {e}" + ), + }; + crate::diag::timeseries_row_image_undecodable(&err, collection, SITE); + err + }) + }) + .collect() +} + +#[cfg(test)] +mod tests { + use super::*; + + fn image(value: &rmpv::Value) -> Vec { + let mut bytes = Vec::new(); + rmpv::encode::write_value(&mut bytes, value).expect("encode image"); + bytes + } + + #[test] + fn every_image_decodes_in_row_order() { + let first = rmpv::Value::Map(vec![("v".into(), 1.into())]); + let second = rmpv::Value::Map(vec![("v".into(), 2.into())]); + let images = [image(&first), image(&second)]; + let refs: Vec<&[u8]> = images.iter().map(Vec::as_slice).collect(); + + let rows = decode_returning_images("ts", &refs).expect("decode"); + + assert_eq!(rows, vec![first, second]); + } + + #[test] + fn an_undecodable_image_refuses_the_statement() { + let good = image(&rmpv::Value::Map(vec![("v".into(), 1.into())])); + // 0xc1 is a reserved MessagePack marker, never a value. + let bad: &[u8] = &[0xc1]; + let refs: Vec<&[u8]> = vec![good.as_slice(), bad]; + + let err = decode_returning_images("ts", &refs).expect_err("row 1 does not decode"); + + assert!( + matches!(&err, crate::Error::SegmentCorrupted { detail } if detail.contains("row 1")), + "{err:?}" + ); + } + + #[test] + fn an_empty_image_is_not_a_row() { + let refs: Vec<&[u8]> = vec![&[]]; + + let err = decode_returning_images("ts", &refs).expect_err("an empty image is no row"); + + assert!( + matches!(err, crate::Error::SegmentCorrupted { .. }), + "{err:?}" + ); + } +} diff --git a/nodedb/src/data/executor/handlers/timeseries/mod.rs b/nodedb/src/data/executor/handlers/timeseries/mod.rs index f73a2c6b7..15bcb6cf9 100644 --- a/nodedb/src/data/executor/handlers/timeseries/mod.rs +++ b/nodedb/src/data/executor/handlers/timeseries/mod.rs @@ -4,6 +4,7 @@ mod admission; pub mod aggregate; +pub mod cell_read; pub mod encode; mod events; pub mod flush; @@ -14,9 +15,11 @@ pub mod ingest_formats; mod ingest_resolved; mod ingest_resolved_fit; mod ingest_resolved_outcome; +mod ingest_resolved_returning; mod ingest_schema; mod msgpack_decode; mod normalize; +pub mod partition_read; pub mod paths; pub mod raw_scan; mod redo_ingest; diff --git a/nodedb/src/data/executor/handlers/timeseries/partition_read.rs b/nodedb/src/data/executor/handlers/timeseries/partition_read.rs new file mode 100644 index 000000000..bce00af3a --- /dev/null +++ b/nodedb/src/data/executor/handlers/timeseries/partition_read.rs @@ -0,0 +1,124 @@ +// SPDX-License-Identifier: BUSL-1.1 + +//! Whole-partition column reads for the timeseries scan paths. +//! +//! The partition writer writes `schema.json` and one `.col` file per schema +//! column, so a schema or column that does not read is corruption. A `.sym` +//! file is written only for a symbol column that has a dictionary, so an +//! absent one is no dictionary, while a present one that does not read is +//! corruption. Corruption refuses the read and files one report here. No +//! caller skips a partition or a column it cannot read. + +use std::collections::HashMap; +use std::path::Path; + +use nodedb_types::timeseries::SymbolDictionary; + +use crate::engine::timeseries::columnar_memtable::{ColumnData, ColumnType, ColumnarSchema}; +use crate::engine::timeseries::columnar_segment::{ColumnarSegmentReader, SegmentError}; + +/// Every column of one on-disk timeseries partition, read whole. +pub(in crate::data::executor) struct TsPartitionColumns { + pub schema: ColumnarSchema, + /// One entry per schema column, in schema order. Every entry is `Some`: + /// the `Option` matches the partition-row emitters this feeds. + pub columns: Vec>, + /// Dictionary of each symbol column that has one, by column index. + pub sym_dicts: HashMap, +} + +/// The corruption error for the partition at `part_dir`, with one report +/// filed. Every timeseries path that finds a listed partition unreadable +/// calls this at the site that finds it. `stage` names what failed and +/// `site` names the read path. +pub(crate) fn partition_corrupt( + part_dir: &Path, + stage: &'static str, + site: &'static str, + detail: impl std::fmt::Display, +) -> crate::Error { + let err = crate::Error::SegmentCorrupted { + detail: format!("timeseries partition {}: {detail}", part_dir.display()), + }; + let partition = part_dir + .file_name() + .and_then(|n| n.to_str()) + .unwrap_or_default(); + crate::diag::timeseries_partition_unreadable(&err, partition, stage, site); + err +} + +/// The error for a partition the registry lists and the disk lacks, or +/// `Ok` when its directory is present. +pub(crate) fn require_partition_dir(part_dir: &Path, site: &'static str) -> crate::Result<()> { + if part_dir.is_dir() { + return Ok(()); + } + Err(partition_corrupt( + part_dir, + "directory", + site, + "the partition registry lists it, and its directory is missing", + )) +} + +/// Read the schema, every column, and every symbol dictionary of the +/// partition at `part_dir`. `site` names the read path in the report. +/// +/// `Err` when the directory is missing: the caller reads only partitions +/// the registry lists. +pub(in crate::data::executor) fn read_ts_partition( + part_dir: &Path, + site: &'static str, +) -> crate::Result { + require_partition_dir(part_dir, site)?; + let report = + |stage: &'static str, err: SegmentError| partition_corrupt(part_dir, stage, site, err); + let schema = + ColumnarSegmentReader::read_schema(part_dir, None).map_err(|e| report("schema", e))?; + // Prefetch all column files into page cache before reading. + let all_col_names: Vec = schema.columns.iter().map(|(n, _)| n.clone()).collect(); + crate::data::io::fadvise::prefetch_partition_columns(part_dir, &all_col_names); + let columns = schema + .columns + .iter() + .map(|(name, ty)| ColumnarSegmentReader::read_column(part_dir, name, *ty, None).map(Some)) + .collect::, _>>() + .map_err(|e| report("column", e))?; + let mut sym_dicts = HashMap::new(); + for (i, (name, ty)) in schema.columns.iter().enumerate() { + if *ty != ColumnType::Symbol || !part_dir.join(format!("{name}.sym")).exists() { + continue; + } + let dict = ColumnarSegmentReader::read_symbol_dict(part_dir, name, None) + .map_err(|e| report("symbol_dict", e))?; + sym_dicts.insert(i, dict); + } + Ok(TsPartitionColumns { + schema, + columns, + sym_dicts, + }) +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn a_partition_without_a_schema_is_refused_as_corruption() { + let dir = tempfile::tempdir().expect("tempdir"); + let result = read_ts_partition(dir.path(), "test"); + assert!(matches!(result, Err(crate::Error::SegmentCorrupted { .. }))); + } + + #[test] + fn a_listed_partition_without_a_directory_is_refused_as_corruption() { + let dir = tempfile::tempdir().expect("tempdir"); + let result = read_ts_partition(&dir.path().join("ts-0001"), "test"); + assert!(matches!( + result, + Err(crate::Error::SegmentCorrupted { ref detail }) if detail.contains("directory is missing") + )); + } +} diff --git a/nodedb/src/data/executor/handlers/timeseries/raw_scan/partition_scan.rs b/nodedb/src/data/executor/handlers/timeseries/raw_scan/partition_scan.rs index 135c66dfb..8c360030c 100644 --- a/nodedb/src/data/executor/handlers/timeseries/raw_scan/partition_scan.rs +++ b/nodedb/src/data/executor/handlers/timeseries/raw_scan/partition_scan.rs @@ -2,12 +2,12 @@ //! Parallel and sequential disk-partition scanning for raw mode. -use std::collections::HashMap; - use crate::bridge::scan_filter::ScanFilter; +use crate::data::executor::handlers::timeseries::partition_read::{ + TsPartitionColumns, read_ts_partition, +}; use crate::engine::timeseries::columnar_agg::timestamp_range_filter; -use crate::engine::timeseries::columnar_memtable::{ColumnData, ColumnType}; -use crate::engine::timeseries::columnar_segment::ColumnarSegmentReader; +use crate::engine::timeseries::columnar_memtable::ColumnType; use super::row_emit::{emit_partition_row, extract_timestamp, row_admitted}; @@ -78,9 +78,15 @@ pub(super) fn scan_partitions_parallel( }) .collect(); + // A scan thread that panicked refuses the scan: dropping its + // result would answer without its partitions' rows. handles .into_iter() - .filter_map(|h| h.join().ok()) + .map(|h| { + h.join().map_err(|_| crate::Error::Internal { + detail: "timeseries partition scan thread panicked".into(), + })? + }) .collect::>>() })?; @@ -153,36 +159,23 @@ pub(super) fn scan_one_partition( has_filters: bool, rls_predicates: &[ScanFilter], ) -> crate::Result> { - let schema = match ColumnarSegmentReader::read_schema(part_dir, None) { - Ok(s) => s, - Err(_) => return Ok(Vec::new()), - }; - - // Prefetch all column files into page cache before reading. + // A partition that does not read refuses the scan. Skipping it would + // answer without the partition's rows. + let TsPartitionColumns { + schema, + columns: col_data, + sym_dicts, + } = read_ts_partition(part_dir, "timeseries_raw_scan")?; let all_col_names: Vec = schema.columns.iter().map(|(n, _)| n.clone()).collect(); - crate::data::io::fadvise::prefetch_partition_columns(part_dir, &all_col_names); - - let col_data: Vec> = schema - .columns - .iter() - .map(|(name, ty)| ColumnarSegmentReader::read_column(part_dir, name, *ty, None).ok()) - .collect(); - - let sym_dicts: HashMap = schema - .columns - .iter() - .enumerate() - .filter(|(_, (_, ty))| *ty == ColumnType::Symbol) - .filter_map(|(i, (name, _))| { - ColumnarSegmentReader::read_symbol_dict(part_dir, name, None) - .ok() - .map(|dict| (i, dict)) - }) - .collect(); - - let ts_col = col_data.get(schema.timestamp_idx).and_then(|d| d.as_ref()); - let Some(ts_col) = ts_col else { - return Ok(Vec::new()); + + let Some(ts_col) = col_data.get(schema.timestamp_idx).and_then(|d| d.as_ref()) else { + return Err(crate::Error::SegmentCorrupted { + detail: format!( + "timeseries partition {}: time column index {} is outside its schema", + part_dir.display(), + schema.timestamp_idx + ), + }); }; let timestamps = ts_col.as_timestamps(); let indices = timestamp_range_filter(timestamps, time_range.0, time_range.1); @@ -229,7 +222,7 @@ pub(super) fn scan_one_partition( if rows.len() >= limit { break; } - let row = emit_partition_row(&schema_vec, &col_data, &sym_dicts, idx as usize)?; + let row = emit_partition_row(part_dir, &schema_vec, &col_data, &sym_dicts, idx as usize)?; if !row_admitted(&row, row_filters, rls_predicates)? { continue; } diff --git a/nodedb/src/data/executor/handlers/timeseries/raw_scan/row_emit.rs b/nodedb/src/data/executor/handlers/timeseries/raw_scan/row_emit.rs index 8fe5f2100..7a6a19569 100644 --- a/nodedb/src/data/executor/handlers/timeseries/raw_scan/row_emit.rs +++ b/nodedb/src/data/executor/handlers/timeseries/raw_scan/row_emit.rs @@ -3,12 +3,18 @@ //! Row emission helpers — build `rmpv::Value` directly. use std::collections::HashMap; +use std::path::Path; use nodedb_types::columnar::schema::TS_SYSTEM; use crate::bridge::scan_filter::ScanFilter; use crate::data::executor::handlers::columnar_read::filter::value_matches_filters; use crate::data::executor::handlers::columnar_read::{emit_column_value, rmpv_time_cell}; +use crate::data::executor::handlers::timeseries::cell_read::{TsCell, read_ts_cell}; +use crate::data::executor::handlers::timeseries::partition_read::partition_corrupt; + +/// The read path the corruption report names. +const SITE: &str = "timeseries_raw_scan"; use crate::engine::timeseries::columnar_memtable::{ColumnData, ColumnType}; use crate::util::rmpv_value::{rmpv_to_value, value_to_rmpv}; @@ -111,13 +117,21 @@ pub(super) fn emit_memtable_row( nodedb_query::msgpack_scan::write_str(&mut buf, col_name); emit_column_value(&mut buf, mt, *col_idx, col_type, col_data, idx)?; } - Ok(crate::util::bounded_msgpack::read_value(&buf).unwrap_or(rmpv::Value::Nil)) + // The row is built above, so a decode error is a bug in that build, never + // a row to answer as NULL. + crate::util::bounded_msgpack::read_value(&buf).map_err(|e| crate::Error::Internal { + detail: format!("timeseries memtable row {idx} does not decode after it is built: {e}"), + }) } /// Emit a single row from a disk partition as rmpv::Value::Map. /// /// `Err` when a stored time cell cannot be read as its column's instant. +/// A column the partition read did not decode, and a cell that does not +/// read as its declared type, are corruption: each files one report naming +/// `part_dir`. pub(super) fn emit_partition_row( + part_dir: &Path, schema: &[(String, ColumnType)], col_data: &[Option], sym_dicts: &HashMap, @@ -125,47 +139,22 @@ pub(super) fn emit_partition_row( ) -> crate::Result { let mut fields: Vec<(rmpv::Value, rmpv::Value)> = Vec::with_capacity(schema.len()); for (col_i, (col_name, col_type)) in schema.iter().enumerate() { - // A column whose file could not be read is emitted as NULL, never - // skipped. Skipping it changed the row's COLUMN SET rather than one - // cell's value, so `SELECT *` on the same row returned different - // columns before and after a flush — the memtable path below always - // emits every column. A missing value is NULL; it is not a missing - // column. - let Some(data) = &col_data[col_i] else { - fields.push(( - rmpv::Value::String(col_name.as_str().into()), - rmpv::Value::Nil, + let Some(data) = col_data.get(col_i).and_then(Option::as_ref) else { + return Err(partition_corrupt( + part_dir, + "column", + SITE, + format!("column '{col_name}' was not decoded"), )); - continue; }; - let val = match col_type { - ColumnType::Timestamp(kind) => rmpv_time_cell(*kind, data.as_timestamps()[idx])?, - ColumnType::Float64 => { - let v = data.as_f64()[idx]; - if v.is_nan() { - rmpv::Value::Nil - } else { - rmpv::Value::F64(v) - } - } - ColumnType::Int64 => { - if let ColumnData::Int64(vals) = data { - rmpv::Value::Integer(vals[idx].into()) - } else { - rmpv::Value::Nil - } - } - ColumnType::Symbol => { - if let ColumnData::Symbol(ids) = data { - sym_dicts - .get(&col_i) - .and_then(|dict| dict.get(ids[idx])) - .map(|s| rmpv::Value::String(s.into())) - .unwrap_or(rmpv::Value::Nil) - } else { - rmpv::Value::Nil - } - } + let cell = read_ts_cell(data, *col_type, col_name, sym_dicts.get(&col_i), idx) + .map_err(|e| partition_corrupt(part_dir, "cell", SITE, e))?; + let val = match cell { + TsCell::Time(kind, millis) => rmpv_time_cell(kind, millis)?, + TsCell::Float(v) => rmpv::Value::F64(v), + TsCell::Int(n) => rmpv::Value::Integer(n.into()), + TsCell::Symbol(s) => rmpv::Value::String(s.into()), + TsCell::Null => rmpv::Value::Nil, }; fields.push((rmpv::Value::String(col_name.as_str().into()), val)); } diff --git a/nodedb/src/data/executor/handlers/timeseries_wal.rs b/nodedb/src/data/executor/handlers/timeseries_wal.rs index 8105d65b9..bb888678a 100644 --- a/nodedb/src/data/executor/handlers/timeseries_wal.rs +++ b/nodedb/src/data/executor/handlers/timeseries_wal.rs @@ -533,7 +533,8 @@ mod tests { .get(&(DatabaseId::new(0), TenantId::new(7), "m".to_string())) .expect("engine") .scan_memtable_rows() - .collect(); + .collect::>() + .expect("read"); assert_eq!(rows, vec![vec![Value::Integer(1), Value::Integer(10)]]); } diff --git a/nodedb/src/data/executor/handlers/transaction/stage_write/stage_rls.rs b/nodedb/src/data/executor/handlers/transaction/stage_write/stage_rls.rs index 9f972ab3c..63dbeca41 100644 --- a/nodedb/src/data/executor/handlers/transaction/stage_write/stage_rls.rs +++ b/nodedb/src/data/executor/handlers/transaction/stage_write/stage_rls.rs @@ -106,9 +106,10 @@ impl CoreLoop { .map_err(|e| ErrorCode::Internal { detail: format!("columnar insert: pk encode failed: {e}"), })?; - engine - .lookup_memtable_row_by_pk(&pk_bytes) - .or_else(|| self.read_flushed_row_by_pk(engine_key, &pk_bytes)) + // A prior row that does not read refuses the statement: the + // policy would otherwise decide against no prior row. + self.read_columnar_row_by_pk(engine_key, &pk_bytes) + .map_err(ErrorCode::from)? } }; let Some(prior) = prior else { @@ -127,10 +128,8 @@ impl CoreLoop { schema .columns .iter() - .map(|col| ndb_field_to_value(merged.get(&col.name), &col.column_type)) + .map(|col| ndb_field_to_value(merged.get(&col.name), col)) .collect::, crate::Error>>() - .map_err(|e| ErrorCode::Internal { - detail: format!("columnar ON CONFLICT coercion: {e}"), - }) + .map_err(ErrorCode::from) } } diff --git a/nodedb/src/data/executor/wal_replay_all.rs b/nodedb/src/data/executor/wal_replay_all.rs index 02b01d801..258edc698 100644 --- a/nodedb/src/data/executor/wal_replay_all.rs +++ b/nodedb/src/data/executor/wal_replay_all.rs @@ -328,6 +328,7 @@ mod tests { .get(&key) .expect("engine") .scan_memtable_rows() + .map(|row| row.expect("read")) .filter_map(|row| match row.first() { Some(Value::Integer(id)) => Some(*id), _ => None, diff --git a/nodedb/src/diag/context/columnar.rs b/nodedb/src/diag/context/columnar.rs new file mode 100644 index 000000000..686234728 --- /dev/null +++ b/nodedb/src/diag/context/columnar.rs @@ -0,0 +1,88 @@ +// SPDX-License-Identifier: BUSL-1.1 + +//! Forensic payloads for columnar segment capture sites. + +use faultbox::DomainContext; +use faultbox::serde_json::{Value, json}; + +/// A flushed columnar segment whose bytes do not open or decode. +pub(in crate::diag) struct ColumnarSegmentCorrupt<'a> { + /// Collection that owns the segment. + pub collection: &'a str, + /// 1-based flushed segment id. + pub segment_id: u64, + /// Decode step that failed: `open`, `column`, or `cell`. + pub stage: &'static str, + /// Path that read the segment. + pub site: &'static str, + /// The error class: the error text before its first colon. + pub error_class: &'a str, +} + +impl DomainContext for ColumnarSegmentCorrupt<'_> { + fn domain_kind(&self) -> &'static str { + "nodedb.columnar_segment_corrupt" + } + + fn grouping_key(&self) -> String { + // The segment is the root cause. The site and row are occurrences, + // so every read of one bad segment files one report with a count. + format!( + "collection={} segment={} stage={}", + self.collection, self.segment_id, self.stage + ) + } + + fn to_json(&self) -> Value { + json!({ + "collection": self.collection, + "segment_id": self.segment_id, + "stage": self.stage, + "site": self.site, + "error_class": self.error_class, + "why_fatal": "the segment is the only in-memory home of its rows between \ + checkpoints, and a checkpoint persists the same bytes. Every \ + read that touches it is refused until it is replaced", + "operator_action": "restore the collection from a snapshot taken before the \ + segment was damaged. A repeated CRC failure on the same \ + segment also places it in quarantine", + }) + } +} + +/// An on-disk timeseries partition whose files do not read. +pub(in crate::diag) struct TimeseriesPartitionUnreadable<'a> { + /// Partition directory name. + pub partition: &'a str, + /// File that failed: `schema`, `column`, or `symbol_dict`. + pub stage: &'static str, + /// Path that read the partition. + pub site: &'static str, + /// The error class: the error text before its first colon. + pub error_class: &'a str, +} + +impl DomainContext for TimeseriesPartitionUnreadable<'_> { + fn domain_kind(&self) -> &'static str { + "nodedb.timeseries_partition_unreadable" + } + + fn grouping_key(&self) -> String { + // The partition is the root cause. The site is the occurrence. + format!("partition={} stage={}", self.partition, self.stage) + } + + fn to_json(&self) -> Value { + json!({ + "partition": self.partition, + "stage": self.stage, + "site": self.site, + "error_class": self.error_class, + "why_fatal": "the partition files are the only copy of its rows, so every \ + scan that reaches it is refused until it is repaired", + "operator_action": "inspect the named partition directory for a missing or \ + damaged file. Restore the collection from a snapshot if \ + the file cannot be recovered", + }) + } +} diff --git a/nodedb/src/diag/recording/columnar.rs b/nodedb/src/diag/recording/columnar.rs new file mode 100644 index 000000000..874ff79c6 --- /dev/null +++ b/nodedb/src/diag/recording/columnar.rs @@ -0,0 +1,66 @@ +// SPDX-License-Identifier: BUSL-1.1 + +//! Capture sites for columnar segment reads. + +use faultbox::{Capture, EventKind, error_chain_of}; + +use super::shared::error_class; +use crate::diag::context; + +/// Report a flushed columnar segment that does not open or decode. +/// +/// Called only from the shared flushed-segment readers in +/// `columnar_read/flushed_segment.rs`, which every segment read goes through. +/// The caller returns the error alongside this report. +pub fn columnar_segment_corrupt( + err: &crate::Error, + collection: &str, + segment_id: u64, + stage: &'static str, + site: &'static str, +) { + let class = error_class(err); + let ctx = context::ColumnarSegmentCorrupt { + collection, + segment_id, + stage, + site, + error_class: &class, + }; + let _ = Capture::new( + EventKind::Corruption, + "flushed columnar segment does not decode, so the read is refused", + ) + .error_chain(error_chain_of(err)) + .domain(&ctx) + .with_backtrace() + .emit(); +} + +/// Report an on-disk timeseries partition whose schema, column, or symbol +/// dictionary does not read. +/// +/// Called only from `read_ts_partition`. The caller returns the error +/// alongside this report. +pub fn timeseries_partition_unreadable( + err: &crate::Error, + partition: &str, + stage: &'static str, + site: &'static str, +) { + let class = error_class(err); + let ctx = context::TimeseriesPartitionUnreadable { + partition, + stage, + site, + error_class: &class, + }; + let _ = Capture::new( + EventKind::Corruption, + "timeseries partition does not read, so the scan is refused", + ) + .error_chain(error_chain_of(err)) + .domain(&ctx) + .with_backtrace() + .emit(); +} diff --git a/nodedb/src/engine/timeseries/grouped_scan/partition.rs b/nodedb/src/engine/timeseries/grouped_scan/partition.rs index 670afab00..5abfc8ce2 100644 --- a/nodedb/src/engine/timeseries/grouped_scan/partition.rs +++ b/nodedb/src/engine/timeseries/grouped_scan/partition.rs @@ -17,8 +17,14 @@ use super::strategies::dispatch_grouping; use super::types::{GroupedAggResult, resolve_schema}; use crate::bridge::envelope::Priority; use crate::bridge::scan_filter::ScanFilter; +use crate::data::executor::handlers::timeseries::partition_read::{ + partition_corrupt, require_partition_dir, +}; use crate::data::io::IoMetrics; +/// The read path the corruption report names. +const SITE: &str = "timeseries_grouped_scan"; + /// Aggregate from a columnar memtable with GROUP BY + optional time_bucket. /// /// `Ok(None)` when a GROUP BY or aggregate column is not in the memtable's @@ -119,22 +125,21 @@ pub struct PartitionAggParams<'a> { /// When `uring_reader` is `Some`, column files are batch-read via io_uring /// (parallel kernel I/O). When `None`, falls back to fadvise + std::fs::read. /// -/// `Ok(None)` when the partition's schema or metadata cannot be read or does -/// not carry a GROUP BY / aggregate column. `Err` when a predicate in -/// `filters` cannot be lowered onto the partition's typed columns: the -/// caller fails the statement rather than aggregating rows the predicate -/// never excluded. -pub fn aggregate_partition( - p: PartitionAggParams<'_>, -) -> Result, UnsupportedPredicate> { +/// `Ok(None)` when the partition's schema does not carry a GROUP BY or +/// aggregate column. `Err` when a predicate in `filters` cannot be lowered +/// onto the partition's typed columns: the caller fails the statement +/// rather than aggregating rows the predicate never excluded. `Err` when the +/// partition does not read: its directory, schema, metadata, sparse index, +/// a needed column, or a symbol dictionary. The partition is never skipped. +pub fn aggregate_partition(p: PartitionAggParams<'_>) -> crate::Result> { let num_aggs = p.aggregates.len(); + let dir = p.partition_dir; - let Ok(schema) = ColumnarSegmentReader::read_schema(p.partition_dir, None) else { - return Ok(None); - }; - let Ok(meta) = ColumnarSegmentReader::read_meta(p.partition_dir, None) else { - return Ok(None); - }; + require_partition_dir(dir, SITE)?; + let schema = ColumnarSegmentReader::read_schema(dir, None) + .map_err(|e| partition_corrupt(dir, "schema", SITE, e))?; + let meta = ColumnarSegmentReader::read_meta(dir, None) + .map_err(|e| partition_corrupt(dir, "meta", SITE, e))?; let row_count = meta.row_count as usize; if row_count == 0 { return Ok(Some(GroupedAggResult::new(num_aggs))); @@ -149,10 +154,10 @@ pub fn aggregate_partition( return Ok(None); }; - // Load sparse index for block-level skip. - let sparse_idx = ColumnarSegmentReader::read_sparse_index(p.partition_dir, None) - .ok() - .flatten(); + // Load sparse index for block-level skip. A partition written without + // one has no index file. + let sparse_idx = ColumnarSegmentReader::read_sparse_index(dir, None) + .map_err(|e| partition_corrupt(dir, "sparse_index", SITE, e))?; // Determine surviving blocks (if sparse index available). let surviving_blocks: Option> = sparse_idx @@ -169,17 +174,23 @@ pub fn aggregate_partition( .as_ref() .is_some_and(|sb| !sb.is_empty() && sb.len() < total_blocks); + let has_time_range = p.time_range.0 > 0 || p.time_range.1 < i64::MAX; + // The time column is read whenever a time range or a bucket needs it, + // named in `needed_columns` or not. + let time_column = (has_time_range || p.bucket_interval_ms > 0).then_some(resolved.ts_idx); + let col_data: Vec> = read_partition_columns(ReadPartitionColumnsParams { - partition_dir: p.partition_dir, + partition_dir: dir, schema_columns: &schema.columns, needed_columns: p.needed_columns, + time_column, meta: &meta, use_block_read, surviving_blocks: surviving_blocks.as_deref(), uring_reader: p.uring_reader, io_priority: p.io_priority, io_metrics: p.io_metrics, - }); + })?; // When block-level read was used, row_count is the number of // decoded rows (only surviving blocks), not the partition total. @@ -200,36 +211,29 @@ pub fn aggregate_partition( row_count }; - let sym_dicts: HashMap = schema - .columns - .iter() - .enumerate() - .filter(|(_, (_, ty))| *ty == ColumnType::Symbol) - .filter_map(|(i, (name, _))| { - if p.needed_columns.is_empty() || p.needed_columns.iter().any(|n| n == name) { - ColumnarSegmentReader::read_symbol_dict(p.partition_dir, name, None) - .ok() - .map(|dict| (i, dict)) - } else { - None - } - }) - .collect(); + // A `.sym` file is written only for a symbol column that has a + // dictionary, so an absent one is no dictionary. A present one that does + // not read is corruption. + let mut sym_dicts: HashMap = HashMap::new(); + for (i, (name, ty)) in schema.columns.iter().enumerate() { + let needed = p.needed_columns.is_empty() || p.needed_columns.iter().any(|n| n == name); + if *ty != ColumnType::Symbol || !needed || !dir.join(format!("{name}.sym")).exists() { + continue; + } + let dict = ColumnarSegmentReader::read_symbol_dict(dir, name, None) + .map_err(|e| partition_corrupt(dir, "symbol_dict", SITE, e))?; + sym_dicts.insert(i, dict); + } // Build bitmask over the decoded data. // When block-level read was used, data is already filtered to surviving // blocks — just need predicate + time range filters on the decoded rows. // When full read was used, need time range + sparse skip + predicate. - let has_time_range = p.time_range.0 > 0 || p.time_range.1 < i64::MAX; - let mut mask = if use_block_read { // Block-level read already filtered by sparse index. // Only need time range filter within surviving blocks. if has_time_range { - let Some(ts_col) = col_data.get(resolved.ts_idx).and_then(|d| d.as_ref()) else { - return Ok(None); - }; - let timestamps = ts_col.as_timestamps(); + let timestamps = time_cells(dir, &col_data, resolved.ts_idx)?; let rt = simd_filter::filter_runtime(); (rt.range_i64)(timestamps, p.time_range.0, p.time_range.1) } else { @@ -241,10 +245,7 @@ pub fn aggregate_partition( let m = if partition_fully_in_range { simd_filter::bitmask_all(effective_row_count) } else { - let Some(ts_col) = col_data.get(resolved.ts_idx).and_then(|d| d.as_ref()) else { - return Ok(None); - }; - let timestamps = ts_col.as_timestamps(); + let timestamps = time_cells(dir, &col_data, resolved.ts_idx)?; let rt = simd_filter::filter_runtime(); (rt.range_i64)(timestamps, p.time_range.0, p.time_range.1) }; @@ -286,10 +287,7 @@ pub fn aggregate_partition( }; let timestamps = if p.bucket_interval_ms > 0 { - col_data - .get(resolved.ts_idx) - .and_then(|d| d.as_ref()) - .map(|d| d.as_timestamps()) + Some(time_cells(dir, &col_data, resolved.ts_idx)?) } else { None }; @@ -315,10 +313,38 @@ pub fn aggregate_partition( Ok(Some(result)) } +/// The time column of a partition read, as millisecond cells. +/// +/// `Err` when the column was not decoded or does not hold time cells: the +/// partition is unreadable, and skipping it would drop its rows. +fn time_cells<'a>( + dir: &Path, + col_data: &'a [Option], + ts_idx: usize, +) -> crate::Result<&'a [i64]> { + match col_data.get(ts_idx).and_then(Option::as_ref) { + Some(ColumnData::Timestamp(v)) => Ok(v), + Some(_) => Err(partition_corrupt( + dir, + "column", + SITE, + format!("time column {ts_idx} does not hold time cells"), + )), + None => Err(partition_corrupt( + dir, + "column", + SITE, + format!("time column {ts_idx} was not decoded"), + )), + } +} + struct ReadPartitionColumnsParams<'a> { partition_dir: &'a Path, schema_columns: &'a [(String, super::super::columnar_memtable::ColumnType)], needed_columns: &'a [String], + /// A column read whether `needed_columns` names it or not. + time_column: Option, meta: &'a nodedb_types::timeseries::PartitionMeta, use_block_read: bool, surviving_blocks: Option<&'a [usize]>, @@ -331,11 +357,17 @@ struct ReadPartitionColumnsParams<'a> { /// /// With `uring_reader`: batch-reads all needed `.col` files in parallel /// via io_uring, then decodes each. Without: fadvise + sequential std::fs::read. -fn read_partition_columns(p: ReadPartitionColumnsParams<'_>) -> Vec> { +/// +/// A column that is not needed is `None`. A needed column that does not read +/// is an error: the partition is corrupt. +fn read_partition_columns( + p: ReadPartitionColumnsParams<'_>, +) -> crate::Result>> { let ReadPartitionColumnsParams { partition_dir, schema_columns, needed_columns, + time_column, meta, use_block_read, surviving_blocks, @@ -347,56 +379,50 @@ fn read_partition_columns(p: ReadPartitionColumnsParams<'_>) -> Vec = schema_columns .iter() .enumerate() - .filter(|(_, (name, _))| { - needed_columns.is_empty() || needed_columns.iter().any(|n| n == name) + .filter(|(i, (name, _))| { + needed_columns.is_empty() + || needed_columns.iter().any(|n| n == name) + || time_column == Some(*i) }) .map(|(i, _)| i) .collect(); + // One column read on its own: whole, or its surviving blocks. + let read_one = |i: usize, blocks: bool| -> crate::Result { + let (name, ty) = &schema_columns[i]; + let codec = meta.column_stats.get(name).map(|s| s.codec); + let read = if blocks { + ColumnarSegmentReader::read_column_blocks( + partition_dir, + name, + *ty, + codec, + surviving_blocks.unwrap_or(&[]), + None, + ) + .map(|(data, _)| data) + } else { + ColumnarSegmentReader::read_column_with_codec(partition_dir, name, *ty, codec, None) + }; + read.map_err(|e| partition_corrupt(partition_dir, "column", SITE, e)) + }; + let mut columns: Vec> = (0..schema_columns.len()).map(|_| None).collect(); + // Block-level reads can't use io_uring batching (need per-block decode). // io_uring batching only benefits full-column reads. - if use_block_read || uring_reader.is_none() { - // Fallback: fadvise + sequential read. - crate::data::io::fadvise::prefetch_partition_columns(partition_dir, needed_columns); - - return schema_columns - .iter() - .enumerate() - .map(|(i, (name, ty))| { - if !needed_indices.contains(&i) { - return None; - } - let codec = meta.column_stats.get(name).map(|s| s.codec); - if use_block_read { - ColumnarSegmentReader::read_column_blocks( - partition_dir, - name, - *ty, - codec, - surviving_blocks.unwrap_or(&[]), - None, - ) - .ok() - .map(|(data, _)| data) - } else { - ColumnarSegmentReader::read_column_with_codec( - partition_dir, - name, - *ty, - codec, - None, - ) - .ok() - } - }) - .collect(); - } - - // io_uring path: batch-read all needed .col files in parallel. - let Some(reader) = uring_reader else { - unreachable!("guarded by is_none() check above"); + let reader = match uring_reader { + Some(reader) if !use_block_read => reader, + _ => { + // Fallback: fadvise + sequential read. + crate::data::io::fadvise::prefetch_partition_columns(partition_dir, needed_columns); + for &i in &needed_indices { + columns[i] = Some(read_one(i, use_block_read)?); + } + return Ok(columns); + } }; + // io_uring path: batch-read all needed .col files in parallel. let col_paths: Vec = needed_indices .iter() .map(|&i| partition_dir.join(format!("{}.col", schema_columns[i].0))) @@ -410,30 +436,69 @@ fn read_partition_columns(p: ReadPartitionColumnsParams<'_>) -> Vec reader.read_files(&path_refs), }; - // Decode each raw buffer into ColumnData. - let mut decoded: HashMap = HashMap::new(); + // Decode each raw buffer into ColumnData. The batch returns an empty + // buffer for a file it could not read. That file reads again on its own, + // so the error names why it does not read. for (buf_idx, &schema_idx) in needed_indices.iter().enumerate() { - let raw = &raw_buffers[buf_idx]; - if raw.is_empty() { - continue; - } - let (name, ty) = &schema_columns[schema_idx]; - let codec = meta.column_stats.get(name).map(|s| s.codec); - if let Ok(data) = ColumnarSegmentReader::decode_column_from_bytes( - partition_dir, - name, - *ty, - codec, - raw, - None, - ) { - decoded.insert(schema_idx, data); - } + let raw = raw_buffers + .get(buf_idx) + .map(Vec::as_slice) + .unwrap_or_default(); + let data = if raw.is_empty() { + read_one(schema_idx, false)? + } else { + let (name, ty) = &schema_columns[schema_idx]; + let codec = meta.column_stats.get(name).map(|s| s.codec); + ColumnarSegmentReader::decode_column_from_bytes( + partition_dir, + name, + *ty, + codec, + raw, + None, + ) + .map_err(|e| partition_corrupt(partition_dir, "column", SITE, e))? + }; + columns[schema_idx] = Some(data); } + Ok(columns) +} - schema_columns - .iter() - .enumerate() - .map(|(i, _)| decoded.remove(&i)) - .collect() +#[cfg(test)] +mod tests { + use super::*; + + fn aggregate(dir: &Path) -> crate::Result> { + aggregate_partition(PartitionAggParams { + partition_dir: dir, + group_by: &[], + aggregates: &[("count".to_string(), "*".to_string())], + filters: &[], + time_range: (0, i64::MAX), + needed_columns: &[], + bucket_interval_ms: 0, + uring_reader: None, + io_priority: None, + io_metrics: None, + }) + } + + #[test] + fn a_listed_partition_without_a_directory_refuses_the_aggregate() { + let dir = tempfile::tempdir().expect("tempdir"); + let result = aggregate(&dir.path().join("ts-0001")); + assert!(matches!( + result, + Err(crate::Error::SegmentCorrupted { ref detail }) if detail.contains("directory is missing") + )); + } + + #[test] + fn a_partition_whose_schema_does_not_read_refuses_the_aggregate() { + let dir = tempfile::tempdir().expect("tempdir"); + assert!(matches!( + aggregate(dir.path()), + Err(crate::Error::SegmentCorrupted { .. }) + )); + } } diff --git a/nodedb/src/error_from_columnar.rs b/nodedb/src/error_from_columnar.rs new file mode 100644 index 000000000..23f3820f2 --- /dev/null +++ b/nodedb/src/error_from_columnar.rs @@ -0,0 +1,114 @@ +// SPDX-License-Identifier: BUSL-1.1 + +//! `nodedb_columnar::ColumnarError` into the crate error. +//! +//! A row the caller sent that the memtable cannot hold is the caller's +//! error: `BadRequest`, or the constraint it breaks. Every other columnar +//! error is a storage fault of the engine. + +use nodedb_columnar::error::ColumnarError; + +use crate::Error; + +impl From for Error { + fn from(e: ColumnarError) -> Self { + match e { + ColumnarError::TypeMismatch { .. } + | ColumnarError::JsonParse { .. } + | ColumnarError::RangeParse { .. } + | ColumnarError::MsgpackSerialize { .. } + | ColumnarError::MsgpackDeserialize { .. } => Self::BadRequest { + detail: e.to_string(), + }, + // A stored cell or segment that does not decode is corruption, + // not a fault the caller can retry around. + ColumnarError::Corruption { .. } + | ColumnarError::MemtableCellCorrupt { .. } + | ColumnarError::WalRowCorrupt { .. } + | ColumnarError::StringCellNotUtf8 { .. } + | ColumnarError::FooterCrcMismatch { .. } + | ColumnarError::TruncatedSegment { .. } + | ColumnarError::InvalidMagic(_) => Self::SegmentCorrupted { + detail: e.to_string(), + }, + ColumnarError::NullViolation(ref column) => Self::RejectedConstraint { + collection: String::new(), + constraint: "not_null".into(), + detail: format!("column '{column}' is NOT NULL"), + }, + ColumnarError::DuplicatePrimaryKey => Self::RejectedConstraint { + collection: String::new(), + constraint: "unique".into(), + detail: e.to_string(), + }, + // `ColumnarError` is `#[non_exhaustive]`: any other variant is a + // fault of the engine or its stored segments. + _ => Self::Storage { + engine: "columnar".into(), + detail: e.to_string(), + }, + } + } +} + +impl From for Error { + /// A quarantined segment is corrupt by definition. Any other open error + /// keeps the class `ColumnarError` maps to. + fn from(e: crate::storage::quarantine::engines::ColumnarOrQuarantine) -> Self { + use crate::storage::quarantine::engines::ColumnarOrQuarantine as Coq; + match e { + Coq::Columnar(inner) => Self::from(inner), + Coq::Quarantined(q) => Self::SegmentCorrupted { + detail: q.to_string(), + }, + } + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn a_value_the_column_cannot_hold_is_a_bad_request() { + let e = Error::from(ColumnarError::TypeMismatch { + column: "v".into(), + expected: "Decimal".into(), + }); + assert!(matches!(e, Error::BadRequest { ref detail } if detail.contains("'v'"))); + } + + #[test] + fn constraint_breaks_keep_their_constraint() { + let e = Error::from(ColumnarError::NullViolation("v".into())); + assert!( + matches!(e, Error::RejectedConstraint { ref constraint, .. } if constraint == "not_null") + ); + let e = Error::from(ColumnarError::DuplicatePrimaryKey); + assert!( + matches!(e, Error::RejectedConstraint { ref constraint, .. } if constraint == "unique") + ); + } + + #[test] + fn a_corrupt_cell_is_segment_corruption() { + let e = Error::from(ColumnarError::MemtableCellCorrupt { + column: "v".into(), + row: 0, + reason: "JSON cell is not MessagePack".into(), + }); + assert!(matches!(e, Error::SegmentCorrupted { ref detail } if detail.contains("'v'"))); + let e = Error::from(ColumnarError::Corruption { + segment_id: None, + reason: "bad cell".into(), + offset: None, + }); + assert!(matches!(e, Error::SegmentCorrupted { .. })); + } + + #[test] + fn an_engine_fault_is_a_storage_error() { + let e = Error::from(ColumnarError::EmptyMemtable); + assert!(matches!(e, Error::Storage { ref engine, .. } if engine == "columnar")); + } +} From b024c4554811d06df6e75f297dd2d11831f6ebf4 Mon Sep 17 00:00:00 2001 From: Farhan Syah Date: Thu, 8 Oct 2026 16:51:29 +0800 Subject: [PATCH 02/24] feat(errors): add typed value-refusal codes and keep handler errors typed Add error variants for a numeric value out of range (22003), invalid text representation (22P02), datatype mismatch (42804), invalid datetime format (22007), datetime field overflow (22008), the node-label limit, and a full-text column fault. Each crosses the cluster wire as its own Data-Plane code and renders its SQLSTATE on every transport. RollbackFailed carries the typed cause of the reverse write that failed. Data-Plane handlers convert errors through ErrorCode::from instead of flattening them into Internal strings, and a bitmap sub-plan that fails fails its join instead of reading as an empty bitmap. --- .../src/rpc_codec/data_plane_error.rs | 89 +++++ nodedb-cluster/src/rpc_codec/mod.rs | 4 +- nodedb-sql/src/error.rs | 26 ++ .../src/error/ctors/read_query_auth.rs | 15 + nodedb-types/src/error/sqlstate.rs | 10 + nodedb/src/bridge/envelope/error_code.rs | 337 +++------------- nodedb/src/bridge/envelope/error_code_from.rs | 322 +++++++++++++++ nodedb/src/bridge/envelope/mod.rs | 1 + .../control/cluster/array_executor/refusal.rs | 6 + .../control/cluster/data_plane_error_wire.rs | 365 ++++++++---------- .../control/cluster/data_plane_fault_wire.rs | 74 ++++ .../control/cluster/execution_error_wire.rs | 192 +++++++++ .../control/cluster/metadata_applier/wedge.rs | 6 + nodedb/src/control/cluster/mod.rs | 2 + .../error_map/class_parity/code_samples.rs | 30 +- .../error_map/class_parity/error_index.rs | 8 +- .../error_map/class_parity/error_parity.rs | 5 + .../error_map/class_parity/error_samples.rs | 17 + nodedb/src/control/gateway/error_map/resp.rs | 6 + nodedb/src/control/planner/plan_error_map.rs | 27 +- .../server/dispatch_utils/write_abort.rs | 21 +- .../control/server/pgwire/types/error_map.rs | 20 + .../src/control/server/shared/ddl/result.rs | 7 +- .../src/control/server/shared/ddl/sqlstate.rs | 57 ++- nodedb/src/control/server/shared/retry.rs | 6 + .../sync/async_dispatch/delta/compensation.rs | 15 +- nodedb/src/control/server/sync/refusal.rs | 30 +- .../data/executor/dispatch/array/mutate.rs | 21 +- .../src/data/executor/dispatch/array/open.rs | 7 +- .../dispatch/meta_retention/handlers.rs | 49 +-- .../data/executor/handlers/bulk_dml/update.rs | 50 ++- .../handlers/columnar_read/scan/execute.rs | 20 +- .../handlers/columnar_write/insert.rs | 14 +- .../data/executor/handlers/control/crdt.rs | 152 +++----- .../handlers/control/crdt_constraints.rs | 21 +- .../executor/handlers/control/crdt_list.rs | 42 +- .../executor/handlers/control/reindex/csr.rs | 57 +-- .../handlers/control/reindex/pending.rs | 2 +- .../executor/handlers/control/snapshot.rs | 80 ++-- nodedb/src/data/executor/handlers/convert.rs | 94 +---- .../executor/handlers/document/index_fetch.rs | 30 +- .../handlers/document/index_maintenance.rs | 37 +- .../executor/handlers/document/read/emit.rs | 94 +---- nodedb/src/data/executor/handlers/facet.rs | 48 ++- .../executor/handlers/graph_algo_edges.rs | 25 +- .../data/executor/handlers/graph_temporal.rs | 9 +- .../executor/handlers/join/grace_drive.rs | 14 +- .../executor/handlers/join/hash_handlers.rs | 190 ++++----- .../executor/handlers/join/nested_loop.rs | 14 +- .../data/executor/handlers/join/sort_merge.rs | 14 +- .../src/data/executor/handlers/kv/atomic.rs | 194 +++++++--- nodedb/src/data/executor/handlers/kv/batch.rs | 21 +- .../data/executor/handlers/kv/crud/delete.rs | 14 +- nodedb/src/data/executor/handlers/kv/field.rs | 29 +- nodedb/src/data/executor/handlers/kv/index.rs | 14 +- .../executor/handlers/kv/predicate/apply.rs | 17 +- nodedb/src/data/executor/handlers/kv/ttl.rs | 7 +- .../executor/handlers/timeseries/ingest.rs | 27 +- .../handlers/timeseries/ingest_dispatch.rs | 7 +- .../handlers/timeseries/ingest_resolved.rs | 38 +- .../data/executor/handlers/timeseries/scan.rs | 7 +- .../stage_write/stage_bulk_delete.rs | 7 +- .../stage_write/stage_bulk_update.rs | 7 +- .../transaction/stage_write/stage_columnar.rs | 18 +- .../executor/handlers/truncate_response.rs | 7 +- .../handlers/unregister_collection.rs | 7 +- .../executor/handlers/update_from_join.rs | 35 +- .../executor/handlers/upsert/exec/dispatch.rs | 7 +- .../executor/handlers/vector_lifecycle.rs | 7 +- .../data/executor/handlers/vector_multi.rs | 14 +- .../data/executor/handlers/vector_sparse.rs | 14 +- .../data/executor/handlers/vector_write.rs | 7 +- nodedb/src/data/executor/snapshot/capture.rs | 9 +- .../src/data/executor/strict_format/encode.rs | 48 ++- nodedb/src/error/conversions.rs | 11 + nodedb/src/error/types.rs | 42 ++ nodedb/src/error_classify/public.rs | 16 + nodedb/src/error_classify/unclassified.rs | 6 + nodedb/src/error_from.rs | 25 ++ nodedb/src/error_from_data_plane.rs | 70 +++- nodedb/src/error_from_graph.rs | 65 ++++ 81 files changed, 1951 insertions(+), 1628 deletions(-) create mode 100644 nodedb/src/bridge/envelope/error_code_from.rs create mode 100644 nodedb/src/control/cluster/data_plane_fault_wire.rs create mode 100644 nodedb/src/control/cluster/execution_error_wire.rs create mode 100644 nodedb/src/error_from_graph.rs 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-sql/src/error.rs b/nodedb-sql/src/error.rs index b02210889..4b47cf6ec 100644 --- a/nodedb-sql/src/error.rs +++ b/nodedb-sql/src/error.rs @@ -66,6 +66,16 @@ pub enum SqlError { #[error("unknown column '{column}' in table '{table}'")] UnknownColumn { table: String, column: String }, + /// The column argument of `text_match` / `bm25_score` cannot serve a + /// full-text search. `fault` names why and carries the SQLSTATE. + #[error("{function}(): column '{column}' of collection '{collection}' {fault}")] + TextColumn { + function: String, + collection: String, + column: String, + fault: nodedb_types::text_search::TextColumnFault, + }, + #[error("ambiguous column '{column}' — qualify with table name")] AmbiguousColumn { column: String }, @@ -102,6 +112,13 @@ pub enum SqlError { #[error("arithmetic overflow evaluating a constant expression: {detail}")] ConstantOverflow { detail: String }, + /// An integer literal past the exact numeric range (the 96-bit + /// `Decimal` mantissa). Rounding it to a float would store another + /// number than the one written. PostgreSQL rejects an out-of-range + /// numeric with SQLSTATE `22003`. + #[error("numeric literal {literal} is out of range")] + NumericLiteralOutOfRange { literal: String }, + /// A write supplied an integer wider than the column's declared type. /// /// nodedb stores every integer as an `i64`, so this is not a storage @@ -133,6 +150,15 @@ pub enum SqlError { declared_type: &'static str, }, + /// A write supplied a decimal whose rounded integer part has more digits + /// than the column's declared `DECIMAL(p,s)` holds. PostgreSQL refuses + /// the same write with SQLSTATE `22003`. + #[error("column '{column}': {source}")] + DecimalOutOfRange { + column: String, + source: nodedb_types::columnar::DecimalOutOfRange, + }, + #[error("unsupported: {detail}")] Unsupported { detail: String }, diff --git a/nodedb-types/src/error/ctors/read_query_auth.rs b/nodedb-types/src/error/ctors/read_query_auth.rs index 332b3f499..f1ded740d 100644 --- a/nodedb-types/src/error/ctors/read_query_auth.rs +++ b/nodedb-types/src/error/ctors/read_query_auth.rs @@ -214,6 +214,21 @@ impl NodeDbError { } } + /// A value does not fit its numeric type: an integer past a column's + /// declared width, or arithmetic that overflows. SQLSTATE `22003` + /// (`numeric_value_out_of_range`). The error names no collection. + /// `detail` is the full message. + pub fn numeric_value_out_of_range(detail: impl Into) -> Self { + Self { + code: ErrorCode::OVERFLOW, + message: detail.into(), + details: ErrorDetails::Overflow { + collection: String::new(), + }, + cause: None, + } + } + /// A statement exceeded a server limit on its own size or depth: a /// recursion depth, a per-transaction staging budget. SQLSTATE `54000` /// (`program_limit_exceeded`). `detail` is the full message. diff --git a/nodedb-types/src/error/sqlstate.rs b/nodedb-types/src/error/sqlstate.rs index 3d2b55525..9876c2132 100644 --- a/nodedb-types/src/error/sqlstate.rs +++ b/nodedb-types/src/error/sqlstate.rs @@ -57,6 +57,14 @@ pub const DATA_EXCEPTION: &str = "22000"; /// `22003` — `numeric_value_out_of_range` pub const NUMERIC_VALUE_OUT_OF_RANGE: &str = "22003"; +/// `22007` — `invalid_datetime_format` (text that does not parse as a +/// `timestamp` or `timestamptz`) +pub const INVALID_DATETIME_FORMAT: &str = "22007"; + +/// `22008` — `datetime_field_overflow` (an instant outside the range a +/// `timestamp` or `timestamptz` holds) +pub const DATETIME_FIELD_OVERFLOW: &str = "22008"; + /// `22012` — `division_by_zero` (`/` or `%` with a zero divisor — /// raised at runtime instead of evaluating to `NULL`) pub const DIVISION_BY_ZERO: &str = "22012"; @@ -397,6 +405,8 @@ mod tests { FEATURE_NOT_SUPPORTED, DATA_EXCEPTION, NUMERIC_VALUE_OUT_OF_RANGE, + INVALID_DATETIME_FORMAT, + DATETIME_FIELD_OVERFLOW, DIVISION_BY_ZERO, INVALID_LIMIT_VALUE, INVALID_TEXT_REPRESENTATION, diff --git a/nodedb/src/bridge/envelope/error_code.rs b/nodedb/src/bridge/envelope/error_code.rs index 24c1167f7..9a348a9c8 100644 --- a/nodedb/src/bridge/envelope/error_code.rs +++ b/nodedb/src/bridge/envelope/error_code.rs @@ -1,7 +1,7 @@ // SPDX-License-Identifier: BUSL-1.1 -//! Deterministic Data-Plane error codes, and their conversions to and from -//! the crate's typed error. +//! Deterministic Data-Plane error codes. `error_code_from` holds the +//! conversions into them. /// Deterministic error codes returned by the Data Plane. /// @@ -117,6 +117,13 @@ pub enum ErrorCode { /// SQLSTATE a planner-side `SqlError::UnknownColumn` produces — a missing /// column reports identically wherever it is detected. UndefinedColumn { column: String }, + /// A full-text search named a field that cannot serve it, detected when + /// the Data Plane resolves the field's index. + TextColumn { + collection: String, + column: String, + fault: nodedb_types::text_search::TextColumnFault, + }, /// Internal error (io_uring failure, corruption, etc.) Internal { detail: String }, /// Operation is not supported on this engine, or not yet implemented for @@ -127,7 +134,16 @@ pub enum ErrorCode { /// applied. The shard state is unknown — the client must treat this as a /// fatal error and the operator must restart the shard (WAL replay restores /// correct state on startup). Never silently continues. - RollbackFailed { entry_index: usize, detail: String }, + /// + /// `entry_index` is the forward-order position of the failed entry in + /// the undo log. `detail` names the reverse write that failed. `cause` + /// is the error that reverse write returned, or `None` when the engine + /// state did not match the entry and no reverse write ran. + RollbackFailed { + entry_index: usize, + detail: String, + cause: Option>, + }, /// The active Calvin executor detected that the declared predicate no /// longer matches the engine state at execution time (OLLP mismatch). /// No write was applied. The OLLP orchestrator retries with a fresh @@ -156,6 +172,10 @@ pub enum ErrorCode { /// different dimensions, an argument of the wrong type, a malformed /// JSONPath. Surfaces as SQLSTATE `22000` (`data_exception`). DataException { detail: String }, + /// A computed value lies outside the range of its result type, such as + /// an exact integer SUM past the decimal range. SQLSTATE `22003` + /// (`numeric_value_out_of_range`). + NumericValueOutOfRange { detail: String }, /// The bridge dispatcher refused the request at a capacity limit, so /// nothing was enqueued or applied. Transient: the same request succeeds /// once capacity frees. `reason` names the limit and its counts. @@ -179,293 +199,26 @@ pub enum ErrorCode { /// A DROP refused because other objects depend on its target. `object` /// names the target. SQLSTATE `2BP01` (dependent_objects_still_exist). DependentObjectsExist { object: String, detail: String }, -} - -/// An expression evaluation failure, as the Data Plane reports it. -/// -/// Exhaustive, so a new evaluator error picks its own code rather than -/// defaulting to one. -impl From for ErrorCode { - fn from(e: nodedb_query::EvalError) -> Self { - match e { - nodedb_query::EvalError::DivisionByZero => Self::DivisionByZero, - nodedb_query::EvalError::UnknownFunction { name } => Self::UndefinedFunction { name }, - e @ (nodedb_query::EvalError::VectorDimensionMismatch { .. } - | nodedb_query::EvalError::ArgumentType { .. } - | nodedb_query::EvalError::InvalidJsonPath { .. }) => Self::DataException { - detail: e.to_string(), - }, - } - } -} - -/// A KV write that binds a row to `Surrogate::ZERO` fails prevalidation, the -/// same refusal the dispatch-time check returns. -impl From for ErrorCode { - fn from(e: crate::engine::kv::UnboundKvWrite) -> Self { - Self::RejectedPrevalidation { - reason: e.to_string(), - } - } -} - -impl From for ErrorCode { - fn from(e: crate::Error) -> Self { - match e { - crate::Error::DeadlineExceeded { .. } => Self::DeadlineExceeded, - crate::Error::RejectedConstraint { - constraint, detail, .. - } => Self::RejectedConstraint { constraint, detail }, - crate::Error::RejectedPrevalidation { reason, .. } => { - Self::RejectedPrevalidation { reason } - } - crate::Error::RetryableRefusal { reason } => Self::RetryableRefusal { reason }, - crate::Error::CollectionNotFound { .. } - | crate::Error::CollectionDeactivated { .. } - | crate::Error::DocumentNotFound { .. } => Self::NotFound, - crate::Error::RejectedAuthz { resource, .. } => Self::RejectedAuthz { resource }, - // The Control Plane gives both `40001` (serialization_failure). - crate::Error::ConflictRetry { .. } | crate::Error::CalvinSerializationConflict => { - Self::ConflictRetry - } - crate::Error::MemoryExhausted { .. } => Self::ResourcesExhausted, - crate::Error::Backpressure { .. } => Self::ResourcesExhausted, - crate::Error::AppendOnlyViolation { collection, .. } => { - Self::AppendOnlyViolation { collection } - } - crate::Error::BalanceViolation { - collection, detail, .. - } => Self::BalanceViolation { collection, detail }, - // A materialized-sum target that cannot be addressed breaks the - // balance invariant the target collection maintains, so it crosses - // the bridge as the same class of violation the Control Plane - // already renders it as — not as a generic `Internal`, which would - // reach the client as SQLSTATE `XX000` and lose the target - // collection, join column, and join value the message names. - crate::Error::MaterializedSumTargetNotFound { - target_collection, - join_column, - join_value, - } => Self::BalanceViolation { - collection: target_collection, - detail: format!( - "no row with primary key '{join_value}', referenced by join column \ - '{join_column}'" - ), - }, - crate::Error::PeriodLocked { collection, .. } => Self::PeriodLocked { collection }, - crate::Error::PeriodLockMisconfigured { - collection, - ref_table, - status_column, - row_identity, - } => Self::PeriodLockMisconfigured { - collection, - ref_table, - status_column, - row_identity, - }, - crate::Error::RetentionViolation { collection, .. } => { - Self::RetentionViolation { collection } - } - crate::Error::LegalHoldActive { collection, .. } => { - Self::LegalHoldActive { collection } - } - crate::Error::StateTransitionViolation { - collection, detail, .. - } => Self::StateTransitionViolation { collection, detail }, - crate::Error::TransitionCheckViolation { collection, detail } => { - Self::TransitionCheckViolation { collection, detail } - } - crate::Error::TypeGuardViolation { - collection, detail, .. - } => Self::TypeGuardViolation { collection, detail }, - crate::Error::TypeMismatch { - collection, detail, .. - } => Self::TypeMismatch { collection, detail }, - crate::Error::InsufficientBalance { - collection, detail, .. - } => Self::InsufficientBalance { collection, detail }, - crate::Error::RateExceeded { - gate, - retry_after_ms, - .. - } => Self::RateExceeded { - gate, - retry_after_ms, - }, - // A capacity refusal enqueued nothing, and the same request - // succeeds once capacity frees. - capacity @ crate::Error::DispatchCapacity { .. } => Self::DispatchCapacity { - reason: capacity.to_string(), - }, - crate::Error::TxnOverlayMemoryExceeded { limit } => { - Self::TxnOverlayMemoryExceeded { limit } - } - crate::Error::DivisionByZero => Self::DivisionByZero, - crate::Error::UndefinedFunction { name } => Self::UndefinedFunction { name }, - crate::Error::DataException { detail } => Self::DataException { detail }, - // `42601` (syntax_error), as the Control Plane gives both. - crate::Error::BadRequest { detail } | crate::Error::PlanError { detail } => { - Self::BadRequest { detail } - } - // `0A000` (feature_not_supported), as the Control Plane gives both. - crate::Error::FeatureNotSupported { detail } => Self::Unsupported { detail }, - unsupported @ crate::Error::CrossCollectionNotColocated { .. } => Self::Unsupported { - detail: unsupported.to_string(), - }, - crate::Error::UndefinedColumn { column } => Self::UndefinedColumn { column }, - // Same condition an undefined column reports at plan time, raised - // here by the strict encoder for a transport the planner never - // sees (native client, `COPY FROM`, CRDT delta merge). - crate::Error::UnknownStrictField { column, .. } => Self::UndefinedColumn { column }, - // Already a Data-Plane verdict: hand back the same code rather - // than re-wrapping it as `Internal` and losing its SQLSTATE. - crate::Error::DataPlane(code) => code, - // Class `22`, the class the Control Plane gives both. - e @ (crate::Error::OffsetRegression { .. } - | crate::Error::BackupTenantMismatch { .. } - | crate::Error::InvalidLimitValue { .. }) => Self::DataException { - detail: e.to_string(), - }, - // `40000`, as the Control Plane gives it. - e @ crate::Error::CalvinParticipantError => Self::TransactionRollback { - detail: e.to_string(), - }, - // `40001`: the client retries the statement. - e @ crate::Error::RetryableSchemaChanged { .. } => Self::RetryableRefusal { - reason: e.to_string(), - }, - // `25001`, as the Control Plane gives all three. - e @ (crate::Error::CrdtApplyForbiddenInTransaction - | crate::Error::NotInTransactionBlock { .. } - | crate::Error::CrossShardInExplicitTransaction) => Self::ActiveSqlTransaction { - detail: e.to_string(), - }, - // `2BP01`, as the Control Plane gives both. The detail is the - // public message the Control Plane renders. - crate::Error::DependentObjectsExist { - root_kind, - root_name, - dependent_count, - dependents, - .. - } => { - let (object, detail) = crate::error_classify::dependent_objects_text( - root_kind, - &root_name, - dependent_count, - &dependents, - ); - Self::DependentObjectsExist { object, detail } - } - crate::Error::RoleInUse { role, dependents } => { - let object = format!("role \"{role}\""); - let detail = crate::Error::RoleInUse { role, dependents }.to_string(); - Self::DependentObjectsExist { object, detail } - } - crate::Error::CrdtAdmissionRetriesExhausted { .. } => Self::ConflictRetry, - // Retryable refusals whose class (`55P03`) no Data-Plane code has. - // The retry contract survives: nothing was applied. - e @ (crate::Error::NoLeader { .. } - | crate::Error::GroupQuorumUnavailable { .. } - | crate::Error::GroupMarksUnavailable { .. } - | crate::Error::BackupCaptureMoved { .. } - | crate::Error::AuthorizationStateBehind { .. } - | crate::Error::LinearizableReadRefused { .. } - | crate::Error::StaleReadNotLeader { .. }) => Self::RetryableRefusal { - reason: e.to_string(), - }, - // Class `57`: the client retries once the leader settles. - e @ crate::Error::NotLeader { .. } => Self::DispatchCapacity { - reason: e.to_string(), - }, - crate::Error::CrdtAdmissionTimeout { .. } => Self::DeadlineExceeded, - e @ crate::Error::VShardAdmissionCapacityExceeded { .. } => Self::RateExceeded { - gate: e.to_string(), - retry_after_ms: 0, - }, - // Class `53`: a configured resource ceiling. - crate::Error::QuotaOvercommit { .. } - | crate::Error::TenantVectorDimExceeded { .. } - | crate::Error::TenantGraphDepthExceeded { .. } => Self::ResourcesExhausted, - // Class `28` has no Data-Plane code. The nearest is the access - // refusal, which keeps it a client error the client cannot retry. - e @ (crate::Error::BackupKeyMismatch | crate::Error::SessionTokenExpired) => { - Self::RejectedAuthz { - resource: e.to_string(), - } - } - // Client errors of class `42`, and client errors whose class - // (`25006`, `55`) no Data-Plane code has. `BadRequest` is the - // class their public code has. - e @ (crate::Error::CrdtAdmissionInvalidPlan { .. } - | crate::Error::CrdtAdmissionCallerFence - | crate::Error::CrdtApplyRequiresAdmission - | crate::Error::CloneWriteRequiresMaterialize { .. } - | crate::Error::ObjectNotInPrerequisiteState { .. } - | crate::Error::MirrorReadOnly { .. } - | crate::Error::UndefinedObject { .. } - | crate::Error::AmbiguousColumn { .. } - | crate::Error::ExecutionLimitExceeded { .. } - | crate::Error::LimitExceeded { .. } - | crate::Error::Promql(_) - | crate::Error::SequencerUnavailable - | crate::Error::SessionCapExceeded { .. } - | crate::Error::SessionIdleTimeout - | crate::Error::SessionKilledByAdmin - | crate::Error::SessionUserDropped - | crate::Error::OidcProviderTenantUnbound - | crate::Error::OidcProviderTenantUnavailable { .. } - | crate::Error::ExternalRoleUndefined { .. } - | crate::Error::OidcNoDefaultDatabase { .. } - | crate::Error::RoleInheritanceCycle { .. } - | crate::Error::RoleInheritanceDepthExceeded { .. }) => Self::BadRequest { - detail: e.to_string(), - }, - // Retry exhaustion takes the code of its cause. - crate::Error::OllpExhausted { cause, .. } => match cause { - crate::OllpExhaustedCause::PredicateDrift => Self::ConflictRetry, - crate::OllpExhaustedCause::PreAdmission(inner) => Self::from(*inner), - crate::OllpExhaustedCause::AdmissionRefused { detail } => Self::RateExceeded { - gate: detail, - retry_after_ms: 0, - }, - }, - // Server-side faults and system defects. `Shaping`, - // `RemoteTyped` and `Ddl` carry a public numeric code that has no - // Data-Plane twin, and none is raised on the Data Plane. - e @ (crate::Error::MaterializedSumResolutionMissing { .. } - | crate::Error::RetryableLeaderChange { .. } - | crate::Error::CommittedResultUnavailable { .. } - | crate::Error::ProposalOutcomeUnknown { .. } - | crate::Error::MetadataLeaderUnavailable - | crate::Error::Wal(_) - | crate::Error::Dispatch { .. } - | crate::Error::Storage { .. } - | crate::Error::ColdStorage { .. } - | crate::Error::Serialization { .. } - | crate::Error::Codec { .. } - | crate::Error::SegmentCorrupted { .. } - | crate::Error::Crdt(_) - | crate::Error::Io(_) - | crate::Error::Config { .. } - | crate::Error::Encryption { .. } - | crate::Error::Bridge { .. } - | crate::Error::VersionCompat { .. } - | crate::Error::RestoreTargetNotEmpty { .. } - | crate::Error::RestoreVerificationFailed { .. } - | crate::Error::Internal { .. } - | crate::Error::Shaping(_) - | crate::Error::RemoteTyped { .. } - | crate::Error::Ddl(_) - | crate::Error::DescriptorVersionAnomaly { .. } - | crate::Error::CollectionPurgeRowMissing { .. } - | crate::Error::CollectionUnstamped { .. } - | crate::Error::CatalogIntegrityViolation { .. } - | crate::Error::CascadeCycle { .. }) => Self::Internal { - detail: e.to_string(), - }, - } - } + /// A label write needs a node label past the partition's node-label cap. + /// Nothing of the statement applied. `node` and `label` name the refused + /// write, and `limit` is the cap. SQLSTATE `54000` (program_limit_exceeded). + NodeLabelLimit { + node: String, + label: String, + limit: usize, + }, + /// Text does not parse as the type of the column it is written to. + /// `detail` names the column, the text and the type. SQLSTATE `22P02` + /// (invalid_text_representation). + InvalidTextRepresentation { detail: String }, + /// A value of the wrong kind for the column it is written to. `detail` + /// names the column, the value and the type. SQLSTATE `42804` + /// (datatype_mismatch). + DatatypeMismatch { detail: String }, + /// Text does not parse as the `TIMESTAMP` or `TIMESTAMPTZ` it is written + /// to. SQLSTATE `22007` (invalid_datetime_format). + InvalidDatetimeFormat { detail: String }, + /// An instant outside the range a `TIMESTAMP` or `TIMESTAMPTZ` holds. + /// SQLSTATE `22008` (datetime_field_overflow). + DatetimeFieldOverflow { detail: String }, } diff --git a/nodedb/src/bridge/envelope/error_code_from.rs b/nodedb/src/bridge/envelope/error_code_from.rs new file mode 100644 index 000000000..0f3dd6a2e --- /dev/null +++ b/nodedb/src/bridge/envelope/error_code_from.rs @@ -0,0 +1,322 @@ +// SPDX-License-Identifier: BUSL-1.1 + +//! Conversions into the Data-Plane [`ErrorCode`]. + +use super::error_code::ErrorCode; + +/// An expression evaluation failure, as the Data Plane reports it. +/// +/// Exhaustive, so a new evaluator error picks its own code rather than +/// defaulting to one. +impl From for ErrorCode { + fn from(e: nodedb_query::EvalError) -> Self { + match e { + nodedb_query::EvalError::DivisionByZero => Self::DivisionByZero, + nodedb_query::EvalError::UnknownFunction { name } => Self::UndefinedFunction { name }, + e @ (nodedb_query::EvalError::VectorDimensionMismatch { .. } + | nodedb_query::EvalError::ArgumentType { .. } + | nodedb_query::EvalError::InvalidJsonPath { .. }) => Self::DataException { + detail: e.to_string(), + }, + e @ nodedb_query::EvalError::NumericOverflow { .. } => Self::NumericValueOutOfRange { + detail: e.to_string(), + }, + } + } +} + +/// A KV write that binds a row to `Surrogate::ZERO` fails prevalidation, the +/// same refusal the dispatch-time check returns. +impl From for ErrorCode { + fn from(e: crate::engine::kv::UnboundKvWrite) -> Self { + Self::RejectedPrevalidation { + reason: e.to_string(), + } + } +} + +impl From for ErrorCode { + fn from(e: crate::Error) -> Self { + match e { + crate::Error::DeadlineExceeded { .. } => Self::DeadlineExceeded, + crate::Error::RejectedConstraint { + constraint, detail, .. + } => Self::RejectedConstraint { constraint, detail }, + crate::Error::RejectedPrevalidation { reason, .. } => { + Self::RejectedPrevalidation { reason } + } + crate::Error::RetryableRefusal { reason } => Self::RetryableRefusal { reason }, + crate::Error::CollectionNotFound { .. } + | crate::Error::CollectionDeactivated { .. } + | crate::Error::DocumentNotFound { .. } => Self::NotFound, + crate::Error::RejectedAuthz { resource, .. } => Self::RejectedAuthz { resource }, + // The Control Plane gives both `40001` (serialization_failure). + crate::Error::ConflictRetry { .. } | crate::Error::CalvinSerializationConflict => { + Self::ConflictRetry + } + crate::Error::MemoryExhausted { .. } => Self::ResourcesExhausted, + crate::Error::Backpressure { .. } => Self::ResourcesExhausted, + crate::Error::AppendOnlyViolation { collection, .. } => { + Self::AppendOnlyViolation { collection } + } + crate::Error::BalanceViolation { + collection, detail, .. + } => Self::BalanceViolation { collection, detail }, + // A materialized-sum target that cannot be addressed breaks the + // balance invariant the target collection maintains, so it crosses + // the bridge as the same class of violation the Control Plane + // already renders it as — not as a generic `Internal`, which would + // reach the client as SQLSTATE `XX000` and lose the target + // collection, join column, and join value the message names. + crate::Error::MaterializedSumTargetNotFound { + target_collection, + join_column, + join_value, + } => Self::BalanceViolation { + collection: target_collection, + detail: format!( + "no row with primary key '{join_value}', referenced by join column \ + '{join_column}'" + ), + }, + crate::Error::PeriodLocked { collection, .. } => Self::PeriodLocked { collection }, + crate::Error::PeriodLockMisconfigured { + collection, + ref_table, + status_column, + row_identity, + } => Self::PeriodLockMisconfigured { + collection, + ref_table, + status_column, + row_identity, + }, + crate::Error::RetentionViolation { collection, .. } => { + Self::RetentionViolation { collection } + } + crate::Error::LegalHoldActive { collection, .. } => { + Self::LegalHoldActive { collection } + } + crate::Error::StateTransitionViolation { + collection, detail, .. + } => Self::StateTransitionViolation { collection, detail }, + crate::Error::TransitionCheckViolation { collection, detail } => { + Self::TransitionCheckViolation { collection, detail } + } + crate::Error::TypeGuardViolation { + collection, detail, .. + } => Self::TypeGuardViolation { collection, detail }, + crate::Error::TypeMismatch { + collection, detail, .. + } => Self::TypeMismatch { collection, detail }, + crate::Error::InsufficientBalance { + collection, detail, .. + } => Self::InsufficientBalance { collection, detail }, + crate::Error::RateExceeded { + gate, + retry_after_ms, + .. + } => Self::RateExceeded { + gate, + retry_after_ms, + }, + // A capacity refusal enqueued nothing, and the same request + // succeeds once capacity frees. + capacity @ crate::Error::DispatchCapacity { .. } => Self::DispatchCapacity { + reason: capacity.to_string(), + }, + crate::Error::TxnOverlayMemoryExceeded { limit } => { + Self::TxnOverlayMemoryExceeded { limit } + } + crate::Error::DivisionByZero => Self::DivisionByZero, + crate::Error::UndefinedFunction { name } => Self::UndefinedFunction { name }, + crate::Error::DataException { detail } => Self::DataException { detail }, + // `42601` (syntax_error), as the Control Plane gives both. + crate::Error::BadRequest { detail } | crate::Error::PlanError { detail } => { + Self::BadRequest { detail } + } + // `0A000` (feature_not_supported), as the Control Plane gives both. + crate::Error::FeatureNotSupported { detail } => Self::Unsupported { detail }, + unsupported @ crate::Error::CrossCollectionNotColocated { .. } => Self::Unsupported { + detail: unsupported.to_string(), + }, + crate::Error::UndefinedColumn { column } => Self::UndefinedColumn { column }, + crate::Error::TextColumn { + collection, + column, + fault, + } => Self::TextColumn { + collection, + column, + fault, + }, + // Same condition an undefined column reports at plan time, raised + // here by the strict encoder for a transport the planner never + // sees (native client, `COPY FROM`, CRDT delta merge). + crate::Error::UnknownStrictField { column, .. } => Self::UndefinedColumn { column }, + // Already a Data-Plane verdict: hand back the same code rather + // than re-wrapping it as `Internal` and losing its SQLSTATE. + crate::Error::DataPlane(code) => code, + // Class `22`, the class the Control Plane gives both. + e @ (crate::Error::OffsetRegression { .. } + | crate::Error::BackupTenantMismatch { .. } + | crate::Error::InvalidLimitValue { .. }) => Self::DataException { + detail: e.to_string(), + }, + // `22003`, as the Control Plane gives it. + crate::Error::NumericValueOutOfRange { detail } => { + Self::NumericValueOutOfRange { detail } + } + // `22P02`, `42804`, `22007` and `22008`, as the Control Plane + // gives them. + crate::Error::InvalidTextRepresentation { detail } => { + Self::InvalidTextRepresentation { detail } + } + crate::Error::DatatypeMismatch { detail } => Self::DatatypeMismatch { detail }, + crate::Error::InvalidDatetimeFormat { detail } => { + Self::InvalidDatetimeFormat { detail } + } + crate::Error::DatetimeFieldOverflow { detail } => { + Self::DatetimeFieldOverflow { detail } + } + // `40000`, as the Control Plane gives it. + e @ crate::Error::CalvinParticipantError => Self::TransactionRollback { + detail: e.to_string(), + }, + // `40001`: the client retries the statement. + e @ crate::Error::RetryableSchemaChanged { .. } => Self::RetryableRefusal { + reason: e.to_string(), + }, + // `25001`, as the Control Plane gives all three. + e @ (crate::Error::CrdtApplyForbiddenInTransaction + | crate::Error::NotInTransactionBlock { .. } + | crate::Error::CrossShardInExplicitTransaction) => Self::ActiveSqlTransaction { + detail: e.to_string(), + }, + // `2BP01`, as the Control Plane gives both. The detail is the + // public message the Control Plane renders. + crate::Error::DependentObjectsExist { + root_kind, + root_name, + dependent_count, + dependents, + .. + } => { + let (object, detail) = crate::error_classify::dependent_objects_text( + root_kind, + &root_name, + dependent_count, + &dependents, + ); + Self::DependentObjectsExist { object, detail } + } + crate::Error::RoleInUse { role, dependents } => { + let object = format!("role \"{role}\""); + let detail = crate::Error::RoleInUse { role, dependents }.to_string(); + Self::DependentObjectsExist { object, detail } + } + crate::Error::CrdtAdmissionRetriesExhausted { .. } => Self::ConflictRetry, + // Retryable refusals whose class (`55P03`) no Data-Plane code has. + // The retry contract survives: nothing was applied. + e @ (crate::Error::NoLeader { .. } + | crate::Error::GroupQuorumUnavailable { .. } + | crate::Error::GroupMarksUnavailable { .. } + | crate::Error::BackupCaptureMoved { .. } + | crate::Error::AuthorizationStateBehind { .. } + | crate::Error::LinearizableReadRefused { .. } + | crate::Error::StaleReadNotLeader { .. }) => Self::RetryableRefusal { + reason: e.to_string(), + }, + // Class `57`: the client retries once the leader settles. + e @ crate::Error::NotLeader { .. } => Self::DispatchCapacity { + reason: e.to_string(), + }, + crate::Error::CrdtAdmissionTimeout { .. } => Self::DeadlineExceeded, + e @ crate::Error::VShardAdmissionCapacityExceeded { .. } => Self::RateExceeded { + gate: e.to_string(), + retry_after_ms: 0, + }, + // Class `53`: a configured resource ceiling. + crate::Error::QuotaOvercommit { .. } + | crate::Error::TenantVectorDimExceeded { .. } + | crate::Error::TenantGraphDepthExceeded { .. } => Self::ResourcesExhausted, + // Class `28` has no Data-Plane code. The nearest is the access + // refusal, which keeps it a client error the client cannot retry. + e @ (crate::Error::BackupKeyMismatch | crate::Error::SessionTokenExpired) => { + Self::RejectedAuthz { + resource: e.to_string(), + } + } + // Client errors of class `42`, and client errors whose class + // (`25006`, `55`) no Data-Plane code has. `BadRequest` is the + // class their public code has. + e @ (crate::Error::CrdtAdmissionInvalidPlan { .. } + | crate::Error::CrdtAdmissionCallerFence + | crate::Error::CrdtApplyRequiresAdmission + | crate::Error::CloneWriteRequiresMaterialize { .. } + | crate::Error::ObjectNotInPrerequisiteState { .. } + | crate::Error::MirrorReadOnly { .. } + | crate::Error::UndefinedObject { .. } + | crate::Error::AmbiguousColumn { .. } + | crate::Error::ExecutionLimitExceeded { .. } + | crate::Error::LimitExceeded { .. } + | crate::Error::Promql(_) + | crate::Error::SequencerUnavailable + | crate::Error::SessionCapExceeded { .. } + | crate::Error::SessionIdleTimeout + | crate::Error::SessionKilledByAdmin + | crate::Error::SessionUserDropped + | crate::Error::OidcProviderTenantUnbound + | crate::Error::OidcProviderTenantUnavailable { .. } + | crate::Error::ExternalRoleUndefined { .. } + | crate::Error::OidcNoDefaultDatabase { .. } + | crate::Error::RoleInheritanceCycle { .. } + | crate::Error::RoleInheritanceDepthExceeded { .. }) => Self::BadRequest { + detail: e.to_string(), + }, + // Retry exhaustion takes the code of its cause. + crate::Error::OllpExhausted { cause, .. } => match cause { + crate::OllpExhaustedCause::PredicateDrift => Self::ConflictRetry, + crate::OllpExhaustedCause::PreAdmission(inner) => Self::from(*inner), + crate::OllpExhaustedCause::AdmissionRefused { detail } => Self::RateExceeded { + gate: detail, + retry_after_ms: 0, + }, + }, + // Server-side faults and system defects. `Shaping`, + // `RemoteTyped` and `Ddl` carry a public numeric code that has no + // Data-Plane twin, and none is raised on the Data Plane. + e @ (crate::Error::MaterializedSumResolutionMissing { .. } + | crate::Error::RetryableLeaderChange { .. } + | crate::Error::CommittedResultUnavailable { .. } + | crate::Error::ProposalOutcomeUnknown { .. } + | crate::Error::MetadataLeaderUnavailable + | crate::Error::Wal(_) + | crate::Error::Dispatch { .. } + | crate::Error::Storage { .. } + | crate::Error::ColdStorage { .. } + | crate::Error::Serialization { .. } + | crate::Error::Codec { .. } + | crate::Error::SegmentCorrupted { .. } + | crate::Error::Crdt(_) + | crate::Error::Io(_) + | crate::Error::Config { .. } + | crate::Error::Encryption { .. } + | crate::Error::Bridge { .. } + | crate::Error::VersionCompat { .. } + | crate::Error::RestoreTargetNotEmpty { .. } + | crate::Error::RestoreVerificationFailed { .. } + | crate::Error::Internal { .. } + | crate::Error::Shaping(_) + | crate::Error::RemoteTyped { .. } + | crate::Error::Ddl(_) + | crate::Error::DescriptorVersionAnomaly { .. } + | crate::Error::CollectionPurgeRowMissing { .. } + | crate::Error::CollectionUnstamped { .. } + | crate::Error::CatalogIntegrityViolation { .. } + | crate::Error::CascadeCycle { .. }) => Self::Internal { + detail: e.to_string(), + }, + } + } +} diff --git a/nodedb/src/bridge/envelope/mod.rs b/nodedb/src/bridge/envelope/mod.rs index 91adf9576..ca2f6ea70 100644 --- a/nodedb/src/bridge/envelope/mod.rs +++ b/nodedb/src/bridge/envelope/mod.rs @@ -3,6 +3,7 @@ //! Request/response envelopes exchanged over the SPSC bridge. pub mod error_code; +pub mod error_code_from; pub mod payload; pub mod request; pub mod response; diff --git a/nodedb/src/control/cluster/array_executor/refusal.rs b/nodedb/src/control/cluster/array_executor/refusal.rs index af94eae15..fc3f23ad5 100644 --- a/nodedb/src/control/cluster/array_executor/refusal.rs +++ b/nodedb/src/control/cluster/array_executor/refusal.rs @@ -104,10 +104,16 @@ pub(super) fn execution_error(context: &str, error: crate::Error) -> ClusterErro | crate::Error::UndefinedObject { .. } | crate::Error::ObjectNotInPrerequisiteState { .. } | crate::Error::UndefinedColumn { .. } + | crate::Error::TextColumn { .. } | crate::Error::AmbiguousColumn { .. } | crate::Error::UnknownStrictField { .. } | crate::Error::DivisionByZero | crate::Error::DataException { .. } + | crate::Error::NumericValueOutOfRange { .. } + | crate::Error::InvalidTextRepresentation { .. } + | crate::Error::DatatypeMismatch { .. } + | crate::Error::InvalidDatetimeFormat { .. } + | crate::Error::DatetimeFieldOverflow { .. } | crate::Error::InvalidLimitValue { .. } | crate::Error::RetryableSchemaChanged { .. } | crate::Error::RetryableLeaderChange { .. } diff --git a/nodedb/src/control/cluster/data_plane_error_wire.rs b/nodedb/src/control/cluster/data_plane_error_wire.rs index d39f47279..1fdaf2a6c 100644 --- a/nodedb/src/control/cluster/data_plane_error_wire.rs +++ b/nodedb/src/control/cluster/data_plane_error_wire.rs @@ -1,185 +1,22 @@ // SPDX-License-Identifier: BUSL-1.1 //! Lossless conversion between the Data-Plane [`ErrorCode`] and its cluster -//! wire mirror [`DataPlaneErrorCode`], plus the one mapping every cross-node -//! executor uses to answer with a local execution error. +//! wire mirror [`DataPlaneErrorCode`]. The mapping every cross-node executor +//! uses to answer with a local execution error lives in +//! [`super::execution_error_wire`] and is re-exported here. //! //! Both matches are exhaustive with no catch-all, so a new `ErrorCode` variant //! fails to compile here until it is mirrored on the wire instead of silently //! degrading to `Internal` and losing its SQLSTATE at the coordinator. -use nodedb_cluster::rpc_codec::{ - DataPlaneCounterFault, DataPlaneErrorCode, DataPlaneSyncHold, TypedClusterError, -}; - -use crate::bridge::envelope::{CounterFault, ErrorCode, SyncHold}; - -/// Map a local-execution [`crate::Error`] to the wire error a remote caller -/// receives. -/// -/// A Data-Plane verdict crosses verbatim as `TypedClusterError::DataPlane`, so -/// the coordinator rebuilds `Error::DataPlane(code)` and renders the SQLSTATE -/// single-node execution renders. Every other error keeps its own numeric -/// classification from `NodeDbError::from(err).code()` — never a hardcoded -/// plan-decode code, which will misname what failed. -pub(crate) fn execution_error_to_typed(err: crate::Error) -> TypedClusterError { - match err { - crate::Error::DataPlane(code) => TypedClusterError::DataPlane { code: code.into() }, - // A statement that ran out of time keeps the wire's own deadline - // variant, which the coordinator rebuilds as `Error::DeadlineExceeded`. - // Folding it into `Internal` will report a client's own timeout as an - // internal failure once it crossed a node boundary. - crate::Error::DeadlineExceeded { .. } => { - TypedClusterError::DeadlineExceeded { elapsed_ms: 0 } - } - // A redirect crosses as the wire's own redirect, with the leader and - // the term this node knows it at, so the coordinator moves its - // routing hint and retries against that leader. - not_leader @ crate::Error::NotLeader { .. } => TypedClusterError::from(not_leader), - // A Calvin abort keeps its verdict and a schema change stays - // retryable, so a routed submit answers as a local one does. - typed @ (crate::Error::CalvinSerializationConflict - | crate::Error::CalvinParticipantError - | crate::Error::RetryableSchemaChanged { .. }) => TypedClusterError::from(typed), - // A Control-Plane constraint refusal crosses verbatim, same as a - // Data-Plane verdict, so the coordinator answers 23502 vs 23505 - // instead of flattening both into one numeric class. - crate::Error::RejectedConstraint { - collection, - constraint, - detail, - } => TypedClusterError::RejectedConstraint { - collection, - constraint, - detail, - }, - // A capacity refusal crosses as its own verdict, so the coordinator - // answers the retryable overload class. - capacity @ crate::Error::DispatchCapacity { .. } => TypedClusterError::DataPlane { - code: DataPlaneErrorCode::DispatchCapacity { - reason: capacity.to_string(), - }, - }, - // Every other error crosses as its public numeric code and message. - // The coordinator renders the SQLSTATE that code maps to. - other @ (crate::Error::TxnOverlayMemoryExceeded { .. } - | crate::Error::RejectedAuthz { .. } - | crate::Error::OffsetRegression { .. } - | crate::Error::ConflictRetry { .. } - | crate::Error::RejectedPrevalidation { .. } - | crate::Error::RetryableRefusal { .. } - | crate::Error::AppendOnlyViolation { .. } - | crate::Error::BalanceViolation { .. } - | crate::Error::MaterializedSumTargetNotFound { .. } - | crate::Error::MaterializedSumResolutionMissing { .. } - | crate::Error::PeriodLocked { .. } - | crate::Error::PeriodLockMisconfigured { .. } - | crate::Error::RetentionViolation { .. } - | crate::Error::LegalHoldActive { .. } - | crate::Error::StateTransitionViolation { .. } - | crate::Error::TransitionCheckViolation { .. } - | crate::Error::TypeGuardViolation { .. } - | crate::Error::TypeMismatch { .. } - | crate::Error::InsufficientBalance { .. } - | crate::Error::RateExceeded { .. } - | crate::Error::CollectionNotFound { .. } - | crate::Error::DocumentNotFound { .. } - | crate::Error::CollectionDeactivated { .. } - | crate::Error::VShardAdmissionCapacityExceeded { .. } - | crate::Error::CrdtAdmissionRetriesExhausted { .. } - | crate::Error::CrdtAdmissionInvalidPlan { .. } - | crate::Error::CrdtAdmissionCallerFence - | crate::Error::CrdtApplyRequiresAdmission - | crate::Error::CrdtApplyForbiddenInTransaction - | crate::Error::NotInTransactionBlock { .. } - | crate::Error::CrdtAdmissionTimeout { .. } - | crate::Error::NoLeader { .. } - | crate::Error::CrossCollectionNotColocated { .. } - | crate::Error::CloneWriteRequiresMaterialize { .. } - | crate::Error::BadRequest { .. } - | crate::Error::BackupTenantMismatch { .. } - | crate::Error::BackupKeyMismatch - | crate::Error::QuotaOvercommit { .. } - | crate::Error::PlanError { .. } - | crate::Error::FeatureNotSupported { .. } - | crate::Error::UndefinedFunction { .. } - | crate::Error::UndefinedObject { .. } - | crate::Error::ObjectNotInPrerequisiteState { .. } - | crate::Error::UndefinedColumn { .. } - | crate::Error::AmbiguousColumn { .. } - | crate::Error::UnknownStrictField { .. } - | crate::Error::DivisionByZero - | crate::Error::DataException { .. } - | crate::Error::InvalidLimitValue { .. } - | crate::Error::RetryableLeaderChange { .. } - | crate::Error::CommittedResultUnavailable { .. } - | crate::Error::ProposalOutcomeUnknown { .. } - | crate::Error::GroupQuorumUnavailable { .. } - | crate::Error::GroupMarksUnavailable { .. } - | crate::Error::BackupCaptureMoved { .. } - | crate::Error::MetadataLeaderUnavailable - | crate::Error::AuthorizationStateBehind { .. } - | crate::Error::LinearizableReadRefused { .. } - | crate::Error::ExecutionLimitExceeded { .. } - | crate::Error::LimitExceeded { .. } - | crate::Error::Wal(_) - | crate::Error::Dispatch { .. } - | crate::Error::Storage { .. } - | crate::Error::ColdStorage { .. } - | crate::Error::Serialization { .. } - | crate::Error::Codec { .. } - | crate::Error::SegmentCorrupted { .. } - | crate::Error::MemoryExhausted { .. } - | crate::Error::Backpressure { .. } - | crate::Error::Crdt(_) - | crate::Error::Io(_) - | crate::Error::Config { .. } - | crate::Error::Encryption { .. } - | crate::Error::Bridge { .. } - | crate::Error::VersionCompat { .. } - | crate::Error::RestoreTargetNotEmpty { .. } - | crate::Error::RestoreVerificationFailed { .. } - | crate::Error::Internal { .. } - | crate::Error::Shaping(_) - | crate::Error::RemoteTyped { .. } - | crate::Error::Ddl(_) - | crate::Error::DescriptorVersionAnomaly { .. } - | crate::Error::CollectionPurgeRowMissing { .. } - | crate::Error::CollectionUnstamped { .. } - | crate::Error::CatalogIntegrityViolation { .. } - | crate::Error::Promql(_) - | crate::Error::DependentObjectsExist { .. } - | crate::Error::RoleInUse { .. } - | crate::Error::CascadeCycle { .. } - | crate::Error::CrossShardInExplicitTransaction - | crate::Error::SequencerUnavailable - | crate::Error::SessionCapExceeded { .. } - | crate::Error::SessionIdleTimeout - | crate::Error::SessionTokenExpired - | crate::Error::SessionKilledByAdmin - | crate::Error::SessionUserDropped - | crate::Error::OidcProviderTenantUnbound - | crate::Error::OidcProviderTenantUnavailable { .. } - | crate::Error::ExternalRoleUndefined { .. } - | crate::Error::OidcNoDefaultDatabase { .. } - | crate::Error::TenantVectorDimExceeded { .. } - | crate::Error::TenantGraphDepthExceeded { .. } - | crate::Error::RoleInheritanceCycle { .. } - | crate::Error::RoleInheritanceDepthExceeded { .. } - | crate::Error::OllpExhausted { .. } - | crate::Error::MirrorReadOnly { .. } - | crate::Error::StaleReadNotLeader { .. }) => numeric_typed(other), - } -} +use nodedb_cluster::rpc_codec::DataPlaneErrorCode; -/// The wire error for a local error with no typed wire carrier: its public -/// numeric code from `NodeDbError::from(err).code()`, and its message. The -/// coordinator rebuilds it as `Error::RemoteTyped`. -pub(crate) fn numeric_typed(err: crate::Error) -> TypedClusterError { - let message = err.to_string(); - let code = u32::from(nodedb_types::error::NodeDbError::from(err).code().0); - TypedClusterError::Internal { code, message } -} +use super::data_plane_fault_wire::{ + counter_fault_from_wire, counter_fault_to_wire, sync_hold_from_wire, sync_hold_to_wire, + text_column_fault_from_wire, text_column_fault_to_wire, +}; +pub(crate) use super::execution_error_wire::{execution_error_to_typed, numeric_typed}; +use crate::bridge::envelope::ErrorCode; /// Widen a pointer-width count to the wire's fixed `u64`. fn to_wire_count(value: usize) -> u64 { @@ -288,9 +125,11 @@ impl From for DataPlaneErrorCode { ErrorCode::RollbackFailed { entry_index, detail, + cause, } => Self::RollbackFailed { entry_index: to_wire_count(entry_index), detail, + cause: cause.map(|cause| Box::new(Self::from(*cause))), }, ErrorCode::OllpRetryRequired => Self::OllpRetryRequired, ErrorCode::TxnOverlayMemoryExceeded { limit } => Self::TxnOverlayMemoryExceeded { @@ -299,6 +138,7 @@ impl From for DataPlaneErrorCode { ErrorCode::DivisionByZero => Self::DivisionByZero, ErrorCode::UndefinedFunction { name } => Self::UndefinedFunction { name }, ErrorCode::DataException { detail } => Self::DataException { detail }, + ErrorCode::NumericValueOutOfRange { detail } => Self::NumericValueOutOfRange { detail }, ErrorCode::DispatchCapacity { reason } => Self::DispatchCapacity { reason }, ErrorCode::ExpiredBeforeExecution => Self::ExpiredBeforeExecution, ErrorCode::BadRequest { detail } => Self::BadRequest { detail }, @@ -307,6 +147,26 @@ impl From for DataPlaneErrorCode { ErrorCode::DependentObjectsExist { object, detail } => { Self::DependentObjectsExist { object, detail } } + ErrorCode::NodeLabelLimit { node, label, limit } => Self::NodeLabelLimit { + node, + label, + limit: to_wire_count(limit), + }, + ErrorCode::InvalidTextRepresentation { detail } => { + Self::InvalidTextRepresentation { detail } + } + ErrorCode::DatatypeMismatch { detail } => Self::DatatypeMismatch { detail }, + ErrorCode::InvalidDatetimeFormat { detail } => Self::InvalidDatetimeFormat { detail }, + ErrorCode::DatetimeFieldOverflow { detail } => Self::DatetimeFieldOverflow { detail }, + ErrorCode::TextColumn { + collection, + column, + fault, + } => Self::TextColumn { + collection, + column, + fault: text_column_fault_to_wire(fault), + }, } } } @@ -420,9 +280,11 @@ impl From for ErrorCode { DataPlaneErrorCode::RollbackFailed { entry_index, detail, + cause, } => Self::RollbackFailed { entry_index: from_wire_count(entry_index), detail, + cause: cause.map(|cause| Box::new(Self::from(*cause))), }, DataPlaneErrorCode::OllpRetryRequired => Self::OllpRetryRequired, DataPlaneErrorCode::TxnOverlayMemoryExceeded { limit } => { @@ -433,6 +295,9 @@ impl From for ErrorCode { DataPlaneErrorCode::DivisionByZero => Self::DivisionByZero, DataPlaneErrorCode::UndefinedFunction { name } => Self::UndefinedFunction { name }, DataPlaneErrorCode::DataException { detail } => Self::DataException { detail }, + DataPlaneErrorCode::NumericValueOutOfRange { detail } => { + Self::NumericValueOutOfRange { detail } + } DataPlaneErrorCode::DispatchCapacity { reason } => Self::DispatchCapacity { reason }, DataPlaneErrorCode::ExpiredBeforeExecution => Self::ExpiredBeforeExecution, DataPlaneErrorCode::BadRequest { detail } => Self::BadRequest { detail }, @@ -445,52 +310,60 @@ impl From for ErrorCode { DataPlaneErrorCode::DependentObjectsExist { object, detail } => { Self::DependentObjectsExist { object, detail } } + DataPlaneErrorCode::NodeLabelLimit { node, label, limit } => Self::NodeLabelLimit { + node, + label, + limit: from_wire_count(limit), + }, + DataPlaneErrorCode::InvalidTextRepresentation { detail } => { + Self::InvalidTextRepresentation { detail } + } + DataPlaneErrorCode::DatatypeMismatch { detail } => Self::DatatypeMismatch { detail }, + DataPlaneErrorCode::InvalidDatetimeFormat { detail } => { + Self::InvalidDatetimeFormat { detail } + } + DataPlaneErrorCode::DatetimeFieldOverflow { detail } => { + Self::DatetimeFieldOverflow { detail } + } + DataPlaneErrorCode::TextColumn { + collection, + column, + fault, + } => Self::TextColumn { + collection, + column, + fault: text_column_fault_from_wire(fault), + }, } } } -/// The wire form of a sync hold. -fn sync_hold_to_wire(hold: SyncHold) -> DataPlaneSyncHold { - match hold { - SyncHold::Duplicate => DataPlaneSyncHold::Duplicate, - SyncHold::Fenced => DataPlaneSyncHold::Fenced, - SyncHold::Gap { expected } => DataPlaneSyncHold::Gap { expected }, - } -} - -/// The sync hold a wire form names. -fn sync_hold_from_wire(hold: DataPlaneSyncHold) -> SyncHold { - match hold { - DataPlaneSyncHold::Duplicate => SyncHold::Duplicate, - DataPlaneSyncHold::Fenced => SyncHold::Fenced, - DataPlaneSyncHold::Gap { expected } => SyncHold::Gap { expected }, - } -} - -/// The wire form of a counter fault. Both types live in other crates, so the -/// mapping is a function, not a `From` impl. -fn counter_fault_to_wire(fault: CounterFault) -> DataPlaneCounterFault { - match fault { - CounterFault::NotAnInteger => DataPlaneCounterFault::NotAnInteger, - CounterFault::NotAFloat => DataPlaneCounterFault::NotAFloat, - CounterFault::IntegerOverflow => DataPlaneCounterFault::IntegerOverflow, - CounterFault::NonFinite => DataPlaneCounterFault::NonFinite, - } -} - -/// The counter fault a wire form carries. -fn counter_fault_from_wire(fault: DataPlaneCounterFault) -> CounterFault { - match fault { - DataPlaneCounterFault::NotAnInteger => CounterFault::NotAnInteger, - DataPlaneCounterFault::NotAFloat => CounterFault::NotAFloat, - DataPlaneCounterFault::IntegerOverflow => CounterFault::IntegerOverflow, - DataPlaneCounterFault::NonFinite => CounterFault::NonFinite, - } -} - #[cfg(test)] mod tests { use super::*; + use crate::bridge::envelope::{CounterFault, SyncHold}; + use nodedb_cluster::rpc_codec::TypedClusterError; + use nodedb_types::text_search::TextColumnFault; + + #[test] + fn text_column_code_roundtrips_verbatim() { + for fault in [ + TextColumnFault::Undeclared, + TextColumnFault::NotText { + data_type: "INT".into(), + }, + TextColumnFault::NotAColumn, + TextColumnFault::NotIndexed, + ] { + let original = ErrorCode::TextColumn { + collection: "docs".into(), + column: "title".into(), + fault, + }; + let wire = DataPlaneErrorCode::from(original.clone()); + assert_eq!(ErrorCode::from(wire), original); + } + } #[test] fn division_by_zero_survives_the_wire_hop() { @@ -508,6 +381,49 @@ mod tests { assert_eq!(ErrorCode::from(wire), original); } + /// A value refusal crosses the hop as its own verdict, from a Data-Plane + /// code and from a Control-Plane error alike. + #[test] + fn value_refusals_cross_the_hop_verbatim() { + let detail = "column 'n': cannot parse 'x' as INT".to_string(); + for original in [ + ErrorCode::InvalidTextRepresentation { + detail: detail.clone(), + }, + ErrorCode::DatatypeMismatch { + detail: detail.clone(), + }, + ErrorCode::InvalidDatetimeFormat { + detail: detail.clone(), + }, + ErrorCode::DatetimeFieldOverflow { + detail: detail.clone(), + }, + ] { + let wire = DataPlaneErrorCode::from(original.clone()); + assert_eq!(ErrorCode::from(wire), original); + } + match execution_error_to_typed(crate::Error::InvalidTextRepresentation { + detail: detail.clone(), + }) { + TypedClusterError::DataPlane { code } => assert_eq!( + code, + DataPlaneErrorCode::InvalidTextRepresentation { + detail: detail.clone() + } + ), + other => panic!("expected DataPlane, got {other:?}"), + } + match execution_error_to_typed(crate::Error::DatatypeMismatch { + detail: detail.clone(), + }) { + TypedClusterError::DataPlane { code } => { + assert_eq!(code, DataPlaneErrorCode::DatatypeMismatch { detail }) + } + other => panic!("expected DataPlane, got {other:?}"), + } + } + #[test] fn execution_error_keeps_a_data_plane_verdict_typed() { let typed = execution_error_to_typed(crate::Error::DataPlane(ErrorCode::DivisionByZero)); @@ -599,4 +515,29 @@ mod tests { let wire = DataPlaneErrorCode::from(original.clone()); assert_eq!(ErrorCode::from(wire), original); } + + /// A failed rollback crosses the hop with the typed cause of its + /// reverse write, nested codes included. + #[test] + fn rollback_failed_keeps_its_typed_cause_across_the_hop() { + for cause in [ + None, + Some(Box::new(ErrorCode::Internal { + detail: "storage error (sparse): commit".into(), + })), + Some(Box::new(ErrorCode::RollbackFailed { + entry_index: 1, + detail: "inner".into(), + cause: Some(Box::new(ErrorCode::DivisionByZero)), + })), + ] { + let original = ErrorCode::RollbackFailed { + entry_index: 4, + detail: "restoring row r1".into(), + cause, + }; + let wire = DataPlaneErrorCode::from(original.clone()); + assert_eq!(ErrorCode::from(wire), original); + } + } } diff --git a/nodedb/src/control/cluster/data_plane_fault_wire.rs b/nodedb/src/control/cluster/data_plane_fault_wire.rs new file mode 100644 index 000000000..a43c67416 --- /dev/null +++ b/nodedb/src/control/cluster/data_plane_fault_wire.rs @@ -0,0 +1,74 @@ +// SPDX-License-Identifier: BUSL-1.1 + +//! Wire forms of the payload enums a Data-Plane [`ErrorCode`] carries. +//! +//! Each local type and its wire mirror live in different crates, so every +//! mapping is a function, not a `From` impl. Every match is exhaustive, so a +//! new variant fails to compile here until it is mirrored on the wire. +//! +//! [`ErrorCode`]: crate::bridge::envelope::ErrorCode + +use nodedb_cluster::rpc_codec::{ + DataPlaneCounterFault, DataPlaneSyncHold, DataPlaneTextColumnFault, +}; +use nodedb_types::text_search::TextColumnFault; + +use crate::bridge::envelope::{CounterFault, SyncHold}; + +/// The wire form of a sync hold. +pub(super) fn sync_hold_to_wire(hold: SyncHold) -> DataPlaneSyncHold { + match hold { + SyncHold::Duplicate => DataPlaneSyncHold::Duplicate, + SyncHold::Fenced => DataPlaneSyncHold::Fenced, + SyncHold::Gap { expected } => DataPlaneSyncHold::Gap { expected }, + } +} + +/// The sync hold a wire form names. +pub(super) fn sync_hold_from_wire(hold: DataPlaneSyncHold) -> SyncHold { + match hold { + DataPlaneSyncHold::Duplicate => SyncHold::Duplicate, + DataPlaneSyncHold::Fenced => SyncHold::Fenced, + DataPlaneSyncHold::Gap { expected } => SyncHold::Gap { expected }, + } +} + +/// The wire form of a counter fault. +pub(super) fn counter_fault_to_wire(fault: CounterFault) -> DataPlaneCounterFault { + match fault { + CounterFault::NotAnInteger => DataPlaneCounterFault::NotAnInteger, + CounterFault::NotAFloat => DataPlaneCounterFault::NotAFloat, + CounterFault::IntegerOverflow => DataPlaneCounterFault::IntegerOverflow, + CounterFault::NonFinite => DataPlaneCounterFault::NonFinite, + } +} + +/// The counter fault a wire form carries. +pub(super) fn counter_fault_from_wire(fault: DataPlaneCounterFault) -> CounterFault { + match fault { + DataPlaneCounterFault::NotAnInteger => CounterFault::NotAnInteger, + DataPlaneCounterFault::NotAFloat => CounterFault::NotAFloat, + DataPlaneCounterFault::IntegerOverflow => CounterFault::IntegerOverflow, + DataPlaneCounterFault::NonFinite => CounterFault::NonFinite, + } +} + +/// The wire form of a text-column fault. +pub(super) fn text_column_fault_to_wire(fault: TextColumnFault) -> DataPlaneTextColumnFault { + match fault { + TextColumnFault::Undeclared => DataPlaneTextColumnFault::Undeclared, + TextColumnFault::NotText { data_type } => DataPlaneTextColumnFault::NotText { data_type }, + TextColumnFault::NotAColumn => DataPlaneTextColumnFault::NotAColumn, + TextColumnFault::NotIndexed => DataPlaneTextColumnFault::NotIndexed, + } +} + +/// The text-column fault a wire form carries. +pub(super) fn text_column_fault_from_wire(fault: DataPlaneTextColumnFault) -> TextColumnFault { + match fault { + DataPlaneTextColumnFault::Undeclared => TextColumnFault::Undeclared, + DataPlaneTextColumnFault::NotText { data_type } => TextColumnFault::NotText { data_type }, + DataPlaneTextColumnFault::NotAColumn => TextColumnFault::NotAColumn, + DataPlaneTextColumnFault::NotIndexed => TextColumnFault::NotIndexed, + } +} diff --git a/nodedb/src/control/cluster/execution_error_wire.rs b/nodedb/src/control/cluster/execution_error_wire.rs new file mode 100644 index 000000000..7d490a588 --- /dev/null +++ b/nodedb/src/control/cluster/execution_error_wire.rs @@ -0,0 +1,192 @@ +// SPDX-License-Identifier: BUSL-1.1 + +//! The one mapping every cross-node executor uses to answer a remote caller +//! with a local execution error. +//! +//! The match is exhaustive with no catch-all, so a new `Error` variant fails +//! to compile here until it picks how it crosses the node hop. + +use nodedb_cluster::rpc_codec::{DataPlaneErrorCode, TypedClusterError}; + +/// Map a local-execution [`crate::Error`] to the wire error a remote caller +/// receives. +/// +/// A Data-Plane verdict crosses verbatim as `TypedClusterError::DataPlane`, so +/// the coordinator rebuilds `Error::DataPlane(code)` and renders the SQLSTATE +/// single-node execution renders. Every other error keeps its own numeric +/// classification from `NodeDbError::from(err).code()` — never a hardcoded +/// plan-decode code, which will misname what failed. +pub(crate) fn execution_error_to_typed(err: crate::Error) -> TypedClusterError { + match err { + crate::Error::DataPlane(code) => TypedClusterError::DataPlane { code: code.into() }, + // A statement that ran out of time keeps the wire's own deadline + // variant, which the coordinator rebuilds as `Error::DeadlineExceeded`. + // Folding it into `Internal` will report a client's own timeout as an + // internal failure once it crossed a node boundary. + crate::Error::DeadlineExceeded { .. } => { + TypedClusterError::DeadlineExceeded { elapsed_ms: 0 } + } + // A redirect crosses as the wire's own redirect, with the leader and + // the term this node knows it at, so the coordinator moves its + // routing hint and retries against that leader. + not_leader @ crate::Error::NotLeader { .. } => TypedClusterError::from(not_leader), + // A Calvin abort keeps its verdict and a schema change stays + // retryable, so a routed submit answers as a local one does. + typed @ (crate::Error::CalvinSerializationConflict + | crate::Error::CalvinParticipantError + | crate::Error::RetryableSchemaChanged { .. }) => TypedClusterError::from(typed), + // A Control-Plane constraint refusal crosses verbatim, same as a + // Data-Plane verdict, so the coordinator answers 23502 vs 23505 + // instead of flattening both into one numeric class. + crate::Error::RejectedConstraint { + collection, + constraint, + detail, + } => TypedClusterError::RejectedConstraint { + collection, + constraint, + detail, + }, + // A capacity refusal crosses as its own verdict, so the coordinator + // answers the retryable overload class. + capacity @ crate::Error::DispatchCapacity { .. } => TypedClusterError::DataPlane { + code: DataPlaneErrorCode::DispatchCapacity { + reason: capacity.to_string(), + }, + }, + // A value refusal crosses as the Data-Plane verdict of the same name, + // so the coordinator answers its exact SQLSTATE. + crate::Error::InvalidTextRepresentation { detail } => TypedClusterError::DataPlane { + code: DataPlaneErrorCode::InvalidTextRepresentation { detail }, + }, + crate::Error::DatatypeMismatch { detail } => TypedClusterError::DataPlane { + code: DataPlaneErrorCode::DatatypeMismatch { detail }, + }, + crate::Error::InvalidDatetimeFormat { detail } => TypedClusterError::DataPlane { + code: DataPlaneErrorCode::InvalidDatetimeFormat { detail }, + }, + crate::Error::DatetimeFieldOverflow { detail } => TypedClusterError::DataPlane { + code: DataPlaneErrorCode::DatetimeFieldOverflow { detail }, + }, + // Every other error crosses as its public numeric code and message. + // The coordinator renders the SQLSTATE that code maps to. + other @ (crate::Error::TxnOverlayMemoryExceeded { .. } + | crate::Error::RejectedAuthz { .. } + | crate::Error::OffsetRegression { .. } + | crate::Error::ConflictRetry { .. } + | crate::Error::RejectedPrevalidation { .. } + | crate::Error::RetryableRefusal { .. } + | crate::Error::AppendOnlyViolation { .. } + | crate::Error::BalanceViolation { .. } + | crate::Error::MaterializedSumTargetNotFound { .. } + | crate::Error::MaterializedSumResolutionMissing { .. } + | crate::Error::PeriodLocked { .. } + | crate::Error::PeriodLockMisconfigured { .. } + | crate::Error::RetentionViolation { .. } + | crate::Error::LegalHoldActive { .. } + | crate::Error::StateTransitionViolation { .. } + | crate::Error::TransitionCheckViolation { .. } + | crate::Error::TypeGuardViolation { .. } + | crate::Error::TypeMismatch { .. } + | crate::Error::InsufficientBalance { .. } + | crate::Error::RateExceeded { .. } + | crate::Error::CollectionNotFound { .. } + | crate::Error::DocumentNotFound { .. } + | crate::Error::CollectionDeactivated { .. } + | crate::Error::VShardAdmissionCapacityExceeded { .. } + | crate::Error::CrdtAdmissionRetriesExhausted { .. } + | crate::Error::CrdtAdmissionInvalidPlan { .. } + | crate::Error::CrdtAdmissionCallerFence + | crate::Error::CrdtApplyRequiresAdmission + | crate::Error::CrdtApplyForbiddenInTransaction + | crate::Error::NotInTransactionBlock { .. } + | crate::Error::CrdtAdmissionTimeout { .. } + | crate::Error::NoLeader { .. } + | crate::Error::CrossCollectionNotColocated { .. } + | crate::Error::CloneWriteRequiresMaterialize { .. } + | crate::Error::BadRequest { .. } + | crate::Error::BackupTenantMismatch { .. } + | crate::Error::BackupKeyMismatch + | crate::Error::QuotaOvercommit { .. } + | crate::Error::PlanError { .. } + | crate::Error::FeatureNotSupported { .. } + | crate::Error::UndefinedFunction { .. } + | crate::Error::UndefinedObject { .. } + | crate::Error::ObjectNotInPrerequisiteState { .. } + | crate::Error::UndefinedColumn { .. } + | crate::Error::TextColumn { .. } + | crate::Error::AmbiguousColumn { .. } + | crate::Error::UnknownStrictField { .. } + | crate::Error::DivisionByZero + | crate::Error::DataException { .. } + | crate::Error::NumericValueOutOfRange { .. } + | crate::Error::InvalidLimitValue { .. } + | crate::Error::RetryableLeaderChange { .. } + | crate::Error::CommittedResultUnavailable { .. } + | crate::Error::ProposalOutcomeUnknown { .. } + | crate::Error::GroupQuorumUnavailable { .. } + | crate::Error::GroupMarksUnavailable { .. } + | crate::Error::BackupCaptureMoved { .. } + | crate::Error::MetadataLeaderUnavailable + | crate::Error::AuthorizationStateBehind { .. } + | crate::Error::LinearizableReadRefused { .. } + | crate::Error::ExecutionLimitExceeded { .. } + | crate::Error::LimitExceeded { .. } + | crate::Error::Wal(_) + | crate::Error::Dispatch { .. } + | crate::Error::Storage { .. } + | crate::Error::ColdStorage { .. } + | crate::Error::Serialization { .. } + | crate::Error::Codec { .. } + | crate::Error::SegmentCorrupted { .. } + | crate::Error::MemoryExhausted { .. } + | crate::Error::Backpressure { .. } + | crate::Error::Crdt(_) + | crate::Error::Io(_) + | crate::Error::Config { .. } + | crate::Error::Encryption { .. } + | crate::Error::Bridge { .. } + | crate::Error::VersionCompat { .. } + | crate::Error::RestoreTargetNotEmpty { .. } + | crate::Error::RestoreVerificationFailed { .. } + | crate::Error::Internal { .. } + | crate::Error::Shaping(_) + | crate::Error::RemoteTyped { .. } + | crate::Error::Ddl(_) + | crate::Error::DescriptorVersionAnomaly { .. } + | crate::Error::CollectionPurgeRowMissing { .. } + | crate::Error::CollectionUnstamped { .. } + | crate::Error::CatalogIntegrityViolation { .. } + | crate::Error::Promql(_) + | crate::Error::DependentObjectsExist { .. } + | crate::Error::RoleInUse { .. } + | crate::Error::CascadeCycle { .. } + | crate::Error::CrossShardInExplicitTransaction + | crate::Error::SequencerUnavailable + | crate::Error::SessionCapExceeded { .. } + | crate::Error::SessionIdleTimeout + | crate::Error::SessionTokenExpired + | crate::Error::SessionKilledByAdmin + | crate::Error::SessionUserDropped + | crate::Error::OidcProviderTenantUnbound + | crate::Error::OidcProviderTenantUnavailable { .. } + | crate::Error::ExternalRoleUndefined { .. } + | crate::Error::OidcNoDefaultDatabase { .. } + | crate::Error::TenantVectorDimExceeded { .. } + | crate::Error::TenantGraphDepthExceeded { .. } + | crate::Error::RoleInheritanceCycle { .. } + | crate::Error::RoleInheritanceDepthExceeded { .. } + | crate::Error::OllpExhausted { .. } + | crate::Error::MirrorReadOnly { .. } + | crate::Error::StaleReadNotLeader { .. }) => numeric_typed(other), + } +} + +/// The wire error for a local error with no typed wire carrier: its public +/// numeric code from `NodeDbError::from(err).code()`, and its message. The +/// coordinator rebuilds it as `Error::RemoteTyped`. +pub(crate) fn numeric_typed(err: crate::Error) -> TypedClusterError { + let message = err.to_string(); + let code = u32::from(nodedb_types::error::NodeDbError::from(err).code().0); + TypedClusterError::Internal { code, message } +} diff --git a/nodedb/src/control/cluster/metadata_applier/wedge.rs b/nodedb/src/control/cluster/metadata_applier/wedge.rs index a457fa543..95d7ff076 100644 --- a/nodedb/src/control/cluster/metadata_applier/wedge.rs +++ b/nodedb/src/control/cluster/metadata_applier/wedge.rs @@ -113,10 +113,16 @@ pub fn classify(error: &crate::Error) -> ApplyFailureClass { | crate::Error::UndefinedObject { .. } | crate::Error::ObjectNotInPrerequisiteState { .. } | crate::Error::UndefinedColumn { .. } + | crate::Error::TextColumn { .. } | crate::Error::AmbiguousColumn { .. } | crate::Error::UnknownStrictField { .. } | crate::Error::DivisionByZero | crate::Error::DataException { .. } + | crate::Error::NumericValueOutOfRange { .. } + | crate::Error::InvalidTextRepresentation { .. } + | crate::Error::DatatypeMismatch { .. } + | crate::Error::InvalidDatetimeFormat { .. } + | crate::Error::DatetimeFieldOverflow { .. } | crate::Error::InvalidLimitValue { .. } | crate::Error::RetryableSchemaChanged { .. } | crate::Error::RetryableLeaderChange { .. } diff --git a/nodedb/src/control/cluster/mod.rs b/nodedb/src/control/cluster/mod.rs index 95d86d319..8f8e0abdf 100644 --- a/nodedb/src/control/cluster/mod.rs +++ b/nodedb/src/control/cluster/mod.rs @@ -14,7 +14,9 @@ pub mod calvin; pub mod calvin_snapshot; pub mod core_stall; pub mod data_plane_error_wire; +pub mod data_plane_fault_wire; pub mod decommission_bridge; +pub mod execution_error_wire; pub mod handle; pub mod init; pub mod leased_read; diff --git a/nodedb/src/control/gateway/error_map/class_parity/code_samples.rs b/nodedb/src/control/gateway/error_map/class_parity/code_samples.rs index 6c677fd60..017889eaa 100644 --- a/nodedb/src/control/gateway/error_map/class_parity/code_samples.rs +++ b/nodedb/src/control/gateway/error_map/class_parity/code_samples.rs @@ -8,7 +8,7 @@ use nodedb_types::sync::wire::SyncProvenance; use crate::bridge::envelope::{CounterFault, ErrorCode, SyncHold}; /// The number of `ErrorCode` variants [`variant_index`] numbers. -pub(super) const VARIANT_COUNT: usize = 43; +pub(super) const VARIANT_COUNT: usize = 50; /// A dense index per variant. Exhaustive, so a new variant fails to compile /// here until it gets an index, and [`every_variant_has_a_sample`] then fails @@ -58,6 +58,13 @@ pub(super) fn variant_index(code: &ErrorCode) -> usize { ErrorCode::TransactionRollback { .. } => 41, ErrorCode::ActiveSqlTransaction { .. } => 42, ErrorCode::DependentObjectsExist { .. } => 10, + ErrorCode::TextColumn { .. } => 43, + ErrorCode::NumericValueOutOfRange { .. } => 44, + ErrorCode::NodeLabelLimit { .. } => 45, + ErrorCode::InvalidTextRepresentation { .. } => 46, + ErrorCode::DatatypeMismatch { .. } => 47, + ErrorCode::InvalidDatetimeFormat { .. } => 48, + ErrorCode::DatetimeFieldOverflow { .. } => 49, } } @@ -159,17 +166,33 @@ pub(super) fn samples() -> Vec { max_depth: 100, }, ErrorCode::UndefinedColumn { column: "x".into() }, + ErrorCode::TextColumn { + collection: collection(), + column: "x".into(), + fault: nodedb_types::text_search::TextColumnFault::NotIndexed, + }, + ErrorCode::TextColumn { + collection: collection(), + column: "x".into(), + fault: nodedb_types::text_search::TextColumnFault::NotAColumn, + }, ErrorCode::Internal { detail: text() }, ErrorCode::Unsupported { detail: text() }, ErrorCode::RollbackFailed { entry_index: 0, detail: text(), + cause: Some(Box::new(ErrorCode::Internal { detail: text() })), }, ErrorCode::OllpRetryRequired, ErrorCode::TxnOverlayMemoryExceeded { limit: 1 << 20 }, ErrorCode::DivisionByZero, ErrorCode::UndefinedFunction { name: "f".into() }, ErrorCode::DataException { detail: text() }, + ErrorCode::NumericValueOutOfRange { detail: text() }, + ErrorCode::InvalidTextRepresentation { detail: text() }, + ErrorCode::DatatypeMismatch { detail: text() }, + ErrorCode::InvalidDatetimeFormat { detail: text() }, + ErrorCode::DatetimeFieldOverflow { detail: text() }, ErrorCode::DispatchCapacity { reason: text() }, ErrorCode::ExpiredBeforeExecution, ErrorCode::BadRequest { detail: text() }, @@ -179,6 +202,11 @@ pub(super) fn samples() -> Vec { object: "role \"analyst\"".into(), detail: text(), }, + ErrorCode::NodeLabelLimit { + node: "alice".into(), + label: "Person".into(), + limit: 64, + }, ]; for constraint in [ "not_null", diff --git a/nodedb/src/control/gateway/error_map/class_parity/error_index.rs b/nodedb/src/control/gateway/error_map/class_parity/error_index.rs index 9e6da3d79..d49285b9f 100644 --- a/nodedb/src/control/gateway/error_map/class_parity/error_index.rs +++ b/nodedb/src/control/gateway/error_map/class_parity/error_index.rs @@ -3,7 +3,7 @@ //! A dense index per `crate::Error` variant. /// The number of `crate::Error` variants [`error_variant_index`] numbers. -pub(super) const ERROR_VARIANT_COUNT: usize = 115; +pub(super) const ERROR_VARIANT_COUNT: usize = 121; /// A dense index per `crate::Error` variant. Exhaustive, so a new variant /// fails to compile here until it gets an index, and @@ -61,6 +61,7 @@ pub(crate) fn error_variant_index(err: &crate::Error) -> usize { E::UndefinedObject { .. } => 48, E::ObjectNotInPrerequisiteState { .. } => 49, E::UndefinedColumn { .. } => 50, + E::TextColumn { .. } => 115, E::AmbiguousColumn { .. } => 51, E::UnknownStrictField { .. } => 52, E::DivisionByZero => 53, @@ -127,5 +128,10 @@ pub(crate) fn error_variant_index(err: &crate::Error) -> usize { E::RestoreVerificationFailed { .. } => 113, E::BackupCaptureMoved { .. } => 114, E::CollectionUnstamped { .. } => 37, + E::NumericValueOutOfRange { .. } => 116, + E::InvalidTextRepresentation { .. } => 117, + E::DatatypeMismatch { .. } => 118, + E::InvalidDatetimeFormat { .. } => 119, + E::DatetimeFieldOverflow { .. } => 120, } } diff --git a/nodedb/src/control/gateway/error_map/class_parity/error_parity.rs b/nodedb/src/control/gateway/error_map/class_parity/error_parity.rs index 7c7909189..32470284d 100644 --- a/nodedb/src/control/gateway/error_map/class_parity/error_parity.rs +++ b/nodedb/src/control/gateway/error_map/class_parity/error_parity.rs @@ -109,6 +109,11 @@ fn classified_sqlstates() -> Vec<(usize, &'static str)> { (107, sqlstate::STALE_READ_NOT_LEADER), (108, sqlstate::DEPENDENT_OBJECTS_STILL_EXIST), (109, sqlstate::DEPENDENT_OBJECTS_STILL_EXIST), + (116, sqlstate::NUMERIC_VALUE_OUT_OF_RANGE), + (117, sqlstate::INVALID_TEXT_REPRESENTATION), + (118, sqlstate::DATATYPE_MISMATCH), + (119, sqlstate::INVALID_DATETIME_FORMAT), + (120, sqlstate::DATETIME_FIELD_OVERFLOW), ] } diff --git a/nodedb/src/control/gateway/error_map/class_parity/error_samples.rs b/nodedb/src/control/gateway/error_map/class_parity/error_samples.rs index 27c125fce..0280a484d 100644 --- a/nodedb/src/control/gateway/error_map/class_parity/error_samples.rs +++ b/nodedb/src/control/gateway/error_map/class_parity/error_samples.rs @@ -192,6 +192,18 @@ pub(crate) fn error_samples() -> Vec { detail: text(), }, E::UndefinedColumn { column: "x".into() }, + E::TextColumn { + collection: collection(), + column: "x".into(), + fault: nodedb_types::text_search::TextColumnFault::NotIndexed, + }, + E::TextColumn { + collection: collection(), + column: "n".into(), + fault: nodedb_types::text_search::TextColumnFault::NotText { + data_type: "INT".into(), + }, + }, E::AmbiguousColumn { column: "id".into(), }, @@ -201,6 +213,11 @@ pub(crate) fn error_samples() -> Vec { }, E::DivisionByZero, E::DataException { detail: text() }, + E::NumericValueOutOfRange { detail: text() }, + E::InvalidTextRepresentation { detail: text() }, + E::DatatypeMismatch { detail: text() }, + E::InvalidDatetimeFormat { detail: text() }, + E::DatetimeFieldOverflow { detail: text() }, E::InvalidLimitValue { clause: "LIMIT", value: "-1".into(), diff --git a/nodedb/src/control/gateway/error_map/resp.rs b/nodedb/src/control/gateway/error_map/resp.rs index 086a60bdf..0da6aa8e4 100644 --- a/nodedb/src/control/gateway/error_map/resp.rs +++ b/nodedb/src/control/gateway/error_map/resp.rs @@ -89,10 +89,16 @@ impl GatewayErrorMap { | Error::UndefinedObject { .. } | Error::ObjectNotInPrerequisiteState { .. } | Error::UndefinedColumn { .. } + | Error::TextColumn { .. } | Error::AmbiguousColumn { .. } | Error::UnknownStrictField { .. } | Error::DivisionByZero | Error::DataException { .. } + | Error::NumericValueOutOfRange { .. } + | Error::InvalidTextRepresentation { .. } + | Error::DatatypeMismatch { .. } + | Error::InvalidDatetimeFormat { .. } + | Error::DatetimeFieldOverflow { .. } | Error::InvalidLimitValue { .. } | Error::RetryableLeaderChange { .. } | Error::CommittedResultUnavailable { .. } diff --git a/nodedb/src/control/planner/plan_error_map.rs b/nodedb/src/control/planner/plan_error_map.rs index 29f7dccfd..75168676f 100644 --- a/nodedb/src/control/planner/plan_error_map.rs +++ b/nodedb/src/control/planner/plan_error_map.rs @@ -63,6 +63,16 @@ pub(crate) fn map_plan_error( nodedb_sql::SqlError::AmbiguousColumn { column } => { crate::Error::AmbiguousColumn { column } } + nodedb_sql::SqlError::TextColumn { + collection, + column, + fault, + .. + } => crate::Error::TextColumn { + collection, + column, + fault, + }, // A target/expression count mismatch is a syntax error in PostgreSQL, // so it renders 42601 through `BadRequest`. nodedb_sql::SqlError::Arity { detail } => crate::Error::BadRequest { detail }, @@ -74,11 +84,14 @@ pub(crate) fn map_plan_error( detail: error.to_string(), } } - // A value out of range for its type is a data exception (class `22`), - // the class PostgreSQL and the DDL DEFAULT gate give it. + // A value out of range for its type is `22003` + // (numeric_value_out_of_range), the SQLSTATE PostgreSQL and the DDL + // DEFAULT gate give it. nodedb_sql::SqlError::ConstantOverflow { .. } + | nodedb_sql::SqlError::NumericLiteralOutOfRange { .. } | nodedb_sql::SqlError::IntegerOutOfRange { .. } - | nodedb_sql::SqlError::FloatOutOfRange { .. } => crate::Error::DataException { + | nodedb_sql::SqlError::FloatOutOfRange { .. } + | nodedb_sql::SqlError::DecimalOutOfRange { .. } => crate::Error::NumericValueOutOfRange { detail: error.to_string(), }, // The executor's recursion cap: the program-limit class (`54000`) the @@ -118,15 +131,17 @@ mod tests { use crate::types::TenantId; #[test] - fn a_value_out_of_range_is_a_data_exception() { + fn a_value_out_of_range_is_numeric_value_out_of_range() { let error = nodedb_sql::SqlError::IntegerOutOfRange { column: "qty".into(), value: 1 << 40, declared_type: "integer", }; match map_plan_error(error, TenantId::new(1)) { - crate::Error::DataException { detail } => assert!(detail.contains("out of range")), - other => panic!("expected a data exception, got {other:?}"), + crate::Error::NumericValueOutOfRange { detail } => { + assert!(detail.contains("out of range")) + } + other => panic!("expected numeric value out of range, got {other:?}"), } } diff --git a/nodedb/src/control/server/dispatch_utils/write_abort.rs b/nodedb/src/control/server/dispatch_utils/write_abort.rs index a7e31abd9..ac9cdd029 100644 --- a/nodedb/src/control/server/dispatch_utils/write_abort.rs +++ b/nodedb/src/control/server/dispatch_utils/write_abort.rs @@ -60,6 +60,7 @@ fn is_transient_verdict(code: &ErrorCode) -> bool { ErrorCode::DeadlineExceeded | ErrorCode::ActiveSqlTransaction { .. } | ErrorCode::DependentObjectsExist { .. } + | ErrorCode::NodeLabelLimit { .. } | ErrorCode::RejectedConstraint { .. } | ErrorCode::RejectedPrevalidation { .. } | ErrorCode::SyncRejected { .. } @@ -83,12 +84,18 @@ fn is_transient_verdict(code: &ErrorCode) -> bool { | ErrorCode::InsufficientBalance { .. } | ErrorCode::RecursionDepthExceeded { .. } | ErrorCode::UndefinedColumn { .. } + | ErrorCode::TextColumn { .. } | ErrorCode::Internal { .. } | ErrorCode::Unsupported { .. } | ErrorCode::RollbackFailed { .. } | ErrorCode::DivisionByZero | ErrorCode::UndefinedFunction { .. } | ErrorCode::DataException { .. } + | ErrorCode::NumericValueOutOfRange { .. } + | ErrorCode::InvalidTextRepresentation { .. } + | ErrorCode::DatatypeMismatch { .. } + | ErrorCode::InvalidDatetimeFormat { .. } + | ErrorCode::DatetimeFieldOverflow { .. } | ErrorCode::BadRequest { .. } => false, } } @@ -152,11 +159,22 @@ pub(crate) fn write_definitely_not_applied(code: &ErrorCode) -> bool { // The staging overlay hit its byte budget, so the transaction's writes // were discarded from the overlay and never installed. | ErrorCode::TxnOverlayMemoryExceeded { .. } + // A label write past the node-label cap is refused before any label of + // the statement stays set. + | ErrorCode::NodeLabelLimit { .. } // Expression evaluation failed before producing a value to write. | ErrorCode::DivisionByZero | ErrorCode::UndefinedFunction { .. } | ErrorCode::DataException { .. } - | ErrorCode::UndefinedColumn { .. } => true, + | ErrorCode::NumericValueOutOfRange { .. } + // The column type refused the value before the row encoded. + | ErrorCode::InvalidTextRepresentation { .. } + | ErrorCode::DatatypeMismatch { .. } + | ErrorCode::InvalidDatetimeFormat { .. } + | ErrorCode::DatetimeFieldOverflow { .. } + | ErrorCode::UndefinedColumn { .. } + // A full-text read refused the field before ranking. + | ErrorCode::TextColumn { .. } => true, // NOT established — every one of these can be reported by a request // whose write reached, or can have reached, engine state. Emitting an @@ -244,6 +262,7 @@ mod tests { assert!(!write_definitely_not_applied(&ErrorCode::RollbackFailed { entry_index: 3, detail: "undo failed".into(), + cause: None, })); assert!(!write_definitely_not_applied( &ErrorCode::ResourcesExhausted diff --git a/nodedb/src/control/server/pgwire/types/error_map.rs b/nodedb/src/control/server/pgwire/types/error_map.rs index 80c76eff0..97cfa04b3 100644 --- a/nodedb/src/control/server/pgwire/types/error_map.rs +++ b/nodedb/src/control/server/pgwire/types/error_map.rs @@ -109,6 +109,7 @@ pub fn error_to_sqlstate(err: &crate::Error) -> (&'static str, &'static str, Str sqlstate::UNDEFINED_COLUMN, format!("column \"{column}\" does not exist"), ), + crate::Error::TextColumn { fault, .. } => ("ERROR", fault.sqlstate(), err.to_string()), crate::Error::AmbiguousColumn { column } => ( "ERROR", sqlstate::AMBIGUOUS_COLUMN, @@ -121,6 +122,25 @@ pub fn error_to_sqlstate(err: &crate::Error) -> (&'static str, &'static str, Str crate::Error::DataException { detail } => { ("ERROR", sqlstate::DATA_EXCEPTION, detail.clone()) } + crate::Error::NumericValueOutOfRange { detail } => ( + "ERROR", + sqlstate::NUMERIC_VALUE_OUT_OF_RANGE, + detail.clone(), + ), + crate::Error::InvalidTextRepresentation { detail } => ( + "ERROR", + sqlstate::INVALID_TEXT_REPRESENTATION, + detail.clone(), + ), + crate::Error::DatatypeMismatch { detail } => { + ("ERROR", sqlstate::DATATYPE_MISMATCH, detail.clone()) + } + crate::Error::InvalidDatetimeFormat { detail } => { + ("ERROR", sqlstate::INVALID_DATETIME_FORMAT, detail.clone()) + } + crate::Error::DatetimeFieldOverflow { detail } => { + ("ERROR", sqlstate::DATETIME_FIELD_OVERFLOW, detail.clone()) + } crate::Error::InvalidLimitValue { .. } => { ("ERROR", sqlstate::INVALID_LIMIT_VALUE, err.to_string()) } diff --git a/nodedb/src/control/server/shared/ddl/result.rs b/nodedb/src/control/server/shared/ddl/result.rs index 0b5c52b7a..237f8dbae 100644 --- a/nodedb/src/control/server/shared/ddl/result.rs +++ b/nodedb/src/control/server/shared/ddl/result.rs @@ -248,9 +248,10 @@ pub fn code_for_sqlstate(sqlstate_str: &str) -> ErrorCode { sqlstate::UNDEFINED_COLUMN => ErrorCode::UNDEFINED_COLUMN, sqlstate::AMBIGUOUS_COLUMN => ErrorCode::AMBIGUOUS_COLUMN, sqlstate::DATA_EXCEPTION => ErrorCode::DATA_EXCEPTION, - // A bad parameter value, text representation or datetime format is a - // data exception: class `22`, the class the code renders back. - "22023" | "22P02" | "22007" => ErrorCode::DATA_EXCEPTION, + // A bad parameter value, text representation, datetime format or + // datetime range is a data exception: class `22`, the class the code + // renders back. + "22023" | "22P02" | "22007" | "22008" => ErrorCode::DATA_EXCEPTION, sqlstate::NUMERIC_VALUE_OUT_OF_RANGE => ErrorCode::OVERFLOW, sqlstate::DIVISION_BY_ZERO => ErrorCode::DIVISION_BY_ZERO, sqlstate::INVALID_LIMIT_VALUE => ErrorCode::INVALID_LIMIT_VALUE, diff --git a/nodedb/src/control/server/shared/ddl/sqlstate.rs b/nodedb/src/control/server/shared/ddl/sqlstate.rs index 46cae6f6e..d8499dde2 100644 --- a/nodedb/src/control/server/shared/ddl/sqlstate.rs +++ b/nodedb/src/control/server/shared/ddl/sqlstate.rs @@ -203,6 +203,15 @@ pub fn error_code_to_sqlstate(code: &ErrorCode) -> (&'static str, &'static str, sqlstate::UNDEFINED_COLUMN, format!("column \"{column}\" does not exist"), ), + ErrorCode::TextColumn { + collection, + column, + fault, + } => ( + "ERROR", + fault.sqlstate(), + format!("column \"{column}\" of collection \"{collection}\" {fault}"), + ), ErrorCode::Internal { detail } => ("ERROR", sqlstate::INTERNAL_ERROR, detail.clone()), // Division/modulo by zero. ErrorCode::DivisionByZero => ( @@ -216,6 +225,25 @@ pub fn error_code_to_sqlstate(code: &ErrorCode) -> (&'static str, &'static str, format!("function {name}() does not exist"), ), ErrorCode::DataException { detail } => ("ERROR", sqlstate::DATA_EXCEPTION, detail.clone()), + ErrorCode::NumericValueOutOfRange { detail } => ( + "ERROR", + sqlstate::NUMERIC_VALUE_OUT_OF_RANGE, + detail.clone(), + ), + ErrorCode::InvalidTextRepresentation { detail } => ( + "ERROR", + sqlstate::INVALID_TEXT_REPRESENTATION, + detail.clone(), + ), + ErrorCode::DatatypeMismatch { detail } => { + ("ERROR", sqlstate::DATATYPE_MISMATCH, detail.clone()) + } + ErrorCode::InvalidDatetimeFormat { detail } => { + ("ERROR", sqlstate::INVALID_DATETIME_FORMAT, detail.clone()) + } + ErrorCode::DatetimeFieldOverflow { detail } => { + ("ERROR", sqlstate::DATETIME_FIELD_OVERFLOW, detail.clone()) + } // The same SQLSTATE the Control Plane gives `crate::Error::BadRequest`. ErrorCode::BadRequest { detail } => ("ERROR", sqlstate::SYNTAX_ERROR, detail.clone()), ErrorCode::TransactionRollback { detail } => { @@ -239,14 +267,22 @@ pub fn error_code_to_sqlstate(code: &ErrorCode) -> (&'static str, &'static str, ErrorCode::RollbackFailed { entry_index, detail, - } => ( - "ERROR", - sqlstate::INTERNAL_ERROR, - format!( - "transaction rollback failed at undo entry {entry_index}: {detail}; \ - shard state is unknown — restart required" - ), - ), + cause, + } => { + // The message of the typed cause, as its own code renders it. + let because = cause + .as_deref() + .map(|cause| format!(" ({})", error_code_to_sqlstate(cause).2)) + .unwrap_or_default(); + ( + "ERROR", + sqlstate::INTERNAL_ERROR, + format!( + "transaction rollback failed at undo entry {entry_index}: \ + {detail}{because}; shard state is unknown — restart required" + ), + ) + } // OllpRetryRequired is an internal scheduler signal and must not // reach the pgwire layer as a user-visible error. If it does, surface // it as a serialization failure so clients retry automatically. @@ -268,6 +304,11 @@ pub fn error_code_to_sqlstate(code: &ErrorCode) -> (&'static str, &'static str, split the transaction into smaller batches" ), ), + ErrorCode::NodeLabelLimit { node, label, limit } => ( + "ERROR", + sqlstate::PROGRAM_LIMIT_EXCEEDED, + crate::error_from_data_plane::node_label_limit_message(node, label, *limit), + ), } } diff --git a/nodedb/src/control/server/shared/retry.rs b/nodedb/src/control/server/shared/retry.rs index 624beeb9b..9897ce4a0 100644 --- a/nodedb/src/control/server/shared/retry.rs +++ b/nodedb/src/control/server/shared/retry.rs @@ -116,10 +116,16 @@ impl RetryableSchemaChange for Error { | Error::UndefinedObject { .. } | Error::ObjectNotInPrerequisiteState { .. } | Error::UndefinedColumn { .. } + | Error::TextColumn { .. } | Error::AmbiguousColumn { .. } | Error::UnknownStrictField { .. } | Error::DivisionByZero | Error::DataException { .. } + | Error::NumericValueOutOfRange { .. } + | Error::InvalidTextRepresentation { .. } + | Error::DatatypeMismatch { .. } + | Error::InvalidDatetimeFormat { .. } + | Error::DatetimeFieldOverflow { .. } | Error::InvalidLimitValue { .. } | Error::RetryableLeaderChange { .. } | Error::CommittedResultUnavailable { .. } diff --git a/nodedb/src/control/server/sync/async_dispatch/delta/compensation.rs b/nodedb/src/control/server/sync/async_dispatch/delta/compensation.rs index 92a9820d0..99cfbf7b9 100644 --- a/nodedb/src/control/server/sync/async_dispatch/delta/compensation.rs +++ b/nodedb/src/control/server/sync/async_dispatch/delta/compensation.rs @@ -109,10 +109,16 @@ pub(super) fn compensation_hint_for_dispatch_error(e: &crate::Error) -> Compensa | crate::Error::UndefinedObject { .. } | crate::Error::ObjectNotInPrerequisiteState { .. } | crate::Error::UndefinedColumn { .. } + | crate::Error::TextColumn { .. } | crate::Error::AmbiguousColumn { .. } | crate::Error::UnknownStrictField { .. } | crate::Error::DivisionByZero | crate::Error::DataException { .. } + | crate::Error::NumericValueOutOfRange { .. } + | crate::Error::InvalidTextRepresentation { .. } + | crate::Error::DatatypeMismatch { .. } + | crate::Error::InvalidDatetimeFormat { .. } + | crate::Error::DatetimeFieldOverflow { .. } | crate::Error::InvalidLimitValue { .. } | crate::Error::ExecutionLimitExceeded { .. } | crate::Error::LimitExceeded { .. } @@ -209,6 +215,7 @@ fn compensation_hint_for_code(code: &ErrorCode) -> CompensationHint { | ErrorCode::InsufficientBalance { .. } | ErrorCode::RecursionDepthExceeded { .. } | ErrorCode::UndefinedColumn { .. } + | ErrorCode::TextColumn { .. } | ErrorCode::Internal { .. } | ErrorCode::Unsupported { .. } | ErrorCode::RollbackFailed { .. } @@ -216,9 +223,15 @@ fn compensation_hint_for_code(code: &ErrorCode) -> CompensationHint { | ErrorCode::DivisionByZero | ErrorCode::UndefinedFunction { .. } | ErrorCode::DataException { .. } + | ErrorCode::NumericValueOutOfRange { .. } + | ErrorCode::InvalidTextRepresentation { .. } + | ErrorCode::DatatypeMismatch { .. } + | ErrorCode::InvalidDatetimeFormat { .. } + | ErrorCode::DatetimeFieldOverflow { .. } | ErrorCode::BadRequest { .. } | ErrorCode::ActiveSqlTransaction { .. } - | ErrorCode::DependentObjectsExist { .. }) => CompensationHint::Custom { + | ErrorCode::DependentObjectsExist { .. } + | ErrorCode::NodeLabelLimit { .. }) => CompensationHint::Custom { constraint: "apply_failed".into(), detail: format!("{other:?}"), }, diff --git a/nodedb/src/control/server/sync/refusal.rs b/nodedb/src/control/server/sync/refusal.rs index d8dd39da2..d50436b29 100644 --- a/nodedb/src/control/server/sync/refusal.rs +++ b/nodedb/src/control/server/sync/refusal.rs @@ -53,6 +53,7 @@ pub(super) fn retryable_refusal_reason(error: &crate::Error) -> Option<&str> { | ErrorCode::CollectionDraining { .. } | ErrorCode::RecursionDepthExceeded { .. } | ErrorCode::UndefinedColumn { .. } + | ErrorCode::TextColumn { .. } | ErrorCode::Internal { .. } | ErrorCode::Unsupported { .. } | ErrorCode::RollbackFailed { .. } @@ -61,12 +62,18 @@ pub(super) fn retryable_refusal_reason(error: &crate::Error) -> Option<&str> { | ErrorCode::DivisionByZero | ErrorCode::UndefinedFunction { .. } | ErrorCode::DataException { .. } + | ErrorCode::NumericValueOutOfRange { .. } + | ErrorCode::InvalidTextRepresentation { .. } + | ErrorCode::DatatypeMismatch { .. } + | ErrorCode::InvalidDatetimeFormat { .. } + | ErrorCode::DatetimeFieldOverflow { .. } | ErrorCode::DispatchCapacity { .. } | ErrorCode::ExpiredBeforeExecution | ErrorCode::BadRequest { .. } | ErrorCode::TransactionRollback { .. } | ErrorCode::ActiveSqlTransaction { .. } - | ErrorCode::DependentObjectsExist { .. } => None, + | ErrorCode::DependentObjectsExist { .. } + | ErrorCode::NodeLabelLimit { .. } => None, }, crate::Error::RejectedConstraint { .. } | crate::Error::TxnOverlayMemoryExceeded { .. } @@ -116,10 +123,16 @@ pub(super) fn retryable_refusal_reason(error: &crate::Error) -> Option<&str> { | crate::Error::UndefinedObject { .. } | crate::Error::ObjectNotInPrerequisiteState { .. } | crate::Error::UndefinedColumn { .. } + | crate::Error::TextColumn { .. } | crate::Error::AmbiguousColumn { .. } | crate::Error::UnknownStrictField { .. } | crate::Error::DivisionByZero | crate::Error::DataException { .. } + | crate::Error::NumericValueOutOfRange { .. } + | crate::Error::InvalidTextRepresentation { .. } + | crate::Error::DatatypeMismatch { .. } + | crate::Error::InvalidDatetimeFormat { .. } + | crate::Error::DatetimeFieldOverflow { .. } | crate::Error::InvalidLimitValue { .. } | crate::Error::RetryableSchemaChanged { .. } | crate::Error::RetryableLeaderChange { .. } @@ -271,10 +284,16 @@ fn is_indeterminate(error: &crate::Error) -> bool { | crate::Error::UndefinedObject { .. } | crate::Error::ObjectNotInPrerequisiteState { .. } | crate::Error::UndefinedColumn { .. } + | crate::Error::TextColumn { .. } | crate::Error::AmbiguousColumn { .. } | crate::Error::UnknownStrictField { .. } | crate::Error::DivisionByZero | crate::Error::DataException { .. } + | crate::Error::NumericValueOutOfRange { .. } + | crate::Error::InvalidTextRepresentation { .. } + | crate::Error::DatatypeMismatch { .. } + | crate::Error::InvalidDatetimeFormat { .. } + | crate::Error::DatetimeFieldOverflow { .. } | crate::Error::InvalidLimitValue { .. } | crate::Error::ExecutionLimitExceeded { .. } | crate::Error::LimitExceeded { .. } @@ -361,6 +380,7 @@ fn is_indeterminate_code(code: &ErrorCode) -> bool { | ErrorCode::InsufficientBalance { .. } | ErrorCode::RecursionDepthExceeded { .. } | ErrorCode::UndefinedColumn { .. } + | ErrorCode::TextColumn { .. } | ErrorCode::Internal { .. } | ErrorCode::Unsupported { .. } | ErrorCode::RollbackFailed { .. } @@ -368,9 +388,15 @@ fn is_indeterminate_code(code: &ErrorCode) -> bool { | ErrorCode::DivisionByZero | ErrorCode::UndefinedFunction { .. } | ErrorCode::DataException { .. } + | ErrorCode::NumericValueOutOfRange { .. } + | ErrorCode::InvalidTextRepresentation { .. } + | ErrorCode::DatatypeMismatch { .. } + | ErrorCode::InvalidDatetimeFormat { .. } + | ErrorCode::DatetimeFieldOverflow { .. } | ErrorCode::BadRequest { .. } | ErrorCode::ActiveSqlTransaction { .. } - | ErrorCode::DependentObjectsExist { .. } => false, + | ErrorCode::DependentObjectsExist { .. } + | ErrorCode::NodeLabelLimit { .. } => false, } } diff --git a/nodedb/src/data/executor/dispatch/array/mutate.rs b/nodedb/src/data/executor/dispatch/array/mutate.rs index 7f7c11b3c..ce5be1428 100644 --- a/nodedb/src/data/executor/dispatch/array/mutate.rs +++ b/nodedb/src/data/executor/dispatch/array/mutate.rs @@ -159,12 +159,7 @@ impl CoreLoop { .note_applied(crate::types::Lsn::new(wal_lsn)); } if let Err(e) = self.flush_array(array_id) { - return self.response_error( - task, - ErrorCode::Internal { - detail: format!("array flush: {e}"), - }, - ); + return self.response_error(task, ErrorCode::from(e)); } encode_count_response(self, task, "flushed", 1) } @@ -230,12 +225,7 @@ impl CoreLoop { if self.array_engine.is_open(array_id) && let Err(e) = self.flush_array(array_id) { - return self.response_error( - task, - ErrorCode::Internal { - detail: format!("array rekey flush: {e}"), - }, - ); + return self.response_error(task, ErrorCode::from(e)); } if let Err(e) = self.array_engine.rekey_array(array_id, target) { return self.response_error( @@ -295,12 +285,7 @@ impl CoreLoop { fn encode_count_response(core: &CoreLoop, task: &ExecutionTask, key: &str, n: usize) -> Response { match super::super::super::response_codec::encode_count(key, n) { Ok(bytes) => core.response_with_payload(task, bytes), - Err(e) => core.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ), + Err(e) => core.response_error(task, ErrorCode::from(e)), } } diff --git a/nodedb/src/data/executor/dispatch/array/open.rs b/nodedb/src/data/executor/dispatch/array/open.rs index c69506886..52419d8a7 100644 --- a/nodedb/src/data/executor/dispatch/array/open.rs +++ b/nodedb/src/data/executor/dispatch/array/open.rs @@ -137,12 +137,7 @@ impl CoreLoop { match super::super::super::response_codec::encode_count("opened", 1) { Ok(bytes) => self.response_with_payload(task, bytes), - Err(e) => self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ), + Err(e) => self.response_error(task, ErrorCode::from(e)), } } } diff --git a/nodedb/src/data/executor/dispatch/meta_retention/handlers.rs b/nodedb/src/data/executor/dispatch/meta_retention/handlers.rs index 595ac9508..9e43e718a 100644 --- a/nodedb/src/data/executor/dispatch/meta_retention/handlers.rs +++ b/nodedb/src/data/executor/dispatch/meta_retention/handlers.rs @@ -110,12 +110,7 @@ impl CoreLoop { .unwrap_or_default(); match response_codec::encode_serde(&wm) { Ok(payload) => self.response_with_payload(task, payload), - Err(e) => self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ), + Err(e) => self.response_error(task, ErrorCode::from(e)), } } @@ -138,12 +133,7 @@ impl CoreLoop { }; match response_codec::encode(&entries) { Ok(payload) => self.response_with_payload(task, payload), - Err(e) => self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ), + Err(e) => self.response_error(task, ErrorCode::from(e)), } } @@ -166,12 +156,7 @@ impl CoreLoop { .map(|e| (e.ts, e.value)); match response_codec::encode(&entry) { Ok(payload) => self.response_with_payload(task, payload), - Err(e) => self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ), + Err(e) => self.response_error(task, ErrorCode::from(e)), } } @@ -204,12 +189,7 @@ impl CoreLoop { let payload = (n as u64).to_le_bytes().to_vec(); self.response_with_payload(task, payload) } - Err(e) => self.response_error( - task, - ErrorCode::Internal { - detail: format!("edge_store temporal purge: {e}"), - }, - ), + Err(e) => self.response_error(task, ErrorCode::from(e)), } } @@ -243,12 +223,7 @@ impl CoreLoop { payload.extend_from_slice(&(idx as u64).to_le_bytes()); self.response_with_payload(task, payload) } - Err(e) => self.response_error( - task, - ErrorCode::Internal { - detail: format!("document_strict temporal purge: {e}"), - }, - ), + Err(e) => self.response_error(task, ErrorCode::from(e)), } } @@ -270,12 +245,7 @@ impl CoreLoop { Some(engine) => match engine.purge_history_before(collection, cutoff_system_ms) { Ok(n) => n as u64, Err(e) => { - return self.response_error( - task, - ErrorCode::Internal { - detail: format!("crdt temporal purge: {e}"), - }, - ); + return self.response_error(task, ErrorCode::from(e)); } }, None => 0, @@ -385,12 +355,7 @@ impl CoreLoop { let payload = (n as u64).to_le_bytes().to_vec(); self.response_with_payload(task, payload) } - Err(e) => self.response_error( - task, - ErrorCode::Internal { - detail: format!("columnar temporal purge: {e}"), - }, - ), + Err(e) => self.response_error(task, ErrorCode::from(e)), } } } diff --git a/nodedb/src/data/executor/handlers/bulk_dml/update.rs b/nodedb/src/data/executor/handlers/bulk_dml/update.rs index 27b6a742c..a2b5a9b20 100644 --- a/nodedb/src/data/executor/handlers/bulk_dml/update.rs +++ b/nodedb/src/data/executor/handlers/bulk_dml/update.rs @@ -119,12 +119,7 @@ impl CoreLoop { ) { Ok(ids) => ids, Err(e) => { - return self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ); + return self.response_error(task, ErrorCode::from(e)); } }; @@ -306,7 +301,7 @@ impl CoreLoop { ); // The row's post-image, then one durable redo entry per derived // target row, naming the TARGET collection. - write_set.push(self.stored_row_image( + let image = self.stored_row_image( StoredRow { database_id, tid, @@ -316,7 +311,15 @@ impl CoreLoop { }, &updated_bytes, None, - )); + ); + match image { + Ok(image) => write_set.push(image), + // This row's transaction committed, so a row landed. + Err(e) => { + let code = refusal_after_partial_apply(ErrorCode::from(e)); + return self.refusal_with_landed_rows(task, code, write_set); + } + } write_set.extend(write_hook::target_write_set(&target_writes)); // Published only after the commit succeeded — the same // ordering the reindex helper used when it owned the @@ -372,8 +375,9 @@ impl CoreLoop { // metadata reconstructed only when the live per-row events // were lost — the live path always emits per row. // - // `row_identity` is read again below for `RETURNING`'s `id` field, - // so the event-emit boundary gets a clone rather than the move. + // `row_identity` is read again below for `RETURNING`'s identity + // column, so the event-emit boundary gets a clone rather than the + // move. self.emit_put_event( task, tid, @@ -384,11 +388,15 @@ impl CoreLoop { ); affected += 1; if returning.is_some() { - // `row_identity` only stands in as `id` for a row that - // declares no primary key of its own — overwriting a - // declared key would return a value the client never wrote. + // `row_identity` fills the identity column only for a row + // that lacks it. A declared key keeps the value the client + // wrote. let mut row = nodedb_types::Value::from(doc); - returning_doc::attach_row_id(&mut row, &row_identity); + returning_doc::attach_row_id( + &mut row, + &row_identity, + &self.identity_column(database_id, tid, collection), + ); returned_docs.push(row); } } @@ -400,23 +408,13 @@ impl CoreLoop { let mut response = if let Some(spec) = returning { match returning_rows::build_rows_payload(spec, rls_filters, &returned_docs) { Ok(payload) => self.response_with_payload(task, payload), - Err(e) => self.response_error( - task, - ErrorCode::Internal { - detail: format!("RETURNING encode: {e}"), - }, - ), + Err(e) => self.response_error(task, ErrorCode::from(e)), } } else { let result = serde_json::json!({ "affected": affected }); match response_codec::encode_json_as_msgpack(&result) { Ok(payload) => self.response_with_payload(task, payload), - Err(e) => self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ), + Err(e) => self.response_error(task, ErrorCode::from(e)), } }; response.write_set = write_set; diff --git a/nodedb/src/data/executor/handlers/columnar_read/scan/execute.rs b/nodedb/src/data/executor/handlers/columnar_read/scan/execute.rs index d77b48aee..88a7c64e9 100644 --- a/nodedb/src/data/executor/handlers/columnar_read/scan/execute.rs +++ b/nodedb/src/data/executor/handlers/columnar_read/scan/execute.rs @@ -130,12 +130,7 @@ impl CoreLoop { // Empty result for missing collection. return match response_codec::encode_value_vec(&[]) { Ok(payload) => self.response_with_payload(task, payload), - Err(e) => self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ), + Err(e) => self.response_error(task, ErrorCode::from(e)), }; } }; @@ -235,7 +230,11 @@ impl CoreLoop { // the flushed pass already filled an unsorted limit; otherwise the // loop stops at the limit or the deadline, like the flushed pass. if !block_skipped && (!sort_keys.is_empty() || matched.len() < limit) { - for (row_surrogate, row) in engine.scan_memtable_rows_with_surrogates() { + for scanned in engine.scan_memtable_rows_with_surrogates() { + let (row_surrogate, row) = match scanned { + Ok(scanned) => scanned, + Err(e) => return self.response_error(task, crate::Error::from(e)), + }; // Row-boundary prefilter: skip this row when its surrogate is // absent from the bitmap. Rows without a recorded surrogate // are always included when no prefilter is active; when a @@ -369,12 +368,7 @@ impl CoreLoop { let payload = match response_codec::encode_value_vec(&results) { Ok(payload) => payload, Err(e) => { - return self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ); + return self.response_error(task, ErrorCode::from(e)); } }; diff --git a/nodedb/src/data/executor/handlers/columnar_write/insert.rs b/nodedb/src/data/executor/handlers/columnar_write/insert.rs index ffd1c1d0a..dcd51cb13 100644 --- a/nodedb/src/data/executor/handlers/columnar_write/insert.rs +++ b/nodedb/src/data/executor/handlers/columnar_write/insert.rs @@ -136,13 +136,16 @@ impl CoreLoop { // in-transaction staging path so a staged insert into a brand-new // collection registers the same schema (see // `ensure_columnar_engine_schema` doc comment). - let schema = self.ensure_columnar_engine_schema( + let schema = match self.ensure_columnar_engine_schema( &engine_key, collection, bitemporal, &ndb_rows[0], schema_bytes, - ); + ) { + Ok(schema) => schema, + Err(e) => return self.response_error(task, ErrorCode::from(e)), + }; let outcome = match self.insert_columnar_rows( task, @@ -229,12 +232,7 @@ impl CoreLoop { let json = match response_codec::encode_json_as_msgpack(&result) { Ok(b) => b, Err(e) => { - return self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ); + return self.response_error(task, ErrorCode::from(e)); } }; Response { diff --git a/nodedb/src/data/executor/handlers/control/crdt.rs b/nodedb/src/data/executor/handlers/control/crdt.rs index 1752de5d4..7a05e657e 100644 --- a/nodedb/src/data/executor/handlers/control/crdt.rs +++ b/nodedb/src/data/executor/handlers/control/crdt.rs @@ -24,12 +24,7 @@ impl CoreLoop { Ok(e) => e, Err(e) => { warn!(core = self.core_id, error = %e, "failed to create CRDT engine"); - return self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ); + return self.response_error(task, ErrorCode::from(e)); } }; match engine.read_snapshot(collection, document_id) { @@ -37,12 +32,7 @@ impl CoreLoop { Ok(None) => self.response_error(task, ErrorCode::NotFound), Err(e) => { warn!(core = self.core_id, error = %e, "crdt read snapshot failed"); - self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ) + self.response_error(task, ErrorCode::from(e)) } } } @@ -60,23 +50,13 @@ impl CoreLoop { let engine = match self.get_crdt_engine(task.request.database_id, tenant_id) { Ok(e) => e, Err(e) => { - return self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ); + return self.response_error(task, ErrorCode::from(e)); } }; match engine.read_at_version_json(collection, document_id, version_vector_json) { Ok(Some(json_bytes)) => self.response_with_payload(task, json_bytes), Ok(None) => self.response_error(task, ErrorCode::NotFound), - Err(e) => self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ), + Err(e) => self.response_error(task, ErrorCode::from(e)), } } @@ -90,22 +70,12 @@ impl CoreLoop { let engine = match self.get_crdt_engine(task.request.database_id, tenant_id) { Ok(e) => e, Err(e) => { - return self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ); + return self.response_error(task, ErrorCode::from(e)); } }; match engine.version_vector_json(collection) { Ok(json) => self.response_with_payload(task, json.into_bytes()), - Err(e) => self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ), + Err(e) => self.response_error(task, ErrorCode::from(e)), } } @@ -120,22 +90,12 @@ impl CoreLoop { let engine = match self.get_crdt_engine(task.request.database_id, tenant_id) { Ok(e) => e, Err(e) => { - return self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ); + return self.response_error(task, ErrorCode::from(e)); } }; match engine.export_delta(collection, from_version_json) { Ok(delta) => self.response_with_payload(task, delta), - Err(e) => self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ), + Err(e) => self.response_error(task, ErrorCode::from(e)), } } @@ -152,22 +112,12 @@ impl CoreLoop { let engine = match self.get_crdt_engine(task.request.database_id, tenant_id) { Ok(e) => e, Err(e) => { - return self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ); + return self.response_error(task, ErrorCode::from(e)); } }; match engine.preview_restore_to_version(collection, document_id, target_version_json) { Ok(delta) => self.response_with_payload(task, delta), - Err(e) => self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ), + Err(e) => self.response_error(task, ErrorCode::from(e)), } } @@ -183,21 +133,11 @@ impl CoreLoop { let engine = match self.get_crdt_engine(task.request.database_id, tenant_id) { Ok(e) => e, Err(e) => { - return self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ); + return self.response_error(task, ErrorCode::from(e)); } }; if let Err(e) = engine.compact_at_version(collection, target_version_json) { - return self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ); + return self.response_error(task, ErrorCode::from(e)); } // The Loro docs are in-memory, and boot restores them from the last // checkpoint plus WAL deltas. Without a publish here, a crash restores @@ -208,12 +148,7 @@ impl CoreLoop { .record_flush("crdt", outcome.files_written); self.response_ok(task) } - Err(e) => self.response_error( - task, - ErrorCode::Internal { - detail: format!("history compaction applied but not checkpointed: {e}"), - }, - ), + Err(e) => self.response_error(task, ErrorCode::from(e)), } } @@ -237,12 +172,7 @@ impl CoreLoop { Ok(e) => e, Err(e) => { warn!(core = self.core_id, error = %e, "failed to create CRDT engine"); - return self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ); + return self.response_error(task, ErrorCode::from(e)); } }; match engine.apply_committed_delta_validated( @@ -260,6 +190,9 @@ impl CoreLoop { warn!(core = self.core_id, %reason, "crdt snapshot rejected by constraints"); // Nothing applied, so the record is cancelled and replay never // reaches this rejection. Its dead-letter entry is stored first. + // An entry the store refused is removed from the queue, so + // nothing records the snapshot, as for a queue refusal. The + // store recorded the refusal in the black box. let code = match self.store_crdt_dead_letter( task.request.database_id, tid, @@ -268,21 +201,41 @@ impl CoreLoop { Ok(()) => crate::data::executor::core_loop::crdt_rejection( collection, "snapshot", &reason, ), - Err(error) => ErrorCode::Internal { - detail: format!( + Err(error) => ErrorCode::RetryableRefusal { + reason: format!( "CRDT snapshot for {collection} violates {reason}, and its \ - dead-letter entry could not be stored: {error}" + dead-letter entry could not be stored: {error}; nothing was \ + applied" ), }, }; self.response_error(task, code) } + // Nothing applied and nothing records the snapshot. The apply + // recorded the refusal in the black box. + crate::engine::crdt::tenant_state::ValidatedApplyOutcome::DeadLetterRefused { + violation, + error, + } => self.response_error( + task, + ErrorCode::RetryableRefusal { + reason: format!( + "CRDT snapshot for {collection} violates {violation}, and the \ + dead-letter queue refused it: {error}; nothing was applied" + ), + }, + ), + // The caller sent bytes that do not decode as a snapshot. The + // import ran on a detached candidate, so nothing was applied. crate::engine::crdt::tenant_state::ValidatedApplyOutcome::Malformed => { warn!(core = self.core_id, "crdt snapshot import was malformed"); self.response_error( task, - ErrorCode::Internal { - detail: "malformed CRDT snapshot".into(), + ErrorCode::DataException { + detail: format!( + "CRDT snapshot for {collection} is malformed: its bytes do not \ + decode as a snapshot; nothing was applied" + ), }, ) } @@ -301,6 +254,27 @@ impl CoreLoop { }, ) } + // This node failed to build the candidate; the snapshot itself + // is not at fault, so the refusal is retryable. + crate::engine::crdt::tenant_state::ValidatedApplyOutcome::CandidateUnavailable { + error, + } => { + warn!( + core = self.core_id, + %collection, + %error, + "crdt snapshot import refused: no apply candidate" + ); + self.response_error( + task, + ErrorCode::RetryableRefusal { + reason: format!( + "no apply candidate for CRDT collection {collection}: {error}; \ + nothing was imported" + ), + }, + ) + } } } diff --git a/nodedb/src/data/executor/handlers/control/crdt_constraints.rs b/nodedb/src/data/executor/handlers/control/crdt_constraints.rs index 1b6904d0a..9bb771cde 100644 --- a/nodedb/src/data/executor/handlers/control/crdt_constraints.rs +++ b/nodedb/src/data/executor/handlers/control/crdt_constraints.rs @@ -49,12 +49,7 @@ impl CoreLoop { Ok(e) => e, Err(e) => { warn!(core = self.core_id, error = %e, "failed to create CRDT engine"); - return self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ); + return self.response_error(task, ErrorCode::from(e)); } }; // A `false` return means the incoming version is older than the one @@ -82,12 +77,7 @@ impl CoreLoop { Ok(e) => e, Err(e) => { warn!(core = self.core_id, error = %e, "failed to create CRDT engine"); - return self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ); + return self.response_error(task, ErrorCode::from(e)); } }; // `false` means a stale (older-version) drop was correctly ignored by @@ -117,12 +107,7 @@ impl CoreLoop { Ok(e) => e, Err(e) => { warn!(core = self.core_id, error = %e, "failed to create CRDT engine"); - return self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ); + return self.response_error(task, ErrorCode::from(e)); } }; let constraints = engine.constraints_for_collection(collection); diff --git a/nodedb/src/data/executor/handlers/control/crdt_list.rs b/nodedb/src/data/executor/handlers/control/crdt_list.rs index 7feaf4ee6..3808bc6fb 100644 --- a/nodedb/src/data/executor/handlers/control/crdt_list.rs +++ b/nodedb/src/data/executor/handlers/control/crdt_list.rs @@ -26,12 +26,7 @@ impl CoreLoop { let engine = match self.get_crdt_engine(task.request.database_id, tenant_id) { Ok(e) => e, Err(e) => { - return self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ); + return self.response_error(task, ErrorCode::from(e)); } }; let fields = @@ -47,12 +42,7 @@ impl CoreLoop { }; match engine.list_insert_fields(collection, document_id, list_path, index, &fields) { Ok(()) => self.response_ok(task), - Err(error) => self.response_error( - task, - ErrorCode::Internal { - detail: error.to_string(), - }, - ), + Err(error) => self.response_error(task, ErrorCode::from(error)), } } @@ -70,22 +60,12 @@ impl CoreLoop { let engine = match self.get_crdt_engine(task.request.database_id, tenant_id) { Ok(e) => e, Err(e) => { - return self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ); + return self.response_error(task, ErrorCode::from(e)); } }; match engine.list_delete(collection, document_id, list_path, index) { Ok(()) => self.response_ok(task), - Err(error) => self.response_error( - task, - ErrorCode::Internal { - detail: error.to_string(), - }, - ), + Err(error) => self.response_error(task, ErrorCode::from(error)), } } @@ -104,22 +84,12 @@ impl CoreLoop { let engine = match self.get_crdt_engine(task.request.database_id, tenant_id) { Ok(e) => e, Err(e) => { - return self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ); + return self.response_error(task, ErrorCode::from(e)); } }; match engine.list_move(collection, document_id, list_path, from_index, to_index) { Ok(()) => self.response_ok(task), - Err(error) => self.response_error( - task, - ErrorCode::Internal { - detail: error.to_string(), - }, - ), + Err(error) => self.response_error(task, ErrorCode::from(error)), } } } diff --git a/nodedb/src/data/executor/handlers/control/reindex/csr.rs b/nodedb/src/data/executor/handlers/control/reindex/csr.rs index c5ec2ebe5..800fbea57 100644 --- a/nodedb/src/data/executor/handlers/control/reindex/csr.rs +++ b/nodedb/src/data/executor/handlers/control/reindex/csr.rs @@ -30,40 +30,6 @@ use crate::data::executor::core_loop::CoreLoop; /// it the rebuild is discarded at cutover and the live partition stays. pub const CSR_REBUILD_JOURNAL_MAX_BYTES: usize = 64 << 20; -/// Map a graph-engine error into the crate error. -pub(super) fn graph_err(e: nodedb_graph::GraphError) -> crate::Error { - use nodedb_graph::GraphError; - match e { - // The engine's memory budget, the same class a vector or FTS budget - // refusal has. - GraphError::MemoryBudget(_) => crate::Error::MemoryExhausted { - engine: "graph".to_string(), - }, - // A rebuild already holds the partition's journal: the index is busy. - GraphError::RebuildInProgress => crate::Error::ObjectNotInPrerequisiteState { - object: "graph CSR index".to_string(), - detail: e.to_string(), - }, - other @ (GraphError::LabelOverflow { .. } - | GraphError::NodeOverflow { .. } - | GraphError::WithdrawRefused { .. } - | GraphError::RebuildSuperseded - | GraphError::RebuildJournalOverflow { .. } - | GraphError::RebuildReplayDiverged { .. } - | GraphError::RebuildSnapshotInvalid { .. }) => crate::Error::Storage { - engine: "graph".to_string(), - detail: other.to_string(), - }, - // `GraphError` is `#[non_exhaustive]` and lives in another crate, so - // the compiler requires this arm. A variant this build cannot name is - // a storage fault. - other => crate::Error::Storage { - engine: "graph".to_string(), - detail: other.to_string(), - }, - } -} - impl CoreLoop { /// Start a CSR rebuild for `target` on its own thread. pub(super) fn start_csr_rebuild(&mut self, target: &RebuildTarget) -> crate::Result<()> { @@ -115,9 +81,7 @@ impl CoreLoop { ); return Ok(None); } - let seed = partition - .begin_rebuild(CSR_REBUILD_JOURNAL_MAX_BYTES) - .map_err(graph_err)?; + let seed = partition.begin_rebuild(CSR_REBUILD_JOURNAL_MAX_BYTES)?; info!( target: "nodedb::reindex", core = core_id, @@ -136,9 +100,9 @@ impl CoreLoop { ) -> crate::Result<()> { let memory = self.graph_memory(target); let Some(live) = self.csr.partition_mut(target.database_id, target.tenant_id) else { - return Err(graph_err(nodedb_graph::GraphError::RebuildSuperseded)); + return Err(nodedb_graph::GraphError::RebuildSuperseded.into()); }; - let copy = live.finish_rebuild(rebuilt, memory).map_err(graph_err)?; + let copy = live.finish_rebuild(rebuilt, memory)?; let nodes = copy.node_count(); let edges = copy.edge_count(); self.csr @@ -171,18 +135,3 @@ impl CoreLoop { ) } } - -#[cfg(test)] -mod tests { - use super::*; - - /// A REINDEX refused because a rebuild is running reports a busy index, - /// not a storage fault. - #[test] - fn a_running_rebuild_is_a_busy_index() { - assert!(matches!( - graph_err(nodedb_graph::GraphError::RebuildInProgress), - crate::Error::ObjectNotInPrerequisiteState { .. } - )); - } -} diff --git a/nodedb/src/data/executor/handlers/control/reindex/pending.rs b/nodedb/src/data/executor/handlers/control/reindex/pending.rs index c470e60bb..092839a65 100644 --- a/nodedb/src/data/executor/handlers/control/reindex/pending.rs +++ b/nodedb/src/data/executor/handlers/control/reindex/pending.rs @@ -67,7 +67,7 @@ impl PendingBuild { Self::Csr { token, rx } => match rx.try_recv() { Ok(result) => BuildPoll::Csr { token, - result: result.map_err(super::csr::graph_err), + result: result.map_err(crate::Error::from), }, Err(TryRecvError::Empty) => BuildPoll::Running(Self::Csr { token, rx }), Err(TryRecvError::Disconnected) => BuildPoll::Csr { diff --git a/nodedb/src/data/executor/handlers/control/snapshot.rs b/nodedb/src/data/executor/handlers/control/snapshot.rs index 1bc7a8aad..59c8c3c2d 100644 --- a/nodedb/src/data/executor/handlers/control/snapshot.rs +++ b/nodedb/src/data/executor/handlers/control/snapshot.rs @@ -66,12 +66,7 @@ impl CoreLoop { Ok(e) => e, Err(e) => { warn!(core = self.core_id, error = %e, "failed to create CRDT engine for get policy"); - return self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ); + return self.response_error(task, ErrorCode::from(e)); } }; let policy = engine.get_collection_policy(collection); @@ -98,24 +93,14 @@ impl CoreLoop { Ok(e) => e, Err(e) => { warn!(core = self.core_id, error = %e, "failed to create CRDT engine"); - return self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ); + return self.response_error(task, ErrorCode::from(e)); } }; match engine.set_collection_policy(collection, policy_json) { Ok(()) => self.response_ok(task), Err(e) => { warn!(core = self.core_id, error = %e, "set collection policy failed"); - self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ) + self.response_error(task, ErrorCode::from(e)) } } } @@ -165,12 +150,7 @@ impl CoreLoop { Ok(true) => kept.push((id, bytes)), Ok(false) => {} Err(e) => { - return self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ); + return self.response_error(task, ErrorCode::from(e)); } } } @@ -178,12 +158,7 @@ impl CoreLoop { } Err(e) => { warn!(core = self.core_id, error = %e, "sparse range scan failed"); - return self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ); + return self.response_error(task, ErrorCode::from(e)); } }; @@ -199,18 +174,24 @@ impl CoreLoop { ); match scan_result { Ok(mut docs) => { + let sort_keys = [nodedb_physical::physical_plan::SortKeySpec::column( + field, true, + )]; + let strict_schema = self.strict_schema_for( + task.request.database_id, + crate::types::TenantId::new(tid), + collection, + ); + let decimal_keys = super::super::document::sort::decimal_sort_keys( + &sort_keys, + strict_schema.as_ref(), + ); if let Err(e) = super::super::document::sort::sort_rows( &mut docs, - &[nodedb_physical::physical_plan::SortKeySpec::column( - field, true, - )], + &sort_keys, + &decimal_keys, ) { - return self.response_error( - task, - ErrorCode::Internal { - detail: format!("in-memory sort failed: {e}"), - }, - ); + return self.response_error(task, ErrorCode::from(e)); } docs.truncate(limit); // Raw msgpack passthrough — no decode/re-encode. @@ -224,22 +205,12 @@ impl CoreLoop { match super::super::super::response_codec::encode_raw_document_rows(&rows) { Ok(payload) => return self.response_with_payload(task, payload), Err(e) => { - return self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ); + return self.response_error(task, ErrorCode::from(e)); } } } Err(e) => { - return self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ); + return self.response_error(task, ErrorCode::from(e)); } } } @@ -248,12 +219,7 @@ impl CoreLoop { Ok(payload) => self.response_with_payload(task, payload), Err(e) => { warn!(core = self.core_id, error = %e, "range scan serialization failed"); - self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ) + self.response_error(task, ErrorCode::from(e)) } } } diff --git a/nodedb/src/data/executor/handlers/convert.rs b/nodedb/src/data/executor/handlers/convert.rs index 1e1a561ea..d4d47041a 100644 --- a/nodedb/src/data/executor/handlers/convert.rs +++ b/nodedb/src/data/executor/handlers/convert.rs @@ -86,10 +86,11 @@ impl CoreLoop { /// Convert to strict mode: re-encode each document as a Binary Tuple. /// /// The sparse engine keys a minted row by its storage key, not its - /// client-visible `id`. A `SELECT` synthesizes `id` at read time; this - /// re-encode must do the same before validating and encoding, or a row - /// with no declared primary key loses its identity and the target - /// schema's NOT NULL `id` column rejects it. + /// client-visible identity. A `SELECT` synthesizes the identity under the + /// identity column at read time; this re-encode does the same under the + /// target schema's key column before validating and encoding, or a row + /// that lacks its key loses its identity and the target schema's NOT NULL + /// key column rejects it. A row that holds its key gains no `id`. fn convert_to_strict( &mut self, task: &ExecutionTask, @@ -131,6 +132,8 @@ impl CoreLoop { .iter() .find(|c| c.primary_key) .map(|c| c.name.as_str()); + let identity_column = nodedb_types::declared_key(declared_primary_key) + .unwrap_or(nodedb_types::DEFAULT_IDENTITY_COLUMN); // Scan all existing documents. let database_id = task.request.database_id.as_u64(); @@ -140,12 +143,7 @@ impl CoreLoop { { Ok(d) => d, Err(e) => { - return self.response_error( - task, - ErrorCode::Internal { - detail: format!("scan failed: {e}"), - }, - ); + return self.response_error(task, ErrorCode::from(e)); } }; @@ -153,11 +151,12 @@ impl CoreLoop { for (doc_id, doc_bytes) in &docs { let normalized = sparse_body_to_msgpack(doc_bytes, source_format.as_format_ref()); - // A row with no declared primary key has no client-visible identity - // yet: inject the surrogate's decimal string as `id` so the target - // schema's NOT NULL primary key column has something to validate. + // A row that lacks the target's key column has no client-visible + // identity yet: inject the surrogate's decimal string there so the + // target schema's NOT NULL key column has something to validate. let synth_id = doc_id.to_identity(); - let with_id = msgpack_scan::inject_str_field(&normalized, "id", synth_id.as_str()); + let with_id = + msgpack_scan::inject_str_field(&normalized, identity_column, synth_id.as_str()); // The identity a user recognizes: the declared primary key's value, // or `id`, read from the row itself — never the internal surrogate. let identity = RowIdentity::of_stored_row(&normalized, declared_primary_key, *doc_id); @@ -199,15 +198,7 @@ impl CoreLoop { collection, identity.as_str(), ); - return self.response_error( - task, - ErrorCode::Internal { - detail: format!( - "collection '{collection}': row '{identity}' converted to a tuple \ - that does not decode: {e}" - ), - }, - ); + return self.response_error(task, ErrorCode::from(e)); }; if let Err(e) = self.put_converted_row( ConvertedRow { @@ -219,14 +210,7 @@ impl CoreLoop { &tuple_bytes, &stored_msgpack, ) { - return self.response_error( - task, - ErrorCode::Internal { - detail: format!( - "collection '{collection}': row '{identity}' failed to write converted document_strict body: {e}" - ), - }, - ); + return self.response_error(task, ErrorCode::from(e)); } // Write-through: a point-get after this statement must see the // re-encoded bytes, not a stale cache entry from before the @@ -245,12 +229,7 @@ impl CoreLoop { }); match response_codec::encode_json_as_msgpack(&result) { Ok(payload) => self.response_with_payload(task, payload), - Err(e) => self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ), + Err(e) => self.response_error(task, ErrorCode::from(e)), } } @@ -275,22 +254,12 @@ impl CoreLoop { { Ok(d) => d, Err(e) => { - return self.response_error( - task, - ErrorCode::Internal { - detail: format!("scan failed: {e}"), - }, - ); + return self.response_error(task, ErrorCode::from(e)); } }; let converted = match source_format { SparseBodyFormat::Strict(schema) => { - let declared_primary_key = schema - .columns - .iter() - .find(|c| c.primary_key) - .map(|c| c.name.as_str()); let mut converted = 0u64; for (doc_id, doc_bytes) in &docs { // Before decode, the row's own identity is unreadable — @@ -304,19 +273,8 @@ impl CoreLoop { collection, undecoded_identity.as_str(), ); - return self.response_error( - task, - ErrorCode::Internal { - detail: format!( - "collection '{collection}': row '{undecoded_identity}' failed to convert to {target_type}: {e}" - ), - }, - ); + return self.response_error(task, ErrorCode::from(e)); }; - // The identity a user recognizes: the declared primary - // key's value, or `id`, read from the decoded row. - let identity = RowIdentity::of_stored_row(&mp, declared_primary_key, *doc_id); - if let Err(e) = self.put_converted_row( ConvertedRow { database_id, @@ -327,14 +285,7 @@ impl CoreLoop { &mp, &mp, ) { - return self.response_error( - task, - ErrorCode::Internal { - detail: format!( - "collection '{collection}': row '{identity}' failed to write converted {target_type} body: {e}" - ), - }, - ); + return self.response_error(task, ErrorCode::from(e)); } // Write-through: a point-get after this statement must see // the re-encoded bytes, not a stale cache entry from before @@ -355,12 +306,7 @@ impl CoreLoop { }); match response_codec::encode_json_as_msgpack(&result) { Ok(payload) => self.response_with_payload(task, payload), - Err(e) => self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ), + Err(e) => self.response_error(task, ErrorCode::from(e)), } } } diff --git a/nodedb/src/data/executor/handlers/document/index_fetch.rs b/nodedb/src/data/executor/handlers/document/index_fetch.rs index e9af74832..4017fed67 100644 --- a/nodedb/src/data/executor/handlers/document/index_fetch.rs +++ b/nodedb/src/data/executor/handlers/document/index_fetch.rs @@ -122,12 +122,7 @@ impl CoreLoop { ), } } - Err(e) => self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ), + Err(e) => self.response_error(task, ErrorCode::from(e)), } } @@ -181,12 +176,7 @@ impl CoreLoop { match doc_engine.index_lookup(collection, path, value, bitemporal) { Ok(ids) => ids, Err(e) => { - return self.response_error( - task, - ErrorCode::Internal { - detail: format!("indexed fetch: {e}"), - }, - ); + return self.response_error(task, ErrorCode::from(e)); } }; @@ -264,6 +254,7 @@ impl CoreLoop { } } + let identity_column = self.identity_column(database_id, tid, collection); let mut rows: Vec<(String, Vec)> = Vec::new(); for doc_id in doc_ids.iter().skip(offset).take(limit) { // Bitemporal collections keep the current body on the versioned @@ -295,6 +286,7 @@ impl CoreLoop { &residual, doc_id, &bytes, + &identity_column, ) } { Ok(true) => {} @@ -320,24 +312,14 @@ impl CoreLoop { // fail. A future compaction will purge the orphan. } Err(e) => { - return self.response_error( - task, - ErrorCode::Internal { - detail: format!("fetch doc {doc_id}: {e}"), - }, - ); + return self.response_error(task, ErrorCode::from(e)); } } } match super::super::super::response_codec::encode_raw_document_rows(&rows) { Ok(bytes) => self.response_with_payload(task, bytes), - Err(e) => self.response_error( - task, - ErrorCode::Internal { - detail: format!("indexed fetch encode: {e}"), - }, - ), + Err(e) => self.response_error(task, ErrorCode::from(e)), } } } diff --git a/nodedb/src/data/executor/handlers/document/index_maintenance.rs b/nodedb/src/data/executor/handlers/document/index_maintenance.rs index 3eb549e04..1b1c37db5 100644 --- a/nodedb/src/data/executor/handlers/document/index_maintenance.rs +++ b/nodedb/src/data/executor/handlers/document/index_maintenance.rs @@ -89,12 +89,7 @@ impl CoreLoop { ) { Ok(d) => d, Err(e) => { - return self.response_error( - task, - crate::bridge::envelope::ErrorCode::Internal { - detail: format!("backfill scan: {e}"), - }, - ); + return self.response_error(task, crate::bridge::envelope::ErrorCode::from(e)); } }; @@ -121,12 +116,7 @@ impl CoreLoop { let txn = match self.sparse.begin_write() { Ok(t) => t, Err(e) => { - return self.response_error( - task, - crate::bridge::envelope::ErrorCode::Internal { - detail: format!("backfill txn: {e}"), - }, - ); + return self.response_error(task, crate::bridge::envelope::ErrorCode::from(e)); } }; @@ -194,12 +184,7 @@ impl CoreLoop { // it writes. Buffering the keys and writing them in order turns the // whole backfill into a single forward walk. if let Err(e) = self.sparse.index_put_sorted_in_txn(&txn, &mut pending_keys) { - return self.response_error( - task, - crate::bridge::envelope::ErrorCode::Internal { - detail: format!("backfill index_put: {e}"), - }, - ); + return self.response_error(task, crate::bridge::envelope::ErrorCode::from(e)); } if let Err(e) = txn.commit() { @@ -241,20 +226,12 @@ impl CoreLoop { Ok(removed) => { match super::super::super::response_codec::encode_count("removed", removed) { Ok(bytes) => self.response_with_payload(task, bytes), - Err(e) => self.response_error( - task, - crate::bridge::envelope::ErrorCode::Internal { - detail: format!("drop index encode: {e}"), - }, - ), + Err(e) => { + self.response_error(task, crate::bridge::envelope::ErrorCode::from(e)) + } } } - Err(e) => self.response_error( - task, - crate::bridge::envelope::ErrorCode::Internal { - detail: e.to_string(), - }, - ), + Err(e) => self.response_error(task, crate::bridge::envelope::ErrorCode::from(e)), } } } diff --git a/nodedb/src/data/executor/handlers/document/read/emit.rs b/nodedb/src/data/executor/handlers/document/read/emit.rs index f4b182d52..387546499 100644 --- a/nodedb/src/data/executor/handlers/document/read/emit.rs +++ b/nodedb/src/data/executor/handlers/document/read/emit.rs @@ -2,10 +2,8 @@ //! Response emission helpers for document scans. //! -//! Two emission shapes are kept distinct: transformed rows (decoded → projected -//! → re-encoded via response_codec::encode) and raw rows (msgpack passthrough -//! via encode_raw_document_rows). Both honour the chunked-streaming contract -//! when row count exceeds `stream_chunk_size`. +//! Rows go out as msgpack passthrough via `encode_raw_document_rows`, in +//! chunks when the row count exceeds `stream_chunk_size`. //! //! A chunk boundary is a deadline safe point. A statement that goes over //! mid-stream returns a terminal `DeadlineExceeded` frame instead of its last @@ -16,32 +14,10 @@ use crate::bridge::dispatch::BridgeResponse; use crate::bridge::envelope::{ErrorCode, Response}; use crate::data::executor::core_loop::CoreLoop; -use crate::data::executor::response_codec::{self, DocumentRow}; +use crate::data::executor::response_codec; use crate::data::executor::task::ExecutionTask; impl CoreLoop { - /// Send transformed document rows (decoded → projected → re-encoded). - pub(in crate::data::executor) fn send_document_rows_transformed( - &mut self, - task: &ExecutionTask, - result: &Vec, - chunk_size: usize, - ) -> Response { - if result.len() <= chunk_size { - match response_codec::encode(result) { - Ok(payload) => self.response_with_payload(task, payload), - Err(e) => self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ), - } - } else { - self.stream_chunks_transformed(task, result, chunk_size) - } - } - /// Send raw document rows with msgpack passthrough (no decode+re-encode). pub(in crate::data::executor) fn send_document_rows_raw( &mut self, @@ -52,63 +28,13 @@ impl CoreLoop { if rows.len() <= chunk_size { match response_codec::encode_raw_document_rows(rows) { Ok(payload) => self.response_with_payload(task, payload), - Err(e) => self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ), + Err(e) => self.response_error(task, ErrorCode::from(e)), } } else { self.stream_chunks_raw(task, rows, chunk_size) } } - /// Stream transformed document rows in chunks. - fn stream_chunks_transformed( - &mut self, - task: &ExecutionTask, - result: &[DocumentRow], - chunk_size: usize, - ) -> Response { - let deadline = crate::data::executor::deadline::DeadlineCheck::for_task(task); - let chunks: Vec<_> = result.chunks(chunk_size).collect(); - let last_idx = chunks.len().saturating_sub(1); - for (i, chunk) in chunks.iter().enumerate() { - // Safe point: a chunk boundary. Every partial already pushed stays - // on the ring, but the terminal frame this returns carries - // `DeadlineExceeded`, and the Control Plane surfaces the error - // rather than the rows it collected. See the module docs. - if deadline.expired_now() { - return self.response_error(task, ErrorCode::DeadlineExceeded); - } - let is_last = i == last_idx; - match response_codec::encode(&chunk.to_vec()) { - Ok(payload) => { - if is_last { - return self.response_with_payload(task, payload); - } - let partial = self.response_partial(task, payload); - let _ = self.response_tx.try_push(BridgeResponse { inner: partial }); - } - Err(e) => { - return self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ); - } - } - } - self.response_error( - task, - ErrorCode::Internal { - detail: "streaming response incomplete".into(), - }, - ) - } - /// Stream raw document rows in chunks with msgpack passthrough. fn stream_chunks_raw( &mut self, @@ -120,7 +46,10 @@ impl CoreLoop { let chunks: Vec<_> = rows.chunks(chunk_size).collect(); let last_idx = chunks.len().saturating_sub(1); for (i, chunk) in chunks.iter().enumerate() { - // Safe point: a chunk boundary. See `stream_chunks_transformed`. + // Safe point: a chunk boundary. Every partial already pushed stays + // on the ring, but the terminal frame this returns carries + // `DeadlineExceeded`, and the Control Plane surfaces the error + // rather than the rows it collected. See the module docs. if deadline.expired_now() { return self.response_error(task, ErrorCode::DeadlineExceeded); } @@ -134,12 +63,7 @@ impl CoreLoop { let _ = self.response_tx.try_push(BridgeResponse { inner: partial }); } Err(e) => { - return self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ); + return self.response_error(task, ErrorCode::from(e)); } } } diff --git a/nodedb/src/data/executor/handlers/facet.rs b/nodedb/src/data/executor/handlers/facet.rs index c5088eea0..5c104bc24 100644 --- a/nodedb/src/data/executor/handlers/facet.rs +++ b/nodedb/src/data/executor/handlers/facet.rs @@ -119,12 +119,7 @@ impl CoreLoop { facet_result, )) { Ok(payload) => self.response_with_payload(task, payload), - Err(e) => self.response_error( - task, - crate::bridge::envelope::ErrorCode::Internal { - detail: e.to_string(), - }, - ), + Err(e) => self.response_error(task, crate::bridge::envelope::ErrorCode::from(e)), } } @@ -176,26 +171,27 @@ impl CoreLoop { ); if let Some((start, end)) = nodedb_query::msgpack_scan::extract_field(&mp, 0, field) { - let value_str = - if let Some(s) = nodedb_query::msgpack_scan::read_str(&mp, start) { - s.to_string() - } else if let Some(i) = nodedb_query::msgpack_scan::read_i64(&mp, start) { - i.to_string() - } else if let Some(f) = nodedb_query::msgpack_scan::read_f64(&mp, start) { - f.to_string() - } else if let Some(b) = nodedb_query::msgpack_scan::read_bool(&mp, start) { - b.to_string() - } else if nodedb_query::msgpack_scan::read_null(&mp, start) { - continue; - } else { - // Complex value — stringify via transcoder. - nodedb_types::msgpack_to_json_string(&mp[start..end]).map_err(|e| { - crate::Error::Serialization { - format: "msgpack".to_string(), - detail: format!("facet value of '{field}' in {key}: {e}"), - } - })? - }; + let value_str = if let Some(s) = + nodedb_query::msgpack_scan::read_str(&mp, start) + { + s.to_string() + } else if let Some(i) = nodedb_query::msgpack_scan::read_integer(&mp, start) { + i.to_string() + } else if let Some(f) = nodedb_query::msgpack_scan::read_f64(&mp, start) { + f.to_string() + } else if let Some(b) = nodedb_query::msgpack_scan::read_bool(&mp, start) { + b.to_string() + } else if nodedb_query::msgpack_scan::read_null(&mp, start) { + continue; + } else { + // Complex value — stringify via transcoder. + nodedb_types::msgpack_to_json_string(&mp[start..end]).map_err(|e| { + crate::Error::Serialization { + format: "msgpack".to_string(), + detail: format!("facet value of '{field}' in {key}: {e}"), + } + })? + }; *counts.entry(value_str).or_default() += 1; } } diff --git a/nodedb/src/data/executor/handlers/graph_algo_edges.rs b/nodedb/src/data/executor/handlers/graph_algo_edges.rs index acfc14e7e..28da59f83 100644 --- a/nodedb/src/data/executor/handlers/graph_algo_edges.rs +++ b/nodedb/src/data/executor/handlers/graph_algo_edges.rs @@ -56,27 +56,16 @@ pub(super) fn csr_from_edges( ) -> crate::Result { let mut csr = CsrIndex::new(memory); for edge in edges { - csr.add_node(&edge.src) - .map_err(|e| crate::Error::Internal { - detail: format!("algorithm CSR add src: {e}"), - })?; - csr.add_node(&edge.dst) - .map_err(|e| crate::Error::Internal { - detail: format!("algorithm CSR add dst: {e}"), - })?; + csr.add_node(&edge.src)?; + csr.add_node(&edge.dst)?; } for edge in edges { - let res = if edge.weight != 1.0 { - csr.add_edge_weighted(&edge.src, &edge.label, &edge.dst, edge.weight) + if edge.weight != 1.0 { + csr.add_edge_weighted(&edge.src, &edge.label, &edge.dst, edge.weight)?; } else { - csr.add_edge(&edge.src, &edge.label, &edge.dst) - }; - res.map_err(|e| crate::Error::Internal { - detail: format!("algorithm CSR add edge: {e}"), - })?; + csr.add_edge(&edge.src, &edge.label, &edge.dst)?; + } } - csr.compact().map_err(|e| crate::Error::Internal { - detail: format!("algorithm CSR compact: {e}"), - })?; + csr.compact()?; Ok(csr) } diff --git a/nodedb/src/data/executor/handlers/graph_temporal.rs b/nodedb/src/data/executor/handlers/graph_temporal.rs index b1046864c..3637a91e6 100644 --- a/nodedb/src/data/executor/handlers/graph_temporal.rs +++ b/nodedb/src/data/executor/handlers/graph_temporal.rs @@ -79,7 +79,7 @@ impl CoreLoop { self.graph_txn_overlays.get(&txn_id), &(task.request.database_id, tenant, collection.to_string()), node_id, - edge_label.as_deref(), + edge_label.as_deref().as_slice(), direction, edges, ) @@ -103,12 +103,7 @@ impl CoreLoop { .collect(); match super::super::response_codec::encode(&entries) { Ok(payload) => self.response_with_payload(task, payload), - Err(e) => self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ), + Err(e) => self.response_error(task, ErrorCode::from(e)), } } diff --git a/nodedb/src/data/executor/handlers/join/grace_drive.rs b/nodedb/src/data/executor/handlers/join/grace_drive.rs index 537481f4b..27c8fcf43 100644 --- a/nodedb/src/data/executor/handlers/join/grace_drive.rs +++ b/nodedb/src/data/executor/handlers/join/grace_drive.rs @@ -266,12 +266,7 @@ impl CoreLoop { return self.response_error(join.task, ErrorCode::ResourcesExhausted); } Err(e) => { - return self.response_error( - join.task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ); + return self.response_error(join.task, ErrorCode::from(e)); } }; @@ -283,12 +278,7 @@ impl CoreLoop { } if let Err(e) = join.filter_and_project(&mut results) { - return self.response_error( - join.task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ); + return self.response_error(join.task, ErrorCode::from(e)); } let payload = crate::data::executor::response_codec::encode_binary_rows(&results); diff --git a/nodedb/src/data/executor/handlers/join/hash_handlers.rs b/nodedb/src/data/executor/handlers/join/hash_handlers.rs index bfe2df6c6..a9d7d75e8 100644 --- a/nodedb/src/data/executor/handlers/join/hash_handlers.rs +++ b/nodedb/src/data/executor/handlers/join/hash_handlers.rs @@ -134,16 +134,30 @@ impl CoreLoop { // Evaluate bitmap sub-plans first. These prefilter the local scan for // each side, pushing surrogate exclusion into the document engine before // any msgpack decode occurs. - let left_bm = left_bitmap.map(|sub_plan| { - crate::data::executor::dispatch::bitmap::hashjoin_inline::run_bitmap_subplan( - self, join.task, sub_plan, - ) - }); - let right_bm = right_bitmap.map(|sub_plan| { - crate::data::executor::dispatch::bitmap::hashjoin_inline::run_bitmap_subplan( - self, join.task, sub_plan, - ) - }); + // A failing sub-plan fails the join: an empty bitmap would admit no + // probe row and silently answer with a zero-row join. + let left_bm = match left_bitmap + .map(|sub_plan| { + crate::data::executor::dispatch::bitmap::hashjoin_inline::run_bitmap_subplan( + self, join.task, sub_plan, + ) + }) + .transpose() + { + Ok(bm) => bm, + Err(e) => return self.response_error(join.task, e), + }; + let right_bm = match right_bitmap + .map(|sub_plan| { + crate::data::executor::dispatch::bitmap::hashjoin_inline::run_bitmap_subplan( + self, join.task, sub_plan, + ) + }) + .transpose() + { + Ok(bm) => bm, + Err(e) => return self.response_error(join.task, e), + }; // Memory-bounded completion path. Only when BOTH sides are plain local // scans can we stream them. For every both-local, NON-CROSS join this @@ -231,53 +245,28 @@ impl CoreLoop { (docs, resolved) } else if let Some(bm) = left_bm { - let docs = match crate::data::executor::dispatch::bitmap::hashjoin_inline::prefiltered_scan_plan( - left_collection, - scan_limit, - bm, - ) { - Some(scan_plan) => { - let resp = self.execute_plan(join.task, &scan_plan); - // Forward a failing sub-plan response (e.g. ResourcesExhausted - // from the bitmap scan) instead of swallowing it to an empty - // Vec, which would silently return a zero-row join. - // The prefiltered scan carries no predicate slot of its own, - // so both of this side's filter sets apply to its rows here. - let rows = - match crate::data::executor::response_codec::decode_response_to_docs(&resp) { - Some(d) => d, - None => return resp, - }; - match self.retain_join_side_rows(rows, left_rls_filters, left_scan_filters) { - Ok(d) => d, - Err(e) => { - return self.response_error( - join.task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ); - } - } + // An empty bitmap admits no row: the scan returns none. + let scan_plan = + crate::data::executor::dispatch::bitmap::hashjoin_inline::prefiltered_scan_plan( + left_collection, + scan_limit, + bm, + ); + let resp = self.execute_plan(join.task, &scan_plan); + // Forward a failing sub-plan response (e.g. ResourcesExhausted + // from the bitmap scan) instead of swallowing it to an empty + // Vec, which would silently return a zero-row join. + // The prefiltered scan carries no predicate slot of its own, + // so both of this side's filter sets apply to its rows here. + let rows = match crate::data::executor::response_codec::decode_response_to_docs(&resp) { + Some(d) => d, + None => return resp, + }; + let docs = match self.retain_join_side_rows(rows, left_rls_filters, left_scan_filters) { + Ok(d) => d, + Err(e) => { + return self.response_error(join.task, ErrorCode::from(e)); } - None => match self.scan_join_side(JoinSideScan { - database_id: join.task.request.database_id.as_u64(), - tenant_id: tid, - collection: left_collection, - limit: scan_limit, - rls_filters: left_rls_filters, - scan_filters: left_scan_filters, - }) { - Ok(d) => d, - Err(e) => { - return self.response_error( - join.task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ); - } - }, }; let keys = join.on.iter().map(|(l, _)| l.clone()).collect(); (docs, keys) @@ -292,12 +281,7 @@ impl CoreLoop { }) { Ok(d) => d, Err(e) => { - return self.response_error( - join.task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ); + return self.response_error(join.task, ErrorCode::from(e)); } }; let keys = join.on.iter().map(|(l, _)| l.clone()).collect(); @@ -323,54 +307,28 @@ impl CoreLoop { None => return sub_response, } } else if let Some(bm) = right_bm { - match crate::data::executor::dispatch::bitmap::hashjoin_inline::prefiltered_scan_plan( - right_collection, - scan_limit, - bm, - ) { - Some(scan_plan) => { - let resp = self.execute_plan(join.task, &scan_plan); - // Forward a failing sub-plan response (e.g. ResourcesExhausted - // from the bitmap scan) instead of swallowing it to an empty - // Vec, which would silently return a zero-row join. - // The prefiltered scan carries no predicate slot of its own, - // so both of this side's filter sets apply to its rows here. - let rows = - match crate::data::executor::response_codec::decode_response_to_docs(&resp) - { - Some(d) => d, - None => return resp, - }; - match self.retain_join_side_rows(rows, right_rls_filters, right_scan_filters) { - Ok(d) => d, - Err(e) => { - return self.response_error( - join.task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ); - } - } + // An empty bitmap admits no row: the scan returns none. + let scan_plan = + crate::data::executor::dispatch::bitmap::hashjoin_inline::prefiltered_scan_plan( + right_collection, + scan_limit, + bm, + ); + let resp = self.execute_plan(join.task, &scan_plan); + // Forward a failing sub-plan response (e.g. ResourcesExhausted + // from the bitmap scan) instead of swallowing it to an empty + // Vec, which would silently return a zero-row join. + // The prefiltered scan carries no predicate slot of its own, + // so both of this side's filter sets apply to its rows here. + let rows = match crate::data::executor::response_codec::decode_response_to_docs(&resp) { + Some(d) => d, + None => return resp, + }; + match self.retain_join_side_rows(rows, right_rls_filters, right_scan_filters) { + Ok(d) => d, + Err(e) => { + return self.response_error(join.task, ErrorCode::from(e)); } - None => match self.scan_join_side(JoinSideScan { - database_id: join.task.request.database_id.as_u64(), - tenant_id: tid, - collection: right_collection, - limit: scan_limit, - rls_filters: right_rls_filters, - scan_filters: right_scan_filters, - }) { - Ok(d) => d, - Err(e) => { - return self.response_error( - join.task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ); - } - }, } } else { match self.scan_join_side(JoinSideScan { @@ -383,12 +341,7 @@ impl CoreLoop { }) { Ok(d) => d, Err(e) => { - return self.response_error( - join.task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ); + return self.response_error(join.task, ErrorCode::from(e)); } } }; @@ -475,12 +428,7 @@ impl CoreLoop { } if let Err(e) = join.filter_and_project(&mut results) { - return self.response_error( - join.task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ); + return self.response_error(join.task, ErrorCode::from(e)); } // Deferred user LIMIT: when post-join WHERE filters exist the probe diff --git a/nodedb/src/data/executor/handlers/join/nested_loop.rs b/nodedb/src/data/executor/handlers/join/nested_loop.rs index a48a9ebe9..53c00e17d 100644 --- a/nodedb/src/data/executor/handlers/join/nested_loop.rs +++ b/nodedb/src/data/executor/handlers/join/nested_loop.rs @@ -57,12 +57,7 @@ impl CoreLoop { ) { Ok(d) => d, Err(e) => { - return self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ); + return self.response_error(task, ErrorCode::from(e)); } }; @@ -83,12 +78,7 @@ impl CoreLoop { ) { Ok(d) => d, Err(e) => { - return self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ); + return self.response_error(task, ErrorCode::from(e)); } }; diff --git a/nodedb/src/data/executor/handlers/join/sort_merge.rs b/nodedb/src/data/executor/handlers/join/sort_merge.rs index 286de26f6..4c969ac13 100644 --- a/nodedb/src/data/executor/handlers/join/sort_merge.rs +++ b/nodedb/src/data/executor/handlers/join/sort_merge.rs @@ -80,12 +80,7 @@ impl CoreLoop { ) { Ok(d) => d, Err(e) => { - return self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ); + return self.response_error(task, ErrorCode::from(e)); } }; @@ -106,12 +101,7 @@ impl CoreLoop { ) { Ok(d) => d, Err(e) => { - return self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ); + return self.response_error(task, ErrorCode::from(e)); } }; diff --git a/nodedb/src/data/executor/handlers/kv/atomic.rs b/nodedb/src/data/executor/handlers/kv/atomic.rs index 30acf60be..40c32d3fc 100644 --- a/nodedb/src/data/executor/handlers/kv/atomic.rs +++ b/nodedb/src/data/executor/handlers/kv/atomic.rs @@ -6,6 +6,7 @@ use nodedb_physical::physical_plan::KvCounterShape; use nodedb_query::msgpack_scan::{KvBodyShape, kv_body_shape}; use tracing::debug; +use super::declared_body::fit_and_admit_kv_image; use crate::bridge::envelope::{ErrorCode, Response}; use crate::data::executor::core_loop::CoreLoop; use crate::data::executor::response_codec; @@ -45,7 +46,7 @@ pub(in crate::data::executor) fn atomic_error_code( AtomicError::Encode { detail } => ErrorCode::Internal { detail }, // Nothing was written: the engine consults the gate before it // installs the computed value. - AtomicError::Rejected(error) => (*error).into(), + AtomicError::Declared(error) | AtomicError::Rejected(error) => (*error).into(), AtomicError::Unbound(error) => error.into(), } } @@ -91,10 +92,13 @@ impl CoreLoop { // see `CoreLoop::kv_ttl_now_ms` for the precedence this resolves. let now_ms: u64 = self.kv_ttl_now_ms(task); // The engine computes the post-image and installs it in one pass, so - // the write policy is handed in and decided on the computed bytes - // rather than on a duplicate of the increment arithmetic out here. - let admit = - |image: &[u8]| super::rls::admit_kv_row(rls_write_check, image, key, tid, collection); + // the declared column rule and the write policy are handed in and + // applied to the computed bytes rather than to a duplicate of the + // increment arithmetic out here. + let declared = self.declared_columns_of(did, tid, collection).to_vec(); + let admit = |image: &[u8]| { + fit_and_admit_kv_image(image, &declared, rls_write_check, key, tid, collection) + }; match self.kv_engine.incr( crate::engine::kv::AtomicKeyCtx { database_id: did, @@ -127,12 +131,7 @@ impl CoreLoop { match response_codec::encode_json_as_msgpack(&serde_json::json!({ "value": value })) { Ok(payload) => self.response_with_payload(task, payload), - Err(e) => self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ), + Err(e) => self.response_error(task, ErrorCode::from(e)), } } Err(error) => self.response_atomic_error(task, collection, error), @@ -165,8 +164,10 @@ impl CoreLoop { .map(|ms| ms as u64) .unwrap_or_else(current_ms); // Same engine-internal compute-and-persist as `Incr` — see there. - let admit = - |image: &[u8]| super::rls::admit_kv_row(rls_write_check, image, key, tid, collection); + let declared = self.declared_columns_of(did, tid, collection).to_vec(); + let admit = |image: &[u8]| { + fit_and_admit_kv_image(image, &declared, rls_write_check, key, tid, collection) + }; match self.kv_engine.incr_float( crate::engine::kv::AtomicKeyCtx { database_id: did, @@ -197,12 +198,7 @@ impl CoreLoop { self.note_kv_write_lsn(task, did, tid, collection, key); match response_codec::encode_json_as_msgpack(&incr_float_reply(value, &written)) { Ok(payload) => self.response_with_payload(task, payload), - Err(e) => self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ), + Err(e) => self.response_error(task, ErrorCode::from(e)), } } Err(error) => self.response_atomic_error(task, collection, error), @@ -235,10 +231,12 @@ impl CoreLoop { .map(|ms| ms as u64) .unwrap_or_else(current_ms); // A swap into a typed row stores the row with one column replaced, not - // `new_value` itself, so the policy decides the image the engine - // computes — see `Incr`. - let admit = - |image: &[u8]| super::rls::admit_kv_row(rls_write_check, image, key, tid, collection); + // `new_value` itself, so the declared rule and the policy apply to the + // image the engine computes — see `Incr`. + let declared = self.declared_columns_of(did, tid, collection).to_vec(); + let admit = |image: &[u8]| { + fit_and_admit_kv_image(image, &declared, rls_write_check, key, tid, collection) + }; let result = match self.kv_engine.cas( crate::engine::kv::AtomicKeyCtx { database_id: did, @@ -280,12 +278,7 @@ impl CoreLoop { "current_value": current_b64, })) { Ok(payload) => self.response_with_payload(task, payload), - Err(e) => self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ), + Err(e) => self.response_error(task, ErrorCode::from(e)), } } @@ -319,9 +312,12 @@ impl CoreLoop { .map(|ms| ms as u64) .unwrap_or_else(current_ms); // A write into a typed row stores the row with one column replaced, so - // the policy decides the image the engine computes — see `Incr`. - let admit = - |image: &[u8]| super::rls::admit_kv_row(rls_write_check, image, key, tid, collection); + // the declared rule and the policy apply to the image the engine + // computes — see `Incr`. + let declared = self.declared_columns_of(did, tid, collection).to_vec(); + let admit = |image: &[u8]| { + fit_and_admit_kv_image(image, &declared, rls_write_check, key, tid, collection) + }; let crate::engine::kv::GetSetResult { old, written } = match self.kv_engine.getset( crate::engine::kv::AtomicKeyCtx { database_id: did, @@ -361,12 +357,7 @@ impl CoreLoop { Ok(true) => old.as_deref(), Ok(false) => None, Err(e) => { - return self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ); + return self.response_error(task, ErrorCode::from(e)); } }, None => None, @@ -376,12 +367,7 @@ impl CoreLoop { .map(|v| base64::Engine::encode(&base64::engine::general_purpose::STANDARD, v)); match response_codec::encode_json_as_msgpack(&serde_json::json!({ "old_value": old_b64 })) { Ok(payload) => self.response_with_payload(task, payload), - Err(e) => self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ), + Err(e) => self.response_error(task, ErrorCode::from(e)), } } @@ -612,6 +598,126 @@ mod tests { assert_eq!(stored(&h.core, b"hits"), b"42".to_vec()); } + /// Register `declared` as the declared numeric columns of the collection. + fn declare(core: &mut CoreLoop, declared: &[(&str, &str)]) { + let mut config = crate::engine::document::store::CollectionConfig::new(COLLECTION); + config.declared_columns = declared + .iter() + .map(|(name, declared)| { + nodedb_physical::physical_plan::DeclaredColumn::from_declared(name, declared) + .expect("numeric declaration") + }) + .collect(); + core.doc_configs.insert( + ( + DatabaseId::DEFAULT, + TenantId::new(TID), + COLLECTION.to_string(), + ), + config, + ); + } + + fn is_out_of_range(resp: &crate::bridge::envelope::Response) -> bool { + matches!( + resp.error_code.as_deref(), + Some(crate::bridge::envelope::ErrorCode::NumericValueOutOfRange { .. }) + ) + } + + #[test] + fn incr_past_a_declared_smallint_is_refused_and_writes_nothing() { + let mut h = make_core(); + let row = typed_row(&[ + ("label", Value::String("gold".into())), + ("n", Value::Integer(32767)), + ]); + seed(&mut h.core, b"player", &row); + declare(&mut h.core, &[("n", "SMALLINT")]); + + let t = task(); + let check = RlsWriteCheck::already_decided_elsewhere(); + let resp = h + .core + .execute_kv_incr(ctx(&t, b"player", &check), 1, 0, &KvCounterShape::Raw); + assert!(is_out_of_range(&resp), "{:?}", resp.error_code); + assert!( + h.events.try_recv().is_none(), + "a refused INCR emits no event" + ); + assert_eq!(stored(&h.core, b"player"), row); + + let resp = h + .core + .execute_kv_incr(ctx(&t, b"player", &check), -1, 0, &KvCounterShape::Raw); + assert_eq!(resp.status, Status::Ok, "{:?}", resp.error_code); + assert_eq!( + columns(&stored(&h.core, b"player")).get("n"), + Some(&Value::Integer(32766)) + ); + } + + #[test] + fn incr_float_meets_a_declared_decimal() { + let mut h = make_core(); + declare(&mut h.core, &[("value", "DECIMAL(5,2)")]); + let t = task(); + let check = RlsWriteCheck::already_decided_elsewhere(); + + seed(&mut h.core, b"price", b"999.99"); + let resp = + h.core + .execute_kv_incr_float(ctx(&t, b"price", &check), "1", &KvCounterShape::Raw); + assert!(is_out_of_range(&resp), "{:?}", resp.error_code); + assert_eq!(stored(&h.core, b"price"), b"999.99".to_vec()); + + seed(&mut h.core, b"fee", b"1.50"); + let resp = + h.core + .execute_kv_incr_float(ctx(&t, b"fee", &check), "0.005", &KvCounterShape::Raw); + assert_eq!(resp.status, Status::Ok, "{:?}", resp.error_code); + assert_eq!( + stored(&h.core, b"fee"), + b"1.51".to_vec(), + "the stored value is rounded to the declared scale" + ); + let reply: serde_json::Value = + nodedb_types::json_from_msgpack(resp.payload.as_bytes()).expect("decode reply"); + assert_eq!(reply["value"], serde_json::json!(1.51)); + assert_eq!(reply["text"], serde_json::json!("1.51")); + } + + #[test] + fn cas_and_getset_past_a_declared_smallint_are_refused() { + let mut h = make_core(); + declare(&mut h.core, &[("value", "SMALLINT")]); + seed(&mut h.core, b"slot", b"5"); + let t = task(); + let check = RlsWriteCheck::already_decided_elsewhere(); + + let resp = h + .core + .execute_kv_cas(ctx(&t, b"slot", &check), b"5", b"40000"); + assert!(is_out_of_range(&resp), "{:?}", resp.error_code); + assert_eq!(stored(&h.core, b"slot"), b"5".to_vec()); + + let resp = h + .core + .execute_kv_getset(ctx(&t, b"slot", &check), b"40000", &[]); + assert!(is_out_of_range(&resp), "{:?}", resp.error_code); + assert_eq!(stored(&h.core, b"slot"), b"5".to_vec()); + assert!( + h.events.try_recv().is_none(), + "a refused write emits no event" + ); + + let resp = h + .core + .execute_kv_getset(ctx(&t, b"slot", &check), b"7", &[]); + assert_eq!(resp.status, Status::Ok, "{:?}", resp.error_code); + assert_eq!(stored(&h.core, b"slot"), b"7".to_vec()); + } + #[test] fn incr_on_raw_text_that_is_not_an_integer_answers_the_counter_fault() { let mut h = make_core(); diff --git a/nodedb/src/data/executor/handlers/kv/batch.rs b/nodedb/src/data/executor/handlers/kv/batch.rs index 043cfb0fd..18fd6fe2c 100644 --- a/nodedb/src/data/executor/handlers/kv/batch.rs +++ b/nodedb/src/data/executor/handlers/kv/batch.rs @@ -69,12 +69,7 @@ impl CoreLoop { )), Ok(false) => serde_json::Value::Null, Err(e) => { - return self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ); + return self.response_error(task, ErrorCode::from(e)); } }, None => serde_json::Value::Null, @@ -83,12 +78,7 @@ impl CoreLoop { } match response_codec::encode_json_vec_as_msgpack(&json_results) { Ok(payload) => self.response_with_payload(task, payload), - Err(e) => self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ), + Err(e) => self.response_error(task, ErrorCode::from(e)), } } @@ -170,12 +160,7 @@ impl CoreLoop { } match response_codec::encode_count("inserted", new_count) { Ok(payload) => self.response_with_payload(task, payload), - Err(e) => self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ), + Err(e) => self.response_error(task, ErrorCode::from(e)), } } } diff --git a/nodedb/src/data/executor/handlers/kv/crud/delete.rs b/nodedb/src/data/executor/handlers/kv/crud/delete.rs index 7691a8611..de25fd2e9 100644 --- a/nodedb/src/data/executor/handlers/kv/crud/delete.rs +++ b/nodedb/src/data/executor/handlers/kv/crud/delete.rs @@ -94,12 +94,7 @@ impl CoreLoop { match response_codec::encode_count("deleted", count) { Ok(payload) => self.response_with_payload(task, payload), - Err(e) => self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ), + Err(e) => self.response_error(task, ErrorCode::from(e)), } } @@ -114,12 +109,7 @@ impl CoreLoop { let count = self.kv_engine.truncate(did, tid, collection); match response_codec::encode_count("deleted", count) { Ok(payload) => self.response_with_payload(task, payload), - Err(e) => self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ), + Err(e) => self.response_error(task, ErrorCode::from(e)), } } } diff --git a/nodedb/src/data/executor/handlers/kv/field.rs b/nodedb/src/data/executor/handlers/kv/field.rs index 6bae04703..d5a68451e 100644 --- a/nodedb/src/data/executor/handlers/kv/field.rs +++ b/nodedb/src/data/executor/handlers/kv/field.rs @@ -72,12 +72,7 @@ impl CoreLoop { Ok(true) => {} Ok(false) => return self.response_error(task, ErrorCode::NotFound), Err(e) => { - return self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ); + return self.response_error(task, ErrorCode::from(e)); } } @@ -114,12 +109,7 @@ impl CoreLoop { match response_codec::encode_json_as_msgpack(&serde_json::Value::Object(result)) { Ok(payload) => self.response_with_payload(task, payload), - Err(e) => self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ), + Err(e) => self.response_error(task, ErrorCode::from(e)), } } @@ -160,12 +150,7 @@ impl CoreLoop { &serde_json::json!({ "affected": 0, "fields_added": 0 }), ) { Ok(payload) => self.response_with_payload(task, payload), - Err(e) => self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ), + Err(e) => self.response_error(task, ErrorCode::from(e)), }; } @@ -176,6 +161,7 @@ impl CoreLoop { collection, current.as_deref(), updates, + self.declared_columns_of(did, tid, collection), ) { Ok(c) => c, Err(e) => return self.response_error(task, e), @@ -233,12 +219,7 @@ impl CoreLoop { &serde_json::json!({ "affected": 1, "fields_added": computed.fields_added }), ) { Ok(payload) => self.response_with_payload(task, payload), - Err(e) => self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ), + Err(e) => self.response_error(task, ErrorCode::from(e)), } } } diff --git a/nodedb/src/data/executor/handlers/kv/index.rs b/nodedb/src/data/executor/handlers/kv/index.rs index 2f775f4fe..8f09c6084 100644 --- a/nodedb/src/data/executor/handlers/kv/index.rs +++ b/nodedb/src/data/executor/handlers/kv/index.rs @@ -53,12 +53,7 @@ impl CoreLoop { "write_amp_estimate": format!("{:.0}%", 15.0 + 10.0 * self.kv_engine.index_count(did, tid, collection) as f64), })) { Ok(payload) => self.response_with_payload(task, payload), - Err(e) => self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ), + Err(e) => self.response_error(task, ErrorCode::from(e)), } } @@ -77,12 +72,7 @@ impl CoreLoop { "entries_removed": removed, })) { Ok(payload) => self.response_with_payload(task, payload), - Err(e) => self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ), + Err(e) => self.response_error(task, ErrorCode::from(e)), } } } diff --git a/nodedb/src/data/executor/handlers/kv/predicate/apply.rs b/nodedb/src/data/executor/handlers/kv/predicate/apply.rs index 6678e111d..7ca5020ec 100644 --- a/nodedb/src/data/executor/handlers/kv/predicate/apply.rs +++ b/nodedb/src/data/executor/handlers/kv/predicate/apply.rs @@ -62,12 +62,14 @@ impl CoreLoop { Err(e) => return self.response_error(task, e), }; + let declared = self.declared_columns_of(did, tid, collection); let mut writes: Vec<(Vec, Vec, Vec)> = Vec::with_capacity(matched.len()); for (key, body) in matched { - let computed = match merge_field_updates(collection, Some(body.as_slice()), updates) { - Ok(c) => c, - Err(e) => return self.response_error(task, e), - }; + let computed = + match merge_field_updates(collection, Some(body.as_slice()), updates, declared) { + Ok(c) => c, + Err(e) => return self.response_error(task, e), + }; if let Err(e) = admit_kv_row(rls_write_check, &computed.new_value, &key, tid, collection) { @@ -119,12 +121,7 @@ impl CoreLoop { match response_codec::encode_count("affected", writes.len()) { Ok(payload) => self.response_with_payload(task, payload), - Err(e) => self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ), + Err(e) => self.response_error(task, ErrorCode::from(e)), } } diff --git a/nodedb/src/data/executor/handlers/kv/ttl.rs b/nodedb/src/data/executor/handlers/kv/ttl.rs index 7029bc837..7d71934ba 100644 --- a/nodedb/src/data/executor/handlers/kv/ttl.rs +++ b/nodedb/src/data/executor/handlers/kv/ttl.rs @@ -199,12 +199,7 @@ impl CoreLoop { fn kv_get_ttl_response(&self, task: &ExecutionTask, ttl_ms: i64) -> Response { match response_codec::encode_json_as_msgpack(&serde_json::json!({ "ttl_ms": ttl_ms })) { Ok(payload) => self.response_with_payload(task, payload), - Err(e) => self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ), + Err(e) => self.response_error(task, ErrorCode::from(e)), } } } diff --git a/nodedb/src/data/executor/handlers/timeseries/ingest.rs b/nodedb/src/data/executor/handlers/timeseries/ingest.rs index 4bc90d7ae..9771b5586 100644 --- a/nodedb/src/data/executor/handlers/timeseries/ingest.rs +++ b/nodedb/src/data/executor/handlers/timeseries/ingest.rs @@ -44,11 +44,7 @@ impl CoreLoop { let mut simulation = match existing { Some(memtable) => { ColumnarMemtable::from_snapshot(memtable.export_snapshot(), memtable.config()) - .map_err(|error| ErrorCode::Internal { - detail: format!( - "failed to clone timeseries memtable for admission: {error}" - ), - })? + .map_err(ErrorCode::from)? } None => { let mut schema = self.initial_ts_schema(task, tid, collection, lines); @@ -251,12 +247,7 @@ impl CoreLoop { && let Err(e) = self.flush_ts_collection(tid, task.request.database_id, collection, now_ms) { - return self.response_error( - task, - ErrorCode::Internal { - detail: format!("pre-ingest ts flush failed: {e}"), - }, - ); + return self.response_error(task, ErrorCode::from(e)); } let emits_events = self.ts_ingest_emits_events(task, collection, mode); @@ -378,12 +369,7 @@ impl CoreLoop { && let Err(e) = self.flush_ts_collection(tid, task.request.database_id, collection, now_ms) { - return self.response_error( - task, - ErrorCode::Internal { - detail: format!("post-ingest ts flush failed: {e}"), - }, - ); + return self.response_error(task, ErrorCode::from(e)); } if accepted > 0 { @@ -434,12 +420,7 @@ impl CoreLoop { let json = match response_codec::encode_json_as_msgpack(&result) { Ok(b) => b, Err(e) => { - return self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ); + return self.response_error(task, ErrorCode::from(e)); } }; Response { diff --git a/nodedb/src/data/executor/handlers/timeseries/ingest_dispatch.rs b/nodedb/src/data/executor/handlers/timeseries/ingest_dispatch.rs index 61404eaee..46888bca3 100644 --- a/nodedb/src/data/executor/handlers/timeseries/ingest_dispatch.rs +++ b/nodedb/src/data/executor/handlers/timeseries/ingest_dispatch.rs @@ -162,12 +162,7 @@ impl CoreLoop { let json = match response_codec::encode_json_as_msgpack(&result) { Ok(b) => b, Err(e) => { - return self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ); + return self.response_error(task, ErrorCode::from(e)); } }; return Response { diff --git a/nodedb/src/data/executor/handlers/timeseries/ingest_resolved.rs b/nodedb/src/data/executor/handlers/timeseries/ingest_resolved.rs index e4494d85e..f2c387510 100644 --- a/nodedb/src/data/executor/handlers/timeseries/ingest_resolved.rs +++ b/nodedb/src/data/executor/handlers/timeseries/ingest_resolved.rs @@ -53,6 +53,7 @@ use nodedb_types::timeseries::SeriesKey; use super::admission; use super::ingest_dispatch::{TimeseriesApplyMode, TimeseriesIngestParams}; use super::ingest_resolved_fit::{CollKey, SchemaFit, landing_values}; +use super::ingest_resolved_returning::decode_returning_images; use crate::data::executor::response_codec::IngestRejection; use crate::engine::timeseries::install_counts::TsInstallCount; @@ -111,12 +112,7 @@ impl CoreLoop { self.flush_ts_collection(tid, task.request.database_id, collection, now_ms) { if refusable { - return self.response_error( - task, - ErrorCode::Internal { - detail: format!("pre-ingest ts flush failed: {e}"), - }, - ); + return self.response_error(task, ErrorCode::from(e)); } // A committed entry every other replica applies cannot make room // on this node. @@ -311,12 +307,11 @@ impl CoreLoop { ); } - let returned_rows: Vec = match returning { - Some(_) => images - .iter() - .filter_map(|image| crate::util::bounded_msgpack::read_value(image).ok()) - .collect(), - None => Vec::new(), + // An image that does not decode refuses the statement below, once the + // landed rows' bookkeeping has run. + let returned_rows = match returning { + Some(_) => decode_returning_images(collection, &images), + None => Ok(Vec::new()), }; if mode != TimeseriesApplyMode::RedoInstall { @@ -329,12 +324,7 @@ impl CoreLoop { && let Err(e) = self.flush_ts_collection(tid, task.request.database_id, collection, now_ms) { - return self.response_error( - task, - ErrorCode::Internal { - detail: format!("post-ingest ts flush failed: {e}"), - }, - ); + return self.response_error(task, ErrorCode::from(e)); } if accepted > 0 { // no-determinism: Instant::now runs only for the operational idle/checkpoint timer, outside a committed-redo install. @@ -346,6 +336,10 @@ impl CoreLoop { } if let Some(spec) = returning { + let returned_rows = match returned_rows { + Ok(rows) => rows, + Err(e) => return self.response_error(task, e), + }; // A row set has no place for a rejected row, so the rows the // install rejected travel beside it as the rejected-lines notice. let rejection = (rejected > 0).then(|| IngestRejection { @@ -392,12 +386,7 @@ impl CoreLoop { read_version_lsn: crate::types::Lsn::ZERO, write_set: Vec::new(), }, - Err(e) => self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ), + Err(e) => self.response_error(task, ErrorCode::from(e)), } } @@ -472,6 +461,7 @@ impl CoreLoop { row_id: RowId::Batch, new_value: Some(*image), old_value: None, + image_fault: None, }, ); } diff --git a/nodedb/src/data/executor/handlers/timeseries/scan.rs b/nodedb/src/data/executor/handlers/timeseries/scan.rs index ccbb8b3e1..f986f9b46 100644 --- a/nodedb/src/data/executor/handlers/timeseries/scan.rs +++ b/nodedb/src/data/executor/handlers/timeseries/scan.rs @@ -86,12 +86,7 @@ impl CoreLoop { // Lazy-load partition registry from disk if not yet loaded. if let Err(e) = self.ensure_ts_registry(tid, task.request.database_id, collection) { - return self.response_error( - task, - crate::bridge::envelope::ErrorCode::Internal { - detail: e.to_string(), - }, - ); + return self.response_error(task, crate::bridge::envelope::ErrorCode::from(e)); } // Both predicate sets are decoded before any row is read. A payload diff --git a/nodedb/src/data/executor/handlers/transaction/stage_write/stage_bulk_delete.rs b/nodedb/src/data/executor/handlers/transaction/stage_write/stage_bulk_delete.rs index 3108cdd32..7477b5910 100644 --- a/nodedb/src/data/executor/handlers/transaction/stage_write/stage_bulk_delete.rs +++ b/nodedb/src/data/executor/handlers/transaction/stage_write/stage_bulk_delete.rs @@ -183,12 +183,7 @@ impl CoreLoop { match response_codec::encode_json_as_msgpack(&serde_json::json!({ "affected": affected })) { Ok(payload) => self.response_with_payload(task, payload), - Err(e) => self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ), + Err(e) => self.response_error(task, ErrorCode::from(e)), } } } diff --git a/nodedb/src/data/executor/handlers/transaction/stage_write/stage_bulk_update.rs b/nodedb/src/data/executor/handlers/transaction/stage_write/stage_bulk_update.rs index 7c8c49829..a381b958d 100644 --- a/nodedb/src/data/executor/handlers/transaction/stage_write/stage_bulk_update.rs +++ b/nodedb/src/data/executor/handlers/transaction/stage_write/stage_bulk_update.rs @@ -227,12 +227,7 @@ impl CoreLoop { match response_codec::encode_json_as_msgpack(&serde_json::json!({ "affected": affected })) { Ok(payload) => self.response_with_payload(task, payload), - Err(e) => self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ), + Err(e) => self.response_error(task, ErrorCode::from(e)), } } diff --git a/nodedb/src/data/executor/handlers/transaction/stage_write/stage_columnar.rs b/nodedb/src/data/executor/handlers/transaction/stage_write/stage_columnar.rs index c91c25b45..140b8a67d 100644 --- a/nodedb/src/data/executor/handlers/transaction/stage_write/stage_columnar.rs +++ b/nodedb/src/data/executor/handlers/transaction/stage_write/stage_columnar.rs @@ -147,13 +147,16 @@ impl CoreLoop { // hitting the scan's "missing engine -> empty result" branch. See // `ensure_columnar_engine_schema` doc comment. let engine_preexisted = self.columnar_engines.contains_key(&engine_key); - let schema = self.ensure_columnar_engine_schema( + let schema = match self.ensure_columnar_engine_schema( &engine_key, collection, bitemporal, &ndb_rows[0], schema_bytes, - ); + ) { + Ok(schema) => schema, + Err(e) => return self.response_error(task, ErrorCode::from(e)), + }; // Track engines THIS transaction newly auto-created (never engines // that already existed before the txn started) so `MetaOp::DropTxnOverlay` // can drop the still-empty ones on rollback without touching engines @@ -208,19 +211,12 @@ impl CoreLoop { Some(Value::Integer(i)) => Value::Integer(*i), _ => Value::Integer(i64::MAX), }), - _ => ndb_field_to_value(obj.get(&col.name), &col.column_type), + _ => ndb_field_to_value(obj.get(&col.name), col), }) .collect::, crate::Error>>() { Ok(v) => v, - Err(e) => { - return self.response_error( - task, - ErrorCode::Internal { - detail: format!("columnar insert coercion: {e}"), - }, - ); - } + Err(e) => return self.response_error(task, ErrorCode::from(e)), }; let surrogate = match surrogates.get(row_idx).copied() { diff --git a/nodedb/src/data/executor/handlers/truncate_response.rs b/nodedb/src/data/executor/handlers/truncate_response.rs index 067aea478..5a09d9b65 100644 --- a/nodedb/src/data/executor/handlers/truncate_response.rs +++ b/nodedb/src/data/executor/handlers/truncate_response.rs @@ -16,12 +16,7 @@ impl CoreLoop { ) -> Response { match response_codec::encode_count("truncated", truncated) { Ok(payload) => self.response_with_payload(task, payload), - Err(e) => self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ), + Err(e) => self.response_error(task, ErrorCode::from(e)), } } } diff --git a/nodedb/src/data/executor/handlers/unregister_collection.rs b/nodedb/src/data/executor/handlers/unregister_collection.rs index 4972cfda1..17fe94202 100644 --- a/nodedb/src/data/executor/handlers/unregister_collection.rs +++ b/nodedb/src/data/executor/handlers/unregister_collection.rs @@ -99,12 +99,7 @@ impl CoreLoop { // the DROP does not finalize the catalog-row removal over storage // rows that survive — the resurrection hole on re-CREATE. Err(e) => { - return self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ); + return self.response_error(task, ErrorCode::from(e)); } }; diff --git a/nodedb/src/data/executor/handlers/update_from_join.rs b/nodedb/src/data/executor/handlers/update_from_join.rs index fa74dd17a..61e0ad405 100644 --- a/nodedb/src/data/executor/handlers/update_from_join.rs +++ b/nodedb/src/data/executor/handlers/update_from_join.rs @@ -65,12 +65,7 @@ impl CoreLoop { ) { Ok(m) => m, Err(e) => { - return self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ); + return self.response_error(task, ErrorCode::from(e)); } }; @@ -99,12 +94,7 @@ impl CoreLoop { let result = serde_json::json!({ "affected": 0u64 }); return match encode_json_as_msgpack(&result) { Ok(payload) => self.response_with_payload(task, payload), - Err(e) => self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ), + Err(e) => self.response_error(task, ErrorCode::from(e)), }; } @@ -152,12 +142,7 @@ impl CoreLoop { ) { Ok(r) => r, Err(e) => { - return self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ); + return self.response_error(task, ErrorCode::from(e)); } }; @@ -260,23 +245,13 @@ impl CoreLoop { &outcome.returned_docs, ) { Ok(payload) => self.response_with_payload(task, payload), - Err(e) => self.response_error( - task, - ErrorCode::Internal { - detail: format!("RETURNING encode: {e}"), - }, - ), + Err(e) => self.response_error(task, ErrorCode::from(e)), } } else { let result = serde_json::json!({ "affected": outcome.affected }); match encode_json_as_msgpack(&result) { Ok(payload) => self.response_with_payload(task, payload), - Err(e) => self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ), + Err(e) => self.response_error(task, ErrorCode::from(e)), } }; if !outcome.write_set.is_empty() { diff --git a/nodedb/src/data/executor/handlers/upsert/exec/dispatch.rs b/nodedb/src/data/executor/handlers/upsert/exec/dispatch.rs index 25cf1547e..da5b83053 100644 --- a/nodedb/src/data/executor/handlers/upsert/exec/dispatch.rs +++ b/nodedb/src/data/executor/handlers/upsert/exec/dispatch.rs @@ -160,12 +160,7 @@ impl CoreLoop { strict_schema: strict_schema.as_ref(), }, ), - Err(e) => self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ), + Err(e) => self.response_error(task, ErrorCode::from(e)), } } } diff --git a/nodedb/src/data/executor/handlers/vector_lifecycle.rs b/nodedb/src/data/executor/handlers/vector_lifecycle.rs index a94858853..9473e1c5f 100644 --- a/nodedb/src/data/executor/handlers/vector_lifecycle.rs +++ b/nodedb/src/data/executor/handlers/vector_lifecycle.rs @@ -137,12 +137,7 @@ impl CoreLoop { match super::super::response_codec::encode_count("compacted", removed) { Ok(bytes) => self.response_with_payload(task, bytes), - Err(e) => self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ), + Err(e) => self.response_error(task, ErrorCode::from(e)), } } diff --git a/nodedb/src/data/executor/handlers/vector_multi.rs b/nodedb/src/data/executor/handlers/vector_multi.rs index 97faaa368..e6ea15310 100644 --- a/nodedb/src/data/executor/handlers/vector_multi.rs +++ b/nodedb/src/data/executor/handlers/vector_multi.rs @@ -126,12 +126,7 @@ impl CoreLoop { match super::super::response_codec::encode_count("inserted_vectors", ids.len()) { Ok(bytes) => self.response_with_payload(task, bytes), - Err(e) => self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ), + Err(e) => self.response_error(task, ErrorCode::from(e)), } } @@ -274,12 +269,7 @@ impl CoreLoop { Ok(payload) => self.response_with_payload(task, payload), Err(e) => { warn!(core = self.core_id, error = %e, "multi-vector search encode failed"); - self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ) + self.response_error(task, ErrorCode::from(e)) } } } diff --git a/nodedb/src/data/executor/handlers/vector_sparse.rs b/nodedb/src/data/executor/handlers/vector_sparse.rs index d941a8d87..6e965c897 100644 --- a/nodedb/src/data/executor/handlers/vector_sparse.rs +++ b/nodedb/src/data/executor/handlers/vector_sparse.rs @@ -118,12 +118,7 @@ impl CoreLoop { >::new()) { Ok(payload) => self.response_with_payload(task, payload), - Err(e) => self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ), + Err(e) => self.response_error(task, ErrorCode::from(e)), }; }; @@ -170,12 +165,7 @@ impl CoreLoop { Ok(payload) => self.response_with_payload(task, payload), Err(e) => { warn!(core = self.core_id, error = %e, "sparse search encode failed"); - self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ) + self.response_error(task, ErrorCode::from(e)) } } } diff --git a/nodedb/src/data/executor/handlers/vector_write.rs b/nodedb/src/data/executor/handlers/vector_write.rs index 66d37a1c8..0b961b654 100644 --- a/nodedb/src/data/executor/handlers/vector_write.rs +++ b/nodedb/src/data/executor/handlers/vector_write.rs @@ -74,12 +74,7 @@ impl CoreLoop { } match super::super::response_codec::encode_count("inserted", vectors.len()) { Ok(bytes) => self.response_with_payload(task, bytes), - Err(e) => self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ), + Err(e) => self.response_error(task, ErrorCode::from(e)), } } Err(err) => self.response_error(task, err), diff --git a/nodedb/src/data/executor/snapshot/capture.rs b/nodedb/src/data/executor/snapshot/capture.rs index aff5dbce4..4af085ce6 100644 --- a/nodedb/src/data/executor/snapshot/capture.rs +++ b/nodedb/src/data/executor/snapshot/capture.rs @@ -62,12 +62,7 @@ impl CoreLoop { } Err(e) => { warn!(core = self.core_id, error = %e, "core snapshot capture failed"); - self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ) + self.response_error(task, ErrorCode::from(e)) } } } @@ -197,7 +192,7 @@ mod tests { tid(), "docs", Surrogate::new(7), - "quick brown fox", + &crate::engine::sparse::inverted::test_support::body("quick brown fox"), ) .unwrap(); diff --git a/nodedb/src/data/executor/strict_format/encode.rs b/nodedb/src/data/executor/strict_format/encode.rs index efb064b1d..e9da9d8c7 100644 --- a/nodedb/src/data/executor/strict_format/encode.rs +++ b/nodedb/src/data/executor/strict_format/encode.rs @@ -66,16 +66,50 @@ pub fn value_to_binary_tuple( } Value::Null } - Some(v) => coerce_value(v, &col.column_type, &col.name)?, + Some(v) => coerce_value(v, col)?, }; values.push(typed); } encoder .encode(&values) - .map_err(|e| crate::Error::BadRequest { - detail: format!("Binary Tuple encode: {e}"), - }) + .map_err(|e| tuple_encode_error(e, map)) +} + +/// The error for a row the tuple encoder refuses. +/// +/// A value of a kind the column does not hold is `DatatypeMismatch` +/// (SQLSTATE `42804`). A value of an accepted kind that does not convert, +/// such as text that is no UUID, is `InvalidTextRepresentation` (SQLSTATE +/// `22P02`). Both name the column, the value from `row`, and the type. Any +/// other encoder error is a malformed request. +fn tuple_encode_error( + error: nodedb_strict::StrictError, + row: &std::collections::HashMap, +) -> crate::Error { + match error { + nodedb_strict::StrictError::TypeMismatch { column, expected } => { + crate::Error::DatatypeMismatch { + detail: format!( + "column '{column}': expected {expected}, got {:?}", + row.get(&column).unwrap_or(&Value::Null) + ), + } + } + nodedb_strict::StrictError::InvalidValue { + column, + expected, + detail, + } => crate::Error::InvalidTextRepresentation { + detail: format!( + "column '{column}': {:?} is not a valid {expected} value: {detail}", + row.get(&column).unwrap_or(&Value::Null) + ), + }, + other => crate::Error::BadRequest { + detail: format!("Binary Tuple encode: {other}"), + }, + } } /// Bitemporal variant: decode msgpack to `Value`, then encode as a Binary @@ -170,7 +204,7 @@ pub fn value_to_binary_tuple_bitemporal( } Value::Null } - Some(v) => coerce_value(v, &col.column_type, &col.name)?, + Some(v) => coerce_value(v, col)?, }; user_values.push(typed); } @@ -178,7 +212,5 @@ pub fn value_to_binary_tuple_bitemporal( let encoder = nodedb_strict::TupleEncoder::new(schema); encoder .encode_bitemporal(system_from_ms, valid_from_ms, valid_until_ms, &user_values) - .map_err(|e| crate::Error::BadRequest { - detail: format!("Binary Tuple encode: {e}"), - }) + .map_err(|e| tuple_encode_error(e, map)) } diff --git a/nodedb/src/error/conversions.rs b/nodedb/src/error/conversions.rs index ef0238f92..972eb5261 100644 --- a/nodedb/src/error/conversions.rs +++ b/nodedb/src/error/conversions.rs @@ -87,3 +87,14 @@ impl From for Error { } } } + +/// An array-engine error. The array store is engine storage, so it renders +/// as a storage error of the `array` engine. +impl From for Error { + fn from(e: crate::engine::array::engine::ArrayEngineError) -> Self { + Error::Storage { + engine: "array".to_string(), + detail: e.to_string(), + } + } +} diff --git a/nodedb/src/error/types.rs b/nodedb/src/error/types.rs index 5340ca7db..b023ac085 100644 --- a/nodedb/src/error/types.rs +++ b/nodedb/src/error/types.rs @@ -332,6 +332,16 @@ pub enum Error { #[error("column \"{column}\" does not exist")] UndefinedColumn { column: String }, + /// The column argument of a full-text search cannot serve it. Propagated + /// from `SqlError::TextColumn` or the Data-Plane `ErrorCode::TextColumn`. + /// The pgwire layer renders `fault.sqlstate()`. + #[error("column \"{column}\" of collection \"{collection}\" {fault}")] + TextColumn { + collection: String, + column: String, + fault: nodedb_types::text_search::TextColumnFault, + }, + /// A bare column name resolved against more than one relation in scope. /// Propagated from `SqlError::AmbiguousColumn`; the pgwire layer renders /// it as SQLSTATE `42702` (ambiguous_column). @@ -357,6 +367,38 @@ pub enum Error { #[error("{detail}")] DataException { detail: String }, + /// A value does not fit its numeric type: an integer past a column's + /// declared width, a float that overflows `REAL`, a literal past the exact + /// numeric range, or arithmetic that overflows. Rendered as SQLSTATE + /// `22003` (numeric_value_out_of_range) at the pgwire layer. + #[error("{detail}")] + NumericValueOutOfRange { detail: String }, + + /// Text does not parse as the type of the column it is written to. + /// `detail` names the column, the text and the type. Rendered as + /// SQLSTATE `22P02` (invalid_text_representation) at the pgwire layer. + #[error("{detail}")] + InvalidTextRepresentation { detail: String }, + + /// A value of the wrong kind for the column it is written to: an array + /// for a `BOOL`, a fractional number for an `INT`. `detail` names the + /// column, the value and the type. Rendered as SQLSTATE `42804` + /// (datatype_mismatch) at the pgwire layer. + #[error("{detail}")] + DatatypeMismatch { detail: String }, + + /// Text does not parse as the `TIMESTAMP` or `TIMESTAMPTZ` it is written + /// to. `detail` names the column, the text and the type. Rendered as + /// SQLSTATE `22007` (invalid_datetime_format) at the pgwire layer. + #[error("{detail}")] + InvalidDatetimeFormat { detail: String }, + + /// An instant outside the range a `TIMESTAMP` or `TIMESTAMPTZ` holds. + /// `detail` names the column, the value and the type. Rendered as + /// SQLSTATE `22008` (datetime_field_overflow) at the pgwire layer. + #[error("{detail}")] + DatetimeFieldOverflow { detail: String }, + /// A LIMIT/OFFSET/FETCH bound did not resolve to `[0, usize::MAX]`. /// The pgwire layer renders it as SQLSTATE `2201W`. #[error("invalid {clause} value: {value}")] diff --git a/nodedb/src/error_classify/public.rs b/nodedb/src/error_classify/public.rs index 1f5ab0241..9db05ecda 100644 --- a/nodedb/src/error_classify/public.rs +++ b/nodedb/src/error_classify/public.rs @@ -183,10 +183,26 @@ pub(crate) fn classify(e: &Error) -> NodeDbError { NodeDbError::object_not_ready(object.clone(), detail.clone()) } Error::UndefinedColumn { column } => NodeDbError::undefined_column(column.clone()), + Error::TextColumn { + collection, + column, + fault, + } => crate::error_from_data_plane::text_column_to_public(collection, column, fault), Error::AmbiguousColumn { column } => NodeDbError::ambiguous_column(column.clone()), Error::UnknownStrictField { column, .. } => NodeDbError::undefined_column(column.clone()), Error::DivisionByZero => NodeDbError::division_by_zero(), Error::DataException { detail } => NodeDbError::data_exception(detail.clone()), + Error::NumericValueOutOfRange { detail } => { + NodeDbError::numeric_value_out_of_range(detail.clone()) + } + // The public class `code_for_sqlstate` gives `22P02`, `22007`, + // `22008` and `42804`. + Error::InvalidTextRepresentation { detail } + | Error::InvalidDatetimeFormat { detail } + | Error::DatetimeFieldOverflow { detail } => NodeDbError::data_exception(detail.clone()), + Error::DatatypeMismatch { detail } => { + NodeDbError::from_wire(nodedb_types::error::ErrorCode::BAD_REQUEST, detail.clone()) + } Error::InvalidLimitValue { clause, value } => { NodeDbError::invalid_limit_value(*clause, value.clone()) } diff --git a/nodedb/src/error_classify/unclassified.rs b/nodedb/src/error_classify/unclassified.rs index 217a9814d..26364b1f6 100644 --- a/nodedb/src/error_classify/unclassified.rs +++ b/nodedb/src/error_classify/unclassified.rs @@ -81,10 +81,16 @@ pub(crate) fn is_unclassified_failure(e: &Error) -> bool { | Error::UndefinedObject { .. } | Error::ObjectNotInPrerequisiteState { .. } | Error::UndefinedColumn { .. } + | Error::TextColumn { .. } | Error::AmbiguousColumn { .. } | Error::UnknownStrictField { .. } | Error::DivisionByZero | Error::DataException { .. } + | Error::NumericValueOutOfRange { .. } + | Error::InvalidTextRepresentation { .. } + | Error::DatatypeMismatch { .. } + | Error::InvalidDatetimeFormat { .. } + | Error::DatetimeFieldOverflow { .. } | Error::InvalidLimitValue { .. } | Error::RetryableSchemaChanged { .. } | Error::RetryableLeaderChange { .. } diff --git a/nodedb/src/error_from.rs b/nodedb/src/error_from.rs index 431577681..c139fe658 100644 --- a/nodedb/src/error_from.rs +++ b/nodedb/src/error_from.rs @@ -29,6 +29,9 @@ impl From for Error { | nodedb_query::EvalError::InvalidJsonPath { .. }) => Self::DataException { detail: e.to_string(), }, + e @ nodedb_query::EvalError::NumericOverflow { .. } => Self::NumericValueOutOfRange { + detail: e.to_string(), + }, } } } @@ -341,6 +344,26 @@ impl From for nodedb_cluster::rpc_codec::TypedClusterError { expected_version: 0, actual_version: 0, }, + // A value refusal crosses as the Data-Plane verdict of the same + // name, so the coordinator answers its exact SQLSTATE. + Error::InvalidTextRepresentation { detail } => TypedClusterError::DataPlane { + code: nodedb_cluster::rpc_codec::DataPlaneErrorCode::InvalidTextRepresentation { + detail, + }, + }, + Error::DatatypeMismatch { detail } => TypedClusterError::DataPlane { + code: nodedb_cluster::rpc_codec::DataPlaneErrorCode::DatatypeMismatch { detail }, + }, + Error::InvalidDatetimeFormat { detail } => TypedClusterError::DataPlane { + code: nodedb_cluster::rpc_codec::DataPlaneErrorCode::InvalidDatetimeFormat { + detail, + }, + }, + Error::DatetimeFieldOverflow { detail } => TypedClusterError::DataPlane { + code: nodedb_cluster::rpc_codec::DataPlaneErrorCode::DatetimeFieldOverflow { + detail, + }, + }, // Every other error crosses as its public numeric code, so a // multi-hop forward keeps its class. other @ (Error::TxnOverlayMemoryExceeded { .. } @@ -387,10 +410,12 @@ impl From for nodedb_cluster::rpc_codec::TypedClusterError { | Error::UndefinedObject { .. } | Error::ObjectNotInPrerequisiteState { .. } | Error::UndefinedColumn { .. } + | Error::TextColumn { .. } | Error::AmbiguousColumn { .. } | Error::UnknownStrictField { .. } | Error::DivisionByZero | Error::DataException { .. } + | Error::NumericValueOutOfRange { .. } | Error::InvalidLimitValue { .. } | Error::RetryableLeaderChange { .. } | Error::CommittedResultUnavailable { .. } diff --git a/nodedb/src/error_from_data_plane.rs b/nodedb/src/error_from_data_plane.rs index 1039b52df..308f9d440 100644 --- a/nodedb/src/error_from_data_plane.rs +++ b/nodedb/src/error_from_data_plane.rs @@ -144,6 +144,11 @@ pub(crate) fn data_plane_code_to_public(code: ErrorCode) -> NodeDbError { add a stricter termination condition or raise max_recursion_depth" )), ErrorCode::UndefinedColumn { column } => NodeDbError::undefined_column(column), + ErrorCode::TextColumn { + collection, + column, + fault, + } => text_column_to_public(&collection, &column, &fault), // `0A000` (feature_not_supported). `SQL_NOT_ENABLED` is the class // every bare `0A000` refusal carries. ErrorCode::Unsupported { detail } => { @@ -152,7 +157,18 @@ pub(crate) fn data_plane_code_to_public(code: ErrorCode) -> NodeDbError { ErrorCode::DivisionByZero => NodeDbError::division_by_zero(), ErrorCode::UndefinedFunction { name } => NodeDbError::undefined_function(name), ErrorCode::DataException { detail } => NodeDbError::data_exception(detail), + ErrorCode::NumericValueOutOfRange { detail } => { + NodeDbError::numeric_value_out_of_range(detail) + } ErrorCode::BadRequest { detail } => NodeDbError::bad_request(detail), + // The public class `code_for_sqlstate` gives `22P02`, `22007`, + // `22008` and `42804`. + ErrorCode::InvalidTextRepresentation { detail } + | ErrorCode::InvalidDatetimeFormat { detail } + | ErrorCode::DatetimeFieldOverflow { detail } => NodeDbError::data_exception(detail), + ErrorCode::DatatypeMismatch { detail } => { + NodeDbError::from_wire(PublicCode::BAD_REQUEST, detail) + } ErrorCode::TransactionRollback { detail } => NodeDbError::transaction_rollback(detail), ErrorCode::ActiveSqlTransaction { detail } => NodeDbError::active_sql_transaction(detail), ErrorCode::DependentObjectsExist { object, detail } => { @@ -167,16 +183,33 @@ pub(crate) fn data_plane_code_to_public(code: ErrorCode) -> NodeDbError { split the transaction into smaller batches" )) } + ErrorCode::NodeLabelLimit { node, label, limit } => { + NodeDbError::program_limit_exceeded(node_label_limit_message(&node, &label, limit)) + } // Genuinely internal: the shard is in an unknown or faulted state. // These are the only codes for which NDB-9000 is the truth. ErrorCode::Internal { detail } => NodeDbError::internal(detail), + // The typed cause of the failed reverse write chains onto the error + // and names itself in the message. ErrorCode::RollbackFailed { entry_index, detail, - } => NodeDbError::internal(format!( - "transaction rollback failed at undo entry {entry_index}: {detail}; \ - shard state is unknown — restart required" - )), + cause, + } => { + let cause = cause.map(|cause| data_plane_code_to_public(*cause)); + let because = cause + .as_ref() + .map(|cause| format!(" ({cause})")) + .unwrap_or_default(); + let error = NodeDbError::internal(format!( + "transaction rollback failed at undo entry {entry_index}: {detail}{because}; \ + shard state is unknown — restart required" + )); + match cause { + Some(cause) => error.with_cause(cause), + None => error, + } + } // A scheduler signal that reached a client: nothing was written, and // the retry that the signal asks for succeeds, so it takes the // retriable class the SQL surfaces send (`40001`). @@ -207,6 +240,35 @@ pub(crate) fn rejected_constraint_to_public( } } +/// The public error of a full-text column fault. A field that does not exist +/// as text is an undefined column (`42703`). An argument that is not a text +/// column is a type mismatch (class `42`). Shared by the Data-Plane code and +/// the Control-Plane variant, so both render one message. +pub(crate) fn text_column_to_public( + collection: &str, + column: &str, + fault: &nodedb_types::text_search::TextColumnFault, +) -> NodeDbError { + use nodedb_types::text_search::TextColumnFault; + let code = match fault { + TextColumnFault::Undeclared | TextColumnFault::NotIndexed => PublicCode::UNDEFINED_COLUMN, + TextColumnFault::NotText { .. } | TextColumnFault::NotAColumn => PublicCode::TYPE_MISMATCH, + }; + NodeDbError::from_wire( + code, + format!("column \"{column}\" of collection \"{collection}\" {fault}"), + ) +} + +/// The message of a refused node-label write. Shared by the public error and +/// the SQLSTATE table, so every protocol renders one text. +pub(crate) fn node_label_limit_message(node: &str, label: &str, limit: usize) -> String { + format!( + "label \"{label}\" on node \"{node}\" exceeds the {limit} distinct node-label \ + limit of the graph partition; no label of the statement was applied" + ) +} + #[cfg(test)] mod tests { use super::*; diff --git a/nodedb/src/error_from_graph.rs b/nodedb/src/error_from_graph.rs new file mode 100644 index 000000000..99b4dcbf3 --- /dev/null +++ b/nodedb/src/error_from_graph.rs @@ -0,0 +1,65 @@ +// SPDX-License-Identifier: BUSL-1.1 + +//! `nodedb_graph::GraphError` into the crate error. +//! +//! A refused memory reservation is the graph engine's memory budget. A +//! rebuild that holds the partition journal leaves the index busy. Every +//! other graph error is a storage fault of the engine. + +use nodedb_graph::GraphError; + +use crate::Error; + +impl From for Error { + fn from(e: GraphError) -> Self { + match e { + // The same class a vector or FTS budget refusal has. + GraphError::MemoryBudget(_) => Self::MemoryExhausted { + engine: "graph".to_string(), + }, + // A rebuild already holds the partition's journal. + GraphError::RebuildInProgress => Self::ObjectNotInPrerequisiteState { + object: "graph CSR index".to_string(), + detail: e.to_string(), + }, + GraphError::LabelOverflow { .. } + | GraphError::NodeOverflow { .. } + | GraphError::WithdrawRefused { .. } + | GraphError::RebuildSuperseded + | GraphError::RebuildJournalOverflow { .. } + | GraphError::RebuildReplayDiverged { .. } + | GraphError::RebuildSnapshotInvalid { .. } => Self::Storage { + engine: "graph".to_string(), + detail: e.to_string(), + }, + // `GraphError` is `#[non_exhaustive]` and lives in another crate, + // so the compiler requires this arm. A variant this build cannot + // name is a storage fault. + _ => Self::Storage { + engine: "graph".to_string(), + detail: e.to_string(), + }, + } + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn a_busy_index_is_not_in_prerequisite_state() { + assert!(matches!( + Error::from(GraphError::RebuildInProgress), + Error::ObjectNotInPrerequisiteState { .. } + )); + } + + #[test] + fn an_exhausted_id_space_is_a_graph_storage_fault() { + assert!(matches!( + Error::from(GraphError::NodeOverflow { used: 1 }), + Error::Storage { ref engine, .. } if engine == "graph" + )); + } +} From 527d01cd97ceb6f879de4dbedfa34d1786cd665a Mon Sep 17 00:00:00 2001 From: Farhan Syah Date: Thu, 8 Oct 2026 16:51:29 +0800 Subject: [PATCH 03/24] fix(types): keep wide integers and non-finite floats exact A u64 above i64::MAX reads as an exact Decimal and writes back as a MessagePack uint64 instead of wrapping negative. Numbers compare through one shared exact order: integers never round through f64, an integer against a float compares exactly, and NaN sorts above every number. The order backs filters, sorts, group keys and JSON comparisons, and a DECIMAL sort key orders by value. An integer literal past i64 stays an exact Decimal, or is refused with 22003 past the Decimal range. A bound float parameter stays a float. NaN and +/-Infinity render as their PostgreSQL text instead of JSON null. Timeseries ingest refuses an unsigned field an Int64 column cannot hold. --- nodedb-query/src/fusion.rs | 4 +- nodedb-query/src/json_expr.rs | 94 ++++++- nodedb-query/src/json_ops.rs | 199 ++++++++++++++- nodedb-query/src/lib.rs | 3 + nodedb-query/src/msgpack_scan/compare.rs | 95 ++++++- .../src/msgpack_scan/compare_decimal.rs | 98 ++++++++ nodedb-query/src/msgpack_scan/group_key.rs | 19 +- nodedb-query/src/msgpack_scan/mod.rs | 8 +- nodedb-query/src/msgpack_scan/reader/mod.rs | 6 +- .../src/msgpack_scan/reader/scalar.rs | 80 ++++-- nodedb-query/src/msgpack_scan/reader/value.rs | 25 +- nodedb-query/src/ts_functions/percentile.rs | 2 +- nodedb-query/src/value_ops.rs | 231 +++++++++++++++++- nodedb-sql/src/dsl_bind.rs | 2 +- nodedb-sql/src/params.rs | 36 ++- nodedb-sql/src/planner/const_fold.rs | 32 ++- nodedb-sql/src/resolver/expr/value.rs | 84 ++++++- nodedb-types/src/conversion.rs | 33 +-- nodedb-types/src/json_msgpack/json_value.rs | 9 +- nodedb-types/src/json_msgpack/reader/json.rs | 29 ++- .../src/json_msgpack/reader/native.rs | 42 +++- nodedb-types/src/json_msgpack/transcoder.rs | 24 +- nodedb-types/src/json_msgpack/writer.rs | 16 +- nodedb-types/src/lib.rs | 5 +- nodedb-types/src/numeric_cmp.rs | 229 +++++++++++++++++ nodedb-types/src/value/coerce.rs | 187 ++++++++++---- nodedb-types/src/value/convert.rs | 30 +++ nodedb-types/src/value/float_text.rs | 62 +++++ nodedb-types/src/value/json.rs | 77 +++++- nodedb-types/src/value/mod.rs | 2 + nodedb/src/bridge/json_ops.rs | 163 ------------ .../sql_plan_convert/value/msgpack_write.rs | 43 +++- .../src/control/server/response_shape/cell.rs | 62 +++-- .../shared/ddl/neutral/graph_ops/algo.rs | 6 +- .../executor/handlers/columnar_read/sort.rs | 42 +++- .../handlers/document/sort/compare.rs | 106 +++++++- .../handlers/document/sort/external.rs | 24 +- .../handlers/document/sort/in_memory.rs | 119 ++++----- .../executor/handlers/document/sort/mod.rs | 3 +- .../handlers/timeseries/ingest_formats.rs | 12 +- .../handlers/timeseries/msgpack_decode.rs | 8 +- .../executor/handlers/timeseries/normalize.rs | 124 +++++++++- .../handlers/timeseries/resolve_ingest.rs | 6 +- .../data/executor/handlers/timeseries/sort.rs | 60 ++++- nodedb/src/engine/timeseries/ilp_ingest.rs | 36 +++ .../wire/cases/aggregate_integer_exactness.rs | 211 ++++++++++++++++ .../wire/cases/aggregate_non_finite_float.rs | 94 +++++++ nodedb/tests/wire/cases/sql_u64_literal.rs | 135 ++++++++++ 48 files changed, 2562 insertions(+), 455 deletions(-) create mode 100644 nodedb-query/src/msgpack_scan/compare_decimal.rs create mode 100644 nodedb-types/src/numeric_cmp.rs create mode 100644 nodedb-types/src/value/float_text.rs delete mode 100644 nodedb/src/bridge/json_ops.rs create mode 100644 nodedb/tests/wire/cases/aggregate_integer_exactness.rs create mode 100644 nodedb/tests/wire/cases/aggregate_non_finite_float.rs create mode 100644 nodedb/tests/wire/cases/sql_u64_literal.rs 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/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/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-sql/src/dsl_bind.rs b/nodedb-sql/src/dsl_bind.rs index 3b0e7c7e6..6b9c19b2d 100644 --- a/nodedb-sql/src/dsl_bind.rs +++ b/nodedb-sql/src/dsl_bind.rs @@ -117,7 +117,7 @@ fn placeholder_literal_token(placeholder: &str, params: &[ParamValue]) -> Option ParamValue::Bool(true) => Token::make_keyword("TRUE"), ParamValue::Bool(false) => Token::make_keyword("FALSE"), ParamValue::Int64(n) => Token::Number(n.to_string(), false), - ParamValue::Float64(f) => Token::Number(f.to_string(), false), + ParamValue::Float64(f) => Token::Number(crate::params::float_literal_text(*f), false), ParamValue::Decimal(d) => Token::Number(d.to_string(), false), ParamValue::Text(s) => Token::SingleQuotedString(s.clone()), ParamValue::Timestamp(dt) | ParamValue::Timestamptz(dt) => { diff --git a/nodedb-sql/src/params.rs b/nodedb-sql/src/params.rs index 2657ae7ec..e111818f0 100644 --- a/nodedb-sql/src/params.rs +++ b/nodedb-sql/src/params.rs @@ -84,7 +84,7 @@ fn placeholder_to_value(placeholder: &str, params: &[ParamValue]) -> Option Value::Boolean(true), ParamValue::Bool(false) => Value::Boolean(false), ParamValue::Int64(n) => Value::Number(n.to_string(), false), - ParamValue::Float64(f) => Value::Number(f.to_string(), false), + ParamValue::Float64(f) => Value::Number(float_literal_text(*f), false), ParamValue::Decimal(d) => Value::Number(d.to_string(), false), ParamValue::Text(s) => Value::SingleQuotedString(s.clone()), // Timestamp/Timestamptz: emit as a typed SQL literal so the resolver @@ -94,6 +94,17 @@ fn placeholder_to_value(placeholder: &str, params: &[ParamValue]) -> Option String { + format!("{f:e}") +} + #[cfg(test)] mod tests { use super::*; @@ -305,6 +316,29 @@ mod tests { assert!(!result.contains("$1"), "got: {result}"); } + /// A float parameter resolves to a float, whole or fractional, never to + /// the decimal or integer its plain text reads as. + #[test] + fn a_float_param_resolves_to_a_float() { + use crate::resolver::expr::convert_value; + use crate::types::SqlValue; + for f in [1.5, 2.0, 0.1, -3.25, 1e300, 1e-300, f64::INFINITY] { + let literal = placeholder_to_value("$1", &[ParamValue::Float64(f)]) + .expect("the placeholder names the one parameter"); + assert_eq!( + convert_value(&literal).expect("a float literal resolves"), + SqlValue::Float(f), + "{f}" + ); + } + let literal = placeholder_to_value("$1", &[ParamValue::Float64(f64::NAN)]) + .expect("the placeholder names the one parameter"); + assert!( + matches!(convert_value(&literal), Ok(SqlValue::Float(f)) if f.is_nan()), + "{literal}" + ); + } + #[test] fn bind_window_order_by_placeholder() { let result = bind_and_format( diff --git a/nodedb-sql/src/planner/const_fold.rs b/nodedb-sql/src/planner/const_fold.rs index ec8e8f00e..f8f526a8f 100644 --- a/nodedb-sql/src/planner/const_fold.rs +++ b/nodedb-sql/src/planner/const_fold.rs @@ -97,6 +97,14 @@ pub fn fold_constant_scoped( // than wrapping (release) or panicking (debug). Some(SqlValue::Int(i)) => i.checked_neg().map(SqlValue::Int), Some(SqlValue::Float(f)) => Some(SqlValue::Float(-f)), + // `9223372036854775808` is past `i64::MAX`, so its literal is a + // `Decimal`. Negated it is `i64::MIN`, an integer like every + // other in range. + Some(SqlValue::Decimal(d)) + if d.scale() == 0 && -d == rust_decimal::Decimal::from(i64::MIN) => + { + Some(SqlValue::Int(i64::MIN)) + } Some(SqlValue::Decimal(d)) => Some(SqlValue::Decimal(-d)), _ => None, }), @@ -392,7 +400,8 @@ pub fn fold_function_call_scoped( Err( e @ (nodedb_query::EvalError::VectorDimensionMismatch { .. } | nodedb_query::EvalError::ArgumentType { .. } - | nodedb_query::EvalError::InvalidJsonPath { .. }), + | nodedb_query::EvalError::InvalidJsonPath { .. } + | nodedb_query::EvalError::NumericOverflow { .. }), ) => Err(SqlError::DataException { detail: e.to_string(), }), @@ -484,6 +493,27 @@ mod tests { } } + /// `-9223372036854775808` folds to the integer `i64::MIN`, and a negated + /// literal past it stays an exact `Decimal`. + #[test] + fn negated_integer_literal_past_i64_max_stays_exact() { + let registry = FunctionRegistry::new(); + let neg = |digits: &str| SqlExpr::UnaryOp { + op: UnaryOp::Neg, + expr: Box::new(SqlExpr::Literal(SqlValue::Decimal( + rust_decimal::Decimal::from_str_exact(digits).unwrap(), + ))), + }; + assert_eq!( + fold_constant(&neg("9223372036854775808"), ®istry).unwrap(), + Some(SqlValue::Int(i64::MIN)) + ); + assert_eq!( + fold_constant(&neg("18446744073709551615"), ®istry).unwrap(), + Some(SqlValue::Decimal(-rust_decimal::Decimal::from(u64::MAX))) + ); + } + /// A constant call whose argument the function cannot compute on fails /// the statement at plan time instead of folding to NULL. #[test] diff --git a/nodedb-sql/src/resolver/expr/value.rs b/nodedb-sql/src/resolver/expr/value.rs index 25fb703db..184f067ac 100644 --- a/nodedb-sql/src/resolver/expr/value.rs +++ b/nodedb-sql/src/resolver/expr/value.rs @@ -8,16 +8,28 @@ use crate::types::*; /// Convert a sqlparser `Value` to our `SqlValue`. /// /// Number literal routing: -/// - Pure integers → `SqlValue::Int`. -/// - Numbers with `.`, `e`, or `E` → `SqlValue::Decimal` (exact arithmetic). +/// - Integers in `i64` range → `SqlValue::Int`. +/// - Larger integers → `SqlValue::Decimal`, exact. A `u64` above `i64::MAX` +/// is the `Decimal` `Value::from_u64` gives it, so it is written as a +/// msgpack `uint64`. An integer past the `Decimal` range is +/// [`SqlError::NumericLiteralOutOfRange`]: rounding it to a float would +/// store another number than the one written. +/// - Numbers with an exponent (`e` or `E`) → `SqlValue::Float`. A bound +/// float parameter renders in this form, so it stays a float. +/// - Other numbers with `.` → `SqlValue::Decimal` (exact arithmetic). /// - If decimal parse fails → fallback to `SqlValue::Float`, then `SqlValue::String`. pub fn convert_value(val: &Value) -> Result { match val { Value::Number(n, _) => { if let Ok(i) = n.parse::() { Ok(SqlValue::Int(i)) - } else if n.contains('.') || n.contains('e') || n.contains('E') { - // Fractional or scientific notation: prefer exact Decimal. + } else if n.contains('e') || n.contains('E') { + match n.parse::() { + Ok(f) => Ok(SqlValue::Float(f)), + Err(_) => Ok(SqlValue::String(n.clone())), + } + } else if n.contains('.') { + // A fractional literal: prefer exact Decimal. if let Ok(d) = rust_decimal::Decimal::from_str_exact(n) { Ok(SqlValue::Decimal(d)) } else if let Ok(f) = n.parse::() { @@ -25,6 +37,8 @@ pub fn convert_value(val: &Value) -> Result { } else { Ok(SqlValue::String(n.clone())) } + } else if n.bytes().all(|b| b.is_ascii_digit()) { + integer_literal(n) } else if let Ok(f) = n.parse::() { Ok(SqlValue::Float(f)) } else { @@ -43,6 +57,18 @@ pub fn convert_value(val: &Value) -> Result { } } +/// An all-digit literal past `i64::MAX` as an exact `Decimal`. +fn integer_literal(digits: &str) -> Result { + if let Ok(u) = digits.parse::() { + return Ok(SqlValue::Decimal(rust_decimal::Decimal::from(u))); + } + rust_decimal::Decimal::from_str_exact(digits) + .map(SqlValue::Decimal) + .map_err(|_| SqlError::NumericLiteralOutOfRange { + literal: digits.to_string(), + }) +} + /// Decode an even-length ASCII hex string (the inner text of an `X'...'` /// literal) into its byte sequence. fn decode_hex_literal(s: &str) -> Result> { @@ -91,7 +117,55 @@ pub(super) fn parse_interval_to_micros(s: &str) -> Option { #[cfg(test)] mod tests { - use super::parse_interval_to_micros; + use super::*; + + fn number(n: &str) -> Value { + Value::Number(n.to_string(), false) + } + + #[test] + fn integer_literal_past_i64_stays_exact() { + assert_eq!( + convert_value(&number("9223372036854775807")).unwrap(), + SqlValue::Int(i64::MAX) + ); + assert_eq!( + convert_value(&number("18446744073709551615")).unwrap(), + SqlValue::Decimal(rust_decimal::Decimal::from(u64::MAX)) + ); + // Past `u64`, inside the 96-bit `Decimal` mantissa. + assert_eq!( + convert_value(&number("79228162514264337593543950335")).unwrap(), + SqlValue::Decimal(rust_decimal::Decimal::MAX) + ); + } + + #[test] + fn an_exponent_literal_is_a_float_and_a_fraction_is_a_decimal() { + assert_eq!( + convert_value(&number("1.5e0")).unwrap(), + SqlValue::Float(1.5) + ); + assert_eq!(convert_value(&number("2E0")).unwrap(), SqlValue::Float(2.0)); + assert_eq!( + convert_value(&number("1e300")).unwrap(), + SqlValue::Float(1e300) + ); + assert_eq!( + convert_value(&number("1.5")).unwrap(), + SqlValue::Decimal(rust_decimal::Decimal::new(15, 1)) + ); + } + + #[test] + fn integer_literal_past_decimal_range_is_an_error() { + let err = convert_value(&number("79228162514264337593543950336")).unwrap_err(); + assert!( + matches!(err, SqlError::NumericLiteralOutOfRange { ref literal } + if literal == "79228162514264337593543950336"), + "{err:?}" + ); + } #[test] fn parse_interval_sql_word_forms() { diff --git a/nodedb-types/src/conversion.rs b/nodedb-types/src/conversion.rs index 197064673..73f02e39e 100644 --- a/nodedb-types/src/conversion.rs +++ b/nodedb-types/src/conversion.rs @@ -6,6 +6,19 @@ use crate::Value; +/// A JSON number as the `Value` that keeps its exact number: an `i64` as an +/// `Integer`, a larger `u64` through [`Value::from_u64`], anything else as a +/// `Float`. +fn json_number(n: &serde_json::Number) -> Value { + if let Some(i) = n.as_i64() { + Value::Integer(i) + } else if let Some(u) = n.as_u64() { + Value::from_u64(u) + } else { + n.as_f64().map_or(Value::Null, Value::Float) + } +} + /// Convert a `serde_json::Value` to a `Value` by consuming ownership. /// /// Nested objects are preserved as `Value::Object`. @@ -13,13 +26,7 @@ pub fn json_to_value(v: serde_json::Value) -> Value { 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) => json_number(&n), serde_json::Value::String(s) => Value::String(s), serde_json::Value::Array(arr) => Value::Array(arr.into_iter().map(json_to_value).collect()), serde_json::Value::Object(obj) => Value::Object( @@ -50,13 +57,7 @@ pub fn json_to_value_ref(v: &serde_json::Value) -> Value { 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) => json_number(n), serde_json::Value::String(s) => Value::String(s.clone()), serde_json::Value::Array(arr) => Value::Array(arr.iter().map(json_to_value_ref).collect()), serde_json::Value::Object(map) => Value::Object( @@ -111,6 +112,10 @@ mod tests { Value::Bool(true) ); assert_eq!(json_to_value(serde_json::json!(42)), Value::Integer(42)); + let max = serde_json::json!(u64::MAX); + let exact = Value::Decimal(rust_decimal::Decimal::from(u64::MAX)); + assert_eq!(json_to_value_ref(&max), exact); + assert_eq!(json_to_value(max), exact); assert_eq!( json_to_value(serde_json::json!("hello")), Value::String("hello".into()) diff --git a/nodedb-types/src/json_msgpack/json_value.rs b/nodedb-types/src/json_msgpack/json_value.rs index 28fb02d81..19ae8a549 100644 --- a/nodedb-types/src/json_msgpack/json_value.rs +++ b/nodedb-types/src/json_msgpack/json_value.rs @@ -4,6 +4,8 @@ use zerompk::{ToMessagePack, Write}; +use crate::value::float_text::float_to_json; + /// Newtype wrapper around `serde_json::Value` implementing zerompk traits. #[derive(Debug, Clone, PartialEq)] pub struct JsonValue(pub serde_json::Value); @@ -109,13 +111,14 @@ fn read_json_from_reader<'a, R: zerompk::Read<'a>>(reader: &mut R) -> zerompk::R if let Ok(u) = reader.read_u64() { return Ok(JsonValue(serde_json::Value::Number(u.into()))); } - // Try f64 (0xCB) + // Try f64 (0xCB). A non-finite float is its PostgreSQL text in a JSON + // string, never `null`. if let Ok(f) = reader.read_f64() { - return Ok(JsonValue(serde_json::json!(f))); + return Ok(JsonValue(float_to_json(f))); } // Try f32 (0xCA) if let Ok(f) = reader.read_f32() { - return Ok(JsonValue(serde_json::json!(f as f64))); + return Ok(JsonValue(float_to_json(f64::from(f)))); } // Try string (fixstr 0xA0-0xBF, str8 0xD9, str16 0xDA, str32 0xDB) if let Ok(s) = reader.read_string() { diff --git a/nodedb-types/src/json_msgpack/reader/json.rs b/nodedb-types/src/json_msgpack/reader/json.rs index 61f449e60..4aea713fc 100644 --- a/nodedb-types/src/json_msgpack/reader/json.rs +++ b/nodedb-types/src/json_msgpack/reader/json.rs @@ -9,6 +9,7 @@ use super::super::error::MsgpackResult; use super::super::instant_ext::instant_from_ext; use super::cursor::Cursor; use crate::datetime::NdbDateTime; +use crate::value::float_text::float_to_json; /// Deserialize a `serde_json::Value` from MessagePack bytes. /// @@ -64,16 +65,18 @@ fn read_json_value(c: &mut Cursor<'_>) -> zerompk::Result { )) } + // A non-finite float is its PostgreSQL text in a JSON string, never + // `null`. 0xCA => { let b = c.take_n(4)?; - Ok(serde_json::json!( - f32::from_be_bytes([b[0], b[1], b[2], b[3]]) as f64 - )) + Ok(float_to_json(f64::from(f32::from_be_bytes([ + b[0], b[1], b[2], b[3], + ])))) } 0xCB => { let b = c.take_n(8)?; - Ok(serde_json::json!(f64::from_be_bytes([ - b[0], b[1], b[2], b[3], b[4], b[5], b[6], b[7] + Ok(float_to_json(f64::from_be_bytes([ + b[0], b[1], b[2], b[3], b[4], b[5], b[6], b[7], ]))) } @@ -282,6 +285,22 @@ mod tests { assert_eq!(val, restored); } + /// A non-finite msgpack float reads as its PostgreSQL text in a JSON + /// string, never `null`. + #[test] + fn non_finite_float_reads_as_postgres_text() { + let mut bytes = vec![0x93, 0xCB]; + bytes.extend_from_slice(&f64::NAN.to_be_bytes()); + bytes.push(0xCB); + bytes.extend_from_slice(&f64::INFINITY.to_be_bytes()); + bytes.push(0xCA); + bytes.extend_from_slice(&f32::NEG_INFINITY.to_be_bytes()); + assert_eq!( + json_from_msgpack(&bytes).unwrap(), + json!(["NaN", "Infinity", "-Infinity"]) + ); + } + #[test] fn roundtrip_string() { let val = json!("hello world"); diff --git a/nodedb-types/src/json_msgpack/reader/native.rs b/nodedb-types/src/json_msgpack/reader/native.rs index ba0cc4386..70a900d21 100644 --- a/nodedb-types/src/json_msgpack/reader/native.rs +++ b/nodedb-types/src/json_msgpack/reader/native.rs @@ -64,7 +64,7 @@ pub(crate) fn read_native_value<'de, R: Read<'de>>(reader: &mut R) -> zerompk::R } 0xC2 | 0xC3 => Ok(Value::Bool(reader.read_boolean()?)), 0x00..=0x7F | 0xE0..=0xFF | 0xD0..=0xD3 => Ok(Value::Integer(reader.read_i64()?)), - 0xCC..=0xCF => Ok(Value::Integer(reader.read_u64()? as i64)), + 0xCC..=0xCF => Ok(Value::from_u64(reader.read_u64()?)), 0xCA => Ok(Value::Float(f64::from(reader.read_f32()?))), 0xCB => Ok(Value::Float(reader.read_f64()?)), 0xA0..=0xBF | 0xD9..=0xDB => Ok(Value::String(reader.read_string()?.into_owned())), @@ -213,6 +213,46 @@ mod tests { assert_eq!(value_from_msgpack(&bytes).unwrap(), Value::Integer(256)); } + #[test] + fn unsigned_64_above_i64_max_keeps_its_number() { + let bytes = [0xCF, 0xFF, 0xFF, 0xFF, 0xFF, 0xFF, 0xFF, 0xFF, 0xFF]; + let expected = Value::Decimal(rust_decimal::Decimal::from(u64::MAX)); + assert_eq!(through_cell(&bytes).unwrap(), expected); + assert_eq!(value_from_msgpack(&bytes).unwrap(), expected); + } + + /// A `u64` above `i64::MAX` writes as `uint64` and reads back as the same + /// `Decimal`. Every other decimal stays text. + #[test] + fn wide_u64_decimal_round_trips_as_uint64() { + let wide = Value::from_u64(u64::MAX); + let bytes = value_to_msgpack(&wide).unwrap(); + assert_eq!( + bytes, + [0xCF, 0xFF, 0xFF, 0xFF, 0xFF, 0xFF, 0xFF, 0xFF, 0xFF] + ); + assert_eq!(value_from_msgpack(&bytes).unwrap(), wide); + assert_eq!( + msgpack_to_json_string(&bytes).unwrap(), + "18446744073709551615" + ); + + let just_above = Value::from_u64(i64::MAX as u64 + 1); + let bytes = value_to_msgpack(&just_above).unwrap(); + assert_eq!(bytes[0], 0xCF); + assert_eq!(value_from_msgpack(&bytes).unwrap(), just_above); + + for text in ["5", "18446744073709551615.0", "18446744073709551616", "-1"] { + let d = Value::Decimal(rust_decimal::Decimal::from_str_exact(text).unwrap()); + let bytes = value_to_msgpack(&d).unwrap(); + assert_eq!( + value_from_msgpack(&bytes).unwrap(), + Value::String(text.to_string()), + "{text} stays text" + ); + } + } + #[test] fn unknown_ext_reads_as_null() { let bytes = [0xD4, 0x07, 0x00]; diff --git a/nodedb-types/src/json_msgpack/transcoder.rs b/nodedb-types/src/json_msgpack/transcoder.rs index 228669b59..693e65c55 100644 --- a/nodedb-types/src/json_msgpack/transcoder.rs +++ b/nodedb-types/src/json_msgpack/transcoder.rs @@ -15,6 +15,7 @@ use super::error::MsgpackResult; use super::instant_ext::instant_from_ext; use super::reader::{Cursor, base64_encode}; use crate::datetime::NdbDateTime; +use crate::value::float_text::non_finite_float_text; /// Transcode raw msgpack bytes to a JSON string without intermediate types. /// @@ -203,9 +204,13 @@ fn write_uint(out: &mut String, v: u64) { let _ = write!(out, "{v}"); } +/// Write a float as JSON. A non-finite float is its PostgreSQL text in a +/// JSON string, as PostgreSQL `to_json` renders it, never `null`. fn write_float(out: &mut String, v: f64) { - if v.is_nan() || v.is_infinite() { - out.push_str("null"); + if let Some(text) = non_finite_float_text(v) { + out.push('"'); + out.push_str(text); + out.push('"'); } else if v.fract() == 0.0 && v.abs() < (1i64 << 53) as f64 { let _ = write!(out, "{v:.1}"); } else { @@ -336,6 +341,21 @@ mod tests { assert_eq!(msgpack_to_json_string(&unknown).unwrap(), "null"); } + /// A non-finite float transcodes to its PostgreSQL text in a JSON + /// string, never `null`. A finite float stays a JSON number. + #[test] + fn non_finite_floats_render_postgres_text() { + let mut mp = vec![0x94]; + for f in [f64::NAN, f64::INFINITY, f64::NEG_INFINITY, 2.5] { + mp.push(0xCB); + mp.extend_from_slice(&f.to_be_bytes()); + } + assert_eq!( + msgpack_to_json_string(&mp).unwrap(), + "[\"NaN\",\"Infinity\",\"-Infinity\",2.5]" + ); + } + #[test] fn nested() { let val = serde_json::json!({"a": {"b": [1, 2, {"c": 3}]}}); diff --git a/nodedb-types/src/json_msgpack/writer.rs b/nodedb-types/src/json_msgpack/writer.rs index 2cd85ae50..944805fcb 100644 --- a/nodedb-types/src/json_msgpack/writer.rs +++ b/nodedb-types/src/json_msgpack/writer.rs @@ -4,8 +4,9 @@ //! //! `Value::DateTime` and `Value::NaiveDateTime` are written as the instant //! ext (`fixext8`, see `instant_ext`). `Duration`, `Decimal`, and `Geometry` -//! are written as strings. `Vector` is a float64 array, `ArrayCell` a map, -//! and `Range` / `Record` are `nil`. +//! are written as strings, except a `Decimal` that holds a `u64` above +//! `i64::MAX`, which is a `uint64`. `Vector` is a float64 array, `ArrayCell` +//! a map, and `Range` / `Record` are `nil`. use zerompk::Write; @@ -48,8 +49,10 @@ impl zerompk::ToMessagePack for NativeRef<'_> { /// Write a `nodedb_types::Value` as standard msgpack. /// -/// `Duration`, `Decimal`, and `Geometry` are strings. `Vector` is a float64 -/// array, `ArrayCell` a map, and `Range` / `Record` are `nil`. +/// `Duration`, `Decimal`, and `Geometry` are strings, except a `Decimal` +/// that holds a `u64` above `i64::MAX` (see [`crate::Value::decimal_as_wide_u64`]), +/// which is a `uint64`. `Vector` is a float64 array, `ArrayCell` a map, and +/// `Range` / `Record` are `nil`. pub(crate) fn write_native_value( writer: &mut W, value: &crate::Value, @@ -82,7 +85,10 @@ pub(crate) fn write_native_value( crate::Value::DateTime(dt) => write_instant_ext(writer, InstantKind::Utc, dt.micros), crate::Value::NaiveDateTime(dt) => write_instant_ext(writer, InstantKind::Naive, dt.micros), crate::Value::Duration(d) => writer.write_string(&d.to_string()), - crate::Value::Decimal(d) => writer.write_string(&d.to_string()), + crate::Value::Decimal(d) => match crate::Value::decimal_as_wide_u64(d) { + Some(u) => writer.write_u64(u), + None => writer.write_string(&d.to_string()), + }, crate::Value::Geometry(g) => match sonic_rs::to_string(g) { Ok(s) => writer.write_string(&s), Err(_) => writer.write_nil(), diff --git a/nodedb-types/src/lib.rs b/nodedb-types/src/lib.rs index d43404d4e..a1a207475 100644 --- a/nodedb-types/src/lib.rs +++ b/nodedb-types/src/lib.rs @@ -47,6 +47,7 @@ pub mod lsn; pub mod mirror; pub mod multi_vector; pub mod namespace; +pub mod numeric_cmp; pub mod path_component; pub mod pg_compat; pub mod protocol; @@ -122,7 +123,7 @@ pub use result::{QueryResult, SearchResult, SubGraph}; pub use rls_write_check::{RlsWriteCheck, WriteGateDecision}; pub use row_identity::{ DEFAULT_IDENTITY_COLUMN, HEADLESS_SENTINEL_PREFIX, ROWID_COLUMN, RowIdentity, StorageKey, - extract_pk_value, value_to_pk_string, + declared_key, extract_pk_value, value_to_pk_string, }; pub use sparse_vector::{SparseVector, SparseVectorError}; pub use sql_quote::{quote_ident, quote_literal}; @@ -137,7 +138,7 @@ pub use temporal::{ MAX_POSITIONS_PER_EPOCH, NANOS_PER_MS, OPEN_UPPER, OrdinalClock, SystemTimeScope, ValidTimePredicate, calvin_txn_ordinal, ms_to_ordinal_upper, ordinal_to_ms, }; -pub use text_search::{Bm25Params, QueryMode, TextSearchParams}; +pub use text_search::{Bm25Params, QueryMode, TextColumnFault, TextSearchParams}; pub use trace::{SpanId, TraceId}; pub use typeguard::TypeGuardFieldDef; pub use value::{NotScalar, Value, scalar_to_raw_bytes}; diff --git a/nodedb-types/src/numeric_cmp.rs b/nodedb-types/src/numeric_cmp.rs new file mode 100644 index 000000000..b00ee7fff --- /dev/null +++ b/nodedb-types/src/numeric_cmp.rs @@ -0,0 +1,229 @@ +// SPDX-License-Identifier: Apache-2.0 + +//! Numeric readings and their total order, shared by every comparison path. +//! +//! - An integer reads as an exact `i128`. `i128` holds every `i64` and +//! every `u64`. +//! - A fractional `Decimal` reads as an exact `Decimal`. A `DECIMAL` cell is +//! stored as its text, so numeric text with a fraction reads as a +//! `Decimal` too. +//! - A float reads as an `f64`. +//! - Integer and decimal pairs compare exactly. An integer against a float +//! compares exactly. A decimal against a float compares as `f64`. +//! - NaN sorts above every number and equals NaN, as in PostgreSQL. The +//! order is total, so a sort over it never panics. + +use std::cmp::Ordering; + +use rust_decimal::Decimal; +use rust_decimal::prelude::ToPrimitive; + +/// The numeric reading of a value for comparison and summation. +#[derive(Debug, Clone, Copy, PartialEq)] +pub enum Numeric { + Int(i128), + Decimal(Decimal), + Float(f64), +} + +/// Total order of two floats. NaN sorts above every number and equals NaN. +/// `-0.0` equals `0.0`. +pub fn cmp_f64(a: f64, b: f64) -> Ordering { + match a.partial_cmp(&b) { + Some(order) => order, + None => a.is_nan().cmp(&b.is_nan()), + } +} + +/// A numeric string read as a number: integer text as an exact integer, +/// fractional decimal text as an exact `Decimal`, and any other float text +/// (an exponent, a fraction past the `Decimal` precision, `NaN`, `inf`) as +/// an `f64`. `None` for non-numeric text. +pub fn parse_numeric_str(s: &str) -> Option { + if let Ok(i) = s.parse::() { + return Some(Numeric::Int(i)); + } + let float = s.parse::().ok()?; + Some(match Decimal::from_str_exact(s) { + Ok(d) => decimal_reading(&d), + Err(_) => Numeric::Float(float), + }) +} + +/// An integral `Decimal` as an exact integer, any other as an exact +/// `Decimal`. +pub fn decimal_reading(d: &Decimal) -> Numeric { + if d.is_integer() + && let Some(i) = d.to_i128() + { + return Numeric::Int(i); + } + Numeric::Decimal(*d) +} + +/// `2^127` as `f64`: the first float above every `i128`. +const TWO_POW_127: f64 = 170_141_183_460_469_231_731_687_303_715_884_105_728.0; + +/// Exact order of an `i128` against an `f64`, with no rounding of either. +/// NaN is above every integer. +fn cmp_int_float(i: i128, f: f64) -> Ordering { + if f.is_nan() || f >= TWO_POW_127 { + return Ordering::Less; + } + if f < -TWO_POW_127 { + return Ordering::Greater; + } + // `f` lies in `[-2^127, 2^127)`, so its integral part is an exact `i128`. + let whole = f.trunc(); + let by_whole = i.cmp(&(whole as i128)); + if by_whole != Ordering::Equal { + return by_whole; + } + let fraction = f - whole; + if fraction > 0.0 { + Ordering::Less + } else if fraction < 0.0 { + Ordering::Greater + } else { + Ordering::Equal + } +} + +/// Exact order of an `i128` against a `Decimal`. An integer past the +/// `Decimal` range is past every `Decimal` on its side of zero. +fn cmp_int_decimal(i: i128, d: Decimal) -> Ordering { + match Decimal::try_from_i128_with_scale(i, 0) { + Ok(as_decimal) => as_decimal.cmp(&d), + Err(_) if i > 0 => Ordering::Greater, + Err(_) => Ordering::Less, + } +} + +/// Total order of two numeric readings. See the module docs for the rule. +pub fn cmp_numeric(a: Numeric, b: Numeric) -> Ordering { + match (a, b) { + (Numeric::Int(x), Numeric::Int(y)) => x.cmp(&y), + (Numeric::Int(x), Numeric::Decimal(y)) => cmp_int_decimal(x, y), + (Numeric::Decimal(x), Numeric::Int(y)) => cmp_int_decimal(y, x).reverse(), + (Numeric::Decimal(x), Numeric::Decimal(y)) => x.cmp(&y), + (Numeric::Int(x), Numeric::Float(y)) => cmp_int_float(x, y), + (Numeric::Float(x), Numeric::Int(y)) => cmp_int_float(y, x).reverse(), + (Numeric::Decimal(x), Numeric::Float(y)) => cmp_f64(x.as_f64(), y), + (Numeric::Float(x), Numeric::Decimal(y)) => cmp_f64(x, y.as_f64()), + (Numeric::Float(x), Numeric::Float(y)) => cmp_f64(x, y), + } +} + +/// Equality of two numeric readings for a coerced `=`: [`cmp_numeric`] +/// orders the pair `Equal`. Exact for every pair, as PostgreSQL float `=` +/// is, and NaN equals NaN. +pub fn numeric_eq(a: Numeric, b: Numeric) -> bool { + cmp_numeric(a, b) == Ordering::Equal +} + +#[cfg(test)] +mod tests { + use super::*; + + fn dec(text: &str) -> Decimal { + Decimal::from_str_exact(text).unwrap() + } + + #[test] + fn nan_is_above_every_number_and_equals_nan() { + assert_eq!(cmp_f64(f64::NAN, f64::INFINITY), Ordering::Greater); + assert_eq!(cmp_f64(f64::NEG_INFINITY, f64::NAN), Ordering::Less); + assert_eq!(cmp_f64(f64::NAN, f64::NAN), Ordering::Equal); + assert_eq!(cmp_f64(-0.0, 0.0), Ordering::Equal); + let nan = Numeric::Float(f64::NAN); + assert_eq!(cmp_numeric(nan, Numeric::Int(i128::MAX)), Ordering::Greater); + assert_eq!(cmp_numeric(Numeric::Int(i128::MIN), nan), Ordering::Less); + assert_eq!( + cmp_numeric(nan, Numeric::Decimal(Decimal::MAX)), + Ordering::Greater + ); + assert_eq!( + cmp_numeric(Numeric::Decimal(dec("0.5")), nan), + Ordering::Less + ); + assert_eq!(cmp_numeric(nan, nan), Ordering::Equal); + assert!(numeric_eq(nan, nan)); + } + + /// A sort over NaN, infinities and every reading kind runs to a + /// consistent order. + #[test] + fn mixed_sort_is_total() { + let mut values = [ + Numeric::Float(f64::NAN), + Numeric::Int(3), + Numeric::Float(f64::INFINITY), + Numeric::Decimal(dec("2.5")), + Numeric::Float(f64::NAN), + Numeric::Float(f64::NEG_INFINITY), + Numeric::Int(-1), + ]; + values.sort_by(|a, b| cmp_numeric(*a, *b)); + assert_eq!(values[0], Numeric::Float(f64::NEG_INFINITY)); + assert_eq!(values[1], Numeric::Int(-1)); + assert_eq!(values[2], Numeric::Decimal(dec("2.5"))); + assert_eq!(values[3], Numeric::Int(3)); + assert_eq!(values[4], Numeric::Float(f64::INFINITY)); + assert!(matches!(values[5], Numeric::Float(f) if f.is_nan())); + assert!(matches!(values[6], Numeric::Float(f) if f.is_nan())); + } + + /// Two decimals one hundredth apart past `2^53` compare exactly. + #[test] + fn fractional_decimal_text_compares_exactly() { + let low = parse_numeric_str("12345678901234567.01").unwrap(); + let high = parse_numeric_str("12345678901234567.02").unwrap(); + assert_eq!(low, Numeric::Decimal(dec("12345678901234567.01"))); + assert_eq!(cmp_numeric(low, high), Ordering::Less); + assert_eq!(cmp_numeric(high, low), Ordering::Greater); + assert!(!numeric_eq(low, high)); + } + + #[test] + fn decimal_against_integer_compares_exactly() { + let above = Numeric::Decimal(dec("9007199254740993.5")); + assert_eq!( + cmp_numeric(above, Numeric::Int(9_007_199_254_740_993)), + Ordering::Greater + ); + assert_eq!( + cmp_numeric(Numeric::Int(9_007_199_254_740_994), above), + Ordering::Greater + ); + assert_eq!( + cmp_numeric(Numeric::Int(i128::MAX), above), + Ordering::Greater + ); + assert_eq!(cmp_numeric(Numeric::Int(i128::MIN), above), Ordering::Less); + assert_eq!( + cmp_numeric(Numeric::Decimal(dec("-0.5")), Numeric::Int(0)), + Ordering::Less + ); + } + + #[test] + fn numeric_text_reads_by_kind() { + assert_eq!(parse_numeric_str("12"), Some(Numeric::Int(12))); + assert_eq!(parse_numeric_str("5.0"), Some(Numeric::Int(5))); + assert_eq!(parse_numeric_str("0.1"), Some(Numeric::Decimal(dec("0.1")))); + assert_eq!(parse_numeric_str("1e3"), Some(Numeric::Float(1000.0))); + assert!(matches!(parse_numeric_str("NaN"), Some(Numeric::Float(f)) if f.is_nan())); + assert_eq!(parse_numeric_str("x"), None); + assert_eq!(parse_numeric_str("1_0.5"), None); + } + + #[test] + fn integral_decimal_reads_as_integer() { + assert_eq!(decimal_reading(&dec("42.000")), Numeric::Int(42)); + assert_eq!( + decimal_reading(&Decimal::MAX), + Numeric::Int(Decimal::MAX.mantissa()) + ); + assert_eq!(decimal_reading(&dec("0.25")), Numeric::Decimal(dec("0.25"))); + } +} diff --git a/nodedb-types/src/value/coerce.rs b/nodedb-types/src/value/coerce.rs index 8ad93e798..4cd0d19a3 100644 --- a/nodedb-types/src/value/coerce.rs +++ b/nodedb-types/src/value/coerce.rs @@ -3,22 +3,27 @@ //! Type-coerced equality and ordering for `Value`. //! //! Single source of truth for type coercion in filter/sort evaluation. +//! Numbers compare through the shared order in [`crate::numeric_cmp`]. + +use std::cmp::Ordering; use super::core::Value; +use crate::numeric_cmp::{Numeric, cmp_numeric, decimal_reading, parse_numeric_str}; impl Value { /// Coerced equality: `Value` vs `Value` with numeric/string coercion. /// /// Single source of truth for type coercion in filter evaluation. /// Used by `matches_binary` (msgpack path) and `matches_value` (Value path). + /// + /// Two strings are equal by their text or by the instant they denote. + /// Any other pair where both sides read as numbers is equal when + /// [`cmp_numeric`] orders it `Equal`: exact for integers and decimals, + /// and NaN equals NaN. pub fn eq_coerced(&self, other: &Value) -> bool { match (self, other) { (Value::Null, Value::Null) => true, (Value::Bool(a), Value::Bool(b)) => a == b, - (Value::Integer(a), Value::Integer(b)) => a == b, - (Value::Integer(a), Value::Float(b)) => *a as f64 == *b, - (Value::Float(a), Value::Integer(b)) => *a == *b as f64, - (Value::Float(a), Value::Float(b)) => a == b, (Value::String(a), Value::String(b)) => { a == b || matches!( @@ -26,46 +31,32 @@ impl Value { (Some(x), Some(y)) if x.micros == y.micros ) } - // Coercion: number vs string - (Value::Integer(a), Value::String(s)) => { - s.parse::().is_ok_and(|n| *a == n) - || s.parse::().is_ok_and(|n| *a as f64 == n) - } - (Value::String(s), Value::Integer(b)) => { - s.parse::().is_ok_and(|n| n == *b) - || s.parse::().is_ok_and(|n| n == *b as f64) - } - (Value::Float(a), Value::String(s)) => s.parse::().is_ok_and(|n| *a == n), - (Value::String(s), Value::Float(b)) => s.parse::().is_ok_and(|n| n == *b), // Structural equality on ND cells: same coords and same attrs. (Value::ArrayCell(a), Value::ArrayCell(b)) => a == b, - // Two exact decimals compare exactly; a decimal against any other - // number compares through f64 like the arms above. - (Value::Decimal(a), Value::Decimal(b)) => a == b, - (Value::Decimal(_), _) | (_, Value::Decimal(_)) => { - match (numeric_f64(self), numeric_f64(other)) { + (a, b) => { + if let (Some(x), Some(y)) = (numeric_reading(a), numeric_reading(b)) { + return cmp_numeric(x, y) == Ordering::Equal; + } + match (datetime_micros(a), datetime_micros(b)) { (Some(x), Some(y)) => x == y, _ => false, } } - (a, b) => match (datetime_micros(a), datetime_micros(b)) { - (Some(x), Some(y)) => x == y, - _ => false, - }, } } /// Coerced partial ordering for predicate evaluation. /// - /// Two numbers (or numeric strings) order numerically, two instants (or - /// ISO-8601 strings) by epoch microseconds, two other strings - /// lexicographically, and two ND cells coordinate-major. A pair with no - /// defined order — an integer against an instant, text against a number, - /// a NaN — is `None`, so a range predicate over it matches nothing rather - /// than every row: the row-level counterpart of PostgreSQL refusing to + /// Two numbers (or numeric strings) order by [`cmp_numeric`]: exact for + /// integers and decimals, NaN above every number and equal to NaN, as + /// in PostgreSQL. Two instants (or ISO-8601 strings) order by epoch + /// microseconds, two other strings lexicographically, and two ND cells + /// coordinate-major. A pair with no defined order — an integer against + /// an instant, text against a number, a bool against a number — is + /// `None`, so a range predicate over it matches nothing rather than + /// every row: the row-level counterpart of PostgreSQL refusing to /// compare the two types. - pub fn partial_cmp_coerced(&self, other: &Value) -> Option { - use std::cmp::Ordering; + pub fn partial_cmp_coerced(&self, other: &Value) -> Option { if let (Value::ArrayCell(a), Value::ArrayCell(b)) = (self, other) { for (x, y) in a.coords.iter().zip(b.coords.iter()) { match x.partial_cmp_coerced(y)? { @@ -85,8 +76,8 @@ impl Value { } return Some(a.attrs.len().cmp(&b.attrs.len())); } - if let (Some(a), Some(b)) = (numeric_f64(self), numeric_f64(other)) { - return a.partial_cmp(&b); + if let (Some(a), Some(b)) = (numeric_reading(self), numeric_reading(other)) { + return Some(cmp_numeric(a, b)); } if let (Some(a), Some(b)) = (datetime_micros(self), datetime_micros(other)) { return Some(a.cmp(&b)); @@ -104,21 +95,20 @@ impl Value { /// for ORDER BY / MIN / MAX style paths that need an `Ordering` for every /// pair; a predicate uses `partial_cmp_coerced` so an unordered pair /// matches nothing. - pub fn cmp_coerced(&self, other: &Value) -> std::cmp::Ordering { - self.partial_cmp_coerced(other) - .unwrap_or(std::cmp::Ordering::Equal) + pub fn cmp_coerced(&self, other: &Value) -> Ordering { + self.partial_cmp_coerced(other).unwrap_or(Ordering::Equal) } } -/// The number a value denotes for coerced ordering: an integer, a float, or -/// a string that parses as one. `None` for anything else. -fn numeric_f64(v: &Value) -> Option { - use rust_decimal::prelude::ToPrimitive; +/// The number a value denotes for coerced comparison: an integer, a float, +/// a decimal, or a string that parses as a number. A bool has no numeric +/// reading here. `None` for anything else. +fn numeric_reading(v: &Value) -> Option { match v { - Value::Integer(i) => Some(*i as f64), - Value::Float(f) => Some(*f), - Value::Decimal(d) => d.to_f64(), - Value::String(s) => s.parse::().ok(), + Value::Integer(i) => Some(Numeric::Int(i128::from(*i))), + Value::Float(f) => Some(Numeric::Float(*f)), + Value::Decimal(d) => Some(decimal_reading(d)), + Value::String(s) => parse_numeric_str(s), _ => None, } } @@ -270,7 +260,7 @@ mod tests { /// An integer carries no unit, so it has no order against an instant: /// `WHERE at >= 5` over an instant column matches nothing, in both - /// orientations. The same holds for text against a number and for NaN. + /// orientations. The same holds for text against a number. #[test] fn partial_cmp_coerced_is_none_for_an_unordered_pair() { let instant = Value::NaiveDateTime(crate::NdbDateTime::from_micros(1_583_402_400_000_000)); @@ -285,9 +275,110 @@ mod tests { None ); assert_eq!(Value::Null.partial_cmp_coerced(&Value::Integer(0)), None); + } + + /// NaN sorts above every number and equals NaN, as in PostgreSQL, so + /// `WHERE f > 1` matches a NaN row and `WHERE f = 'NaN'` matches it. + #[test] + fn nan_orders_above_every_number_and_equals_nan() { + let nan = Value::Float(f64::NAN); assert_eq!( - Value::Float(f64::NAN).partial_cmp_coerced(&Value::Float(1.0)), - None + nan.partial_cmp_coerced(&Value::Float(1.0)), + Some(Ordering::Greater) + ); + assert_eq!( + nan.partial_cmp_coerced(&Value::Float(f64::INFINITY)), + Some(Ordering::Greater) + ); + assert_eq!( + Value::Integer(i64::MAX).partial_cmp_coerced(&nan), + Some(Ordering::Less) + ); + assert_eq!( + Value::Decimal(rust_decimal::Decimal::MAX).partial_cmp_coerced(&nan), + Some(Ordering::Less) + ); + assert_eq!( + nan.partial_cmp_coerced(&Value::Float(f64::NAN)), + Some(Ordering::Equal) + ); + assert_eq!( + nan.partial_cmp_coerced(&Value::String("NaN".into())), + Some(Ordering::Equal) + ); + assert!(nan.eq_coerced(&Value::Float(f64::NAN))); + assert!(!nan.eq_coerced(&Value::Float(1.0))); + let mut values = [ + nan.clone(), + Value::Integer(3), + Value::Float(f64::NEG_INFINITY), + Value::Float(0.5), + ]; + values.sort_by(Value::cmp_coerced); + 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())); + } + + /// Two decimals one hundredth apart past `2^53` collapse to one `f64`. + /// They compare exactly, as decimals and as decimal text. + #[test] + fn decimal_pairs_compare_exactly() { + let dec = |s: &str| rust_decimal::Decimal::from_str_exact(s).expect("decimal"); + let low = Value::Decimal(dec("12345678901234567.01")); + let high = Value::Decimal(dec("12345678901234567.02")); + assert_eq!(low.partial_cmp_coerced(&high), Some(Ordering::Less)); + assert_eq!(high.cmp_coerced(&low), Ordering::Greater); + assert!(!low.eq_coerced(&high)); + let high_text = Value::String("12345678901234567.02".into()); + assert_eq!(low.partial_cmp_coerced(&high_text), Some(Ordering::Less)); + assert!(high.eq_coerced(&high_text)); + assert!(!low.eq_coerced(&high_text)); + let half = Value::Decimal(dec("9007199254740993.5")); + assert_eq!( + half.partial_cmp_coerced(&Value::Integer(9_007_199_254_740_993)), + Some(Ordering::Greater) + ); + assert_eq!( + Value::Integer(9_007_199_254_740_994).partial_cmp_coerced(&half), + Some(Ordering::Greater) + ); + } + + /// An `i64` near the limit compares exactly against a float, with no + /// rounding of the integer to `f64`. + #[test] + fn large_integers_compare_exactly_against_floats() { + const TWO_POW_63: f64 = 9_223_372_036_854_775_808.0; + let max = Value::Integer(i64::MAX); + assert_eq!( + max.partial_cmp_coerced(&Value::Float(TWO_POW_63)), + Some(Ordering::Less) + ); + assert!(!max.eq_coerced(&Value::Float(TWO_POW_63))); + assert_eq!( + Value::Integer(i64::MIN).partial_cmp_coerced(&Value::Float(-TWO_POW_63)), + Some(Ordering::Equal) + ); + assert!(Value::Integer(i64::MIN).eq_coerced(&Value::Float(-TWO_POW_63))); + let above = Value::Integer(9_007_199_254_740_993); + let float = Value::Float(9_007_199_254_740_992.0); + assert_eq!(above.partial_cmp_coerced(&float), Some(Ordering::Greater)); + assert_eq!(float.partial_cmp_coerced(&above), Some(Ordering::Less)); + assert!(!above.eq_coerced(&float)); + assert!(!float.eq_coerced(&above)); + assert_eq!( + Value::Integer(i64::MAX - 1).partial_cmp_coerced(&Value::Integer(i64::MAX)), + Some(Ordering::Less) + ); + assert_eq!( + Value::Integer(i64::MIN).partial_cmp_coerced(&Value::Float(f64::NEG_INFINITY)), + Some(Ordering::Greater) + ); + assert!( + !Value::String("9007199254740993".into()) + .eq_coerced(&Value::Integer(9_007_199_254_740_992)) ); } diff --git a/nodedb-types/src/value/convert.rs b/nodedb-types/src/value/convert.rs index 06e885284..fa1cb78d5 100644 --- a/nodedb-types/src/value/convert.rs +++ b/nodedb-types/src/value/convert.rs @@ -26,6 +26,36 @@ impl From for Value { } } +impl Value { + /// An unsigned integer as a `Value` that keeps its exact number. + /// + /// A value up to `i64::MAX` is an `Integer`. A larger value is a + /// `Decimal`, because an `Integer` cannot hold it without wrapping. + pub fn from_u64(u: u64) -> Self { + match i64::try_from(u) { + Ok(i) => Value::Integer(i), + Err(_) => Value::Decimal(rust_decimal::Decimal::from(u)), + } + } + + /// The `u64` that `d` stands for when `d` is a `Decimal` + /// [`Value::from_u64`] produces: scale `0`, above `i64::MAX`, at most + /// `u64::MAX`. `None` for any other decimal. + /// + /// Msgpack writers encode such a decimal as a `uint64`, the number type + /// that holds it, and readers decode a `uint64` through `from_u64`, so + /// the value round-trips to the same `Decimal`. Every other decimal stays + /// text: an `Integer`-range or scaled decimal would not read back as a + /// `Decimal` of the same scale. + pub fn decimal_as_wide_u64(d: &rust_decimal::Decimal) -> Option { + use rust_decimal::prelude::ToPrimitive; + if d.scale() != 0 { + return None; + } + d.to_u64().filter(|u| i64::try_from(*u).is_err()) + } +} + impl From for Value { fn from(f: f64) -> Self { Value::Float(f) diff --git a/nodedb-types/src/value/float_text.rs b/nodedb-types/src/value/float_text.rs new file mode 100644 index 000000000..9fb7f6366 --- /dev/null +++ b/nodedb-types/src/value/float_text.rs @@ -0,0 +1,62 @@ +// SPDX-License-Identifier: Apache-2.0 + +//! PostgreSQL text and JSON forms of a float, including the non-finite +//! values JSON numbers cannot hold. +//! +//! PostgreSQL renders a non-finite float as `NaN`, `Infinity` or +//! `-Infinity`, and `to_json` renders it as that text in a JSON string. +//! A JSON number cannot hold these values, and `serde_json` turns them into +//! `null`. Every path that renders a float to text or JSON calls these +//! functions instead. + +/// The PostgreSQL text of a non-finite float: `NaN`, `Infinity` or +/// `-Infinity`. `None` for a finite float. +pub fn non_finite_float_text(f: f64) -> Option<&'static str> { + (!f.is_finite()).then(|| non_finite_text(f)) +} + +/// A float as a JSON value. A finite float is a JSON number. A non-finite +/// float is its PostgreSQL text in a JSON string, as `to_json` renders it. +pub fn float_to_json(f: f64) -> serde_json::Value { + match serde_json::Number::from_f64(f) { + Some(n) => serde_json::Value::Number(n), + // `from_f64` refuses exactly the non-finite floats. + None => serde_json::Value::String(non_finite_text(f).to_owned()), + } +} + +/// The PostgreSQL text of a float the caller knows is non-finite. +fn non_finite_text(f: f64) -> &'static str { + if f.is_nan() { + "NaN" + } else if f > 0.0 { + "Infinity" + } else { + "-Infinity" + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn non_finite_floats_render_postgres_text() { + assert_eq!(non_finite_float_text(f64::NAN), Some("NaN")); + assert_eq!(non_finite_float_text(f64::INFINITY), Some("Infinity")); + assert_eq!(non_finite_float_text(f64::NEG_INFINITY), Some("-Infinity")); + assert_eq!(non_finite_float_text(1.5), None); + assert_eq!(non_finite_float_text(f64::MAX), None); + } + + #[test] + fn non_finite_floats_are_json_strings() { + assert_eq!(float_to_json(f64::NAN), serde_json::json!("NaN")); + assert_eq!(float_to_json(f64::INFINITY), serde_json::json!("Infinity")); + assert_eq!( + float_to_json(f64::NEG_INFINITY), + serde_json::json!("-Infinity") + ); + assert_eq!(float_to_json(2.5), serde_json::json!(2.5)); + } +} diff --git a/nodedb-types/src/value/json.rs b/nodedb-types/src/value/json.rs index e508efdcc..abc2f68c1 100644 --- a/nodedb-types/src/value/json.rs +++ b/nodedb-types/src/value/json.rs @@ -12,7 +12,9 @@ impl From for serde_json::Value { Value::Null => serde_json::Value::Null, Value::Bool(b) => serde_json::Value::Bool(b), Value::Integer(i) => serde_json::json!(i), - Value::Float(f) => serde_json::json!(f), + // A non-finite float is its PostgreSQL text in a JSON string, + // never `null`. + Value::Float(f) => super::float_text::float_to_json(f), Value::String(s) | Value::Uuid(s) | Value::Ulid(s) | Value::Regex(s) => { serde_json::Value::String(s) } @@ -32,16 +34,31 @@ impl From for serde_json::Value { Value::Decimal(d) => { // Represent as a JSON Number so clients see a numeric type, // not a quoted string. `from_str` handles the decimal notation - // produced by rust_decimal's `to_string`. + // produced by rust_decimal's `to_string`. An integer past the + // 64-bit range parses to a rounded float, so it keeps its + // exact digits as a string. let s = d.to_string(); - serde_json::Number::from_str(&s) - .map(serde_json::Value::Number) - .unwrap_or_else(|_| serde_json::Value::String(s)) + match serde_json::Number::from_str(&s) { + Ok(n) if !(d.scale() == 0 && n.is_f64()) => serde_json::Value::Number(n), + _ => serde_json::Value::String(s), + } } Value::Geometry(g) => serde_json::to_value(g).unwrap_or(serde_json::Value::Null), Value::Range { .. } | Value::Record { .. } => serde_json::Value::Null, Value::Vector(v) => { - serde_json::Value::Array(v.iter().map(|f| serde_json::json!(*f)).collect()) + // A finite element keeps its shortest `f32` text. A + // non-finite one is its PostgreSQL text, never `null`. + serde_json::Value::Array( + v.iter() + .map(|f| { + if f.is_finite() { + serde_json::json!(*f) + } else { + super::float_text::float_to_json(f64::from(*f)) + } + }) + .collect(), + ) } Value::ArrayCell(cell) => { let mut obj = serde_json::json!({ @@ -66,7 +83,7 @@ impl From for Value { if let Some(i) = n.as_i64() { Value::Integer(i) } else if let Some(u) = n.as_u64() { - Value::Integer(u as i64) + Value::from_u64(u) } else if let Some(f) = n.as_f64() { Value::Float(f) } else { @@ -89,6 +106,30 @@ mod tests { use super::*; use crate::array_cell::ArrayCell; + #[test] + fn json_u64_above_i64_max_keeps_its_number() { + let json: serde_json::Value = + sonic_rs::from_str("18446744073709551615").expect("parse u64::MAX"); + assert_eq!( + Value::from(json), + Value::Decimal(rust_decimal::Decimal::from(u64::MAX)) + ); + let small: serde_json::Value = sonic_rs::from_str("42").expect("parse 42"); + assert_eq!(Value::from(small), Value::Integer(42)); + } + + #[test] + fn integer_decimal_past_64_bits_keeps_exact_digits() { + let big = rust_decimal::Decimal::from_i128_with_scale(2 * i128::from(u64::MAX), 0); + let json = serde_json::Value::from(Value::Decimal(big)); + assert_eq!( + json, + serde_json::Value::String("36893488147419103230".into()) + ); + let max = serde_json::Value::from(Value::Decimal(rust_decimal::Decimal::from(u64::MAX))); + assert_eq!(max, serde_json::json!(u64::MAX)); + } + #[test] fn decimal_to_json_is_number_not_string() { let d = rust_decimal::Decimal::new(12345, 2); // 123.45 @@ -205,4 +246,26 @@ mod tests { "ArrayCell round-trips through JSON as Object, got {rt:?}" ); } + + /// A non-finite float renders as its PostgreSQL text in a JSON string, + /// at the top level, nested, and inside a vector. + #[test] + fn non_finite_floats_render_postgres_text_not_null() { + assert_eq!( + serde_json::Value::from(Value::Float(f64::NAN)), + serde_json::json!("NaN") + ); + assert_eq!( + serde_json::Value::from(Value::Array(vec![ + Value::Float(f64::INFINITY), + Value::Float(f64::NEG_INFINITY), + Value::Float(1.5), + ])), + serde_json::json!(["Infinity", "-Infinity", 1.5]) + ); + assert_eq!( + serde_json::Value::from(Value::Vector(vec![0.5, f32::NAN].into())), + serde_json::json!([0.5, "NaN"]) + ); + } } diff --git a/nodedb-types/src/value/mod.rs b/nodedb-types/src/value/mod.rs index 621113b6c..be7f61443 100644 --- a/nodedb-types/src/value/mod.rs +++ b/nodedb-types/src/value/mod.rs @@ -4,10 +4,12 @@ pub mod coerce; pub mod convert; pub mod core; pub mod display; +pub mod float_text; pub mod json; pub mod msgpack; pub mod raw_bytes; pub mod sql_literal; pub use core::Value; +pub use float_text::{float_to_json, non_finite_float_text}; pub use raw_bytes::{NotScalar, scalar_to_raw_bytes}; diff --git a/nodedb/src/bridge/json_ops.rs b/nodedb/src/bridge/json_ops.rs deleted file mode 100644 index af0a0c52f..000000000 --- a/nodedb/src/bridge/json_ops.rs +++ /dev/null @@ -1,163 +0,0 @@ -// SPDX-License-Identifier: BUSL-1.1 - -//! Shared JSON value operations: comparison, coercion, truthiness. -//! -//! Used by both `expr_eval` (computed projections) and `scan_filter` -//! (WHERE predicate evaluation on the Data Plane). - -use std::cmp::Ordering; - -/// Coerce a JSON value to f64. -/// -/// - Numbers: `as_f64()` directly -/// - Strings: parse as f64 (`"5"` → `5.0`) -/// - Booleans: `true` → `1.0`, `false` → `0.0` (when `coerce_bool` is true) -/// - Other types: `None` -pub fn json_to_f64(v: &serde_json::Value, coerce_bool: bool) -> Option { - match v { - serde_json::Value::Number(n) => n.as_f64(), - serde_json::Value::String(s) => s.parse::().ok(), - serde_json::Value::Bool(b) if coerce_bool => Some(if *b { 1.0 } else { 0.0 }), - _ => None, - } -} - -/// Compare two JSON values with type coercion. -/// -/// Tries numeric comparison first (with bool coercion), then falls -/// back to string comparison. -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); - } - // Fallback: string comparison. - let sa = json_to_display_string(a); - let sb = json_to_display_string(b); - sa.cmp(&sb) -} - -/// Compare two optional JSON values (for scan_filter compatibility). -pub fn compare_json_optional( - a: Option<&serde_json::Value>, - b: Option<&serde_json::Value>, -) -> Ordering { - match (a, b) { - (None, None) => Ordering::Equal, - (None, Some(_)) => Ordering::Less, - (Some(_), None) => Ordering::Greater, - (Some(a), Some(b)) => compare_json(a, b), - } -} - -/// 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. -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; - } - false -} - -/// Check if a JSON value is truthy (for boolean contexts). -/// -/// - `true` → true, `false` → false -/// - `null` → false -/// - Numbers: non-zero → true -/// - Strings: non-empty → true -/// - Arrays/Objects: always true -pub fn is_truthy(v: &serde_json::Value) -> bool { - match v { - serde_json::Value::Bool(b) => *b, - serde_json::Value::Null => false, - serde_json::Value::Number(n) => n.as_f64().unwrap_or(0.0) != 0.0, - serde_json::Value::String(s) => !s.is_empty(), - _ => true, - } -} - -/// Convert a JSON value to a display string. -/// -/// - Strings: returned as-is (no quotes) -/// - Null: empty string -/// - Numbers/Bools: `.to_string()` -/// - Objects/Arrays: JSON serialization -pub fn json_to_display_string(v: &serde_json::Value) -> String { - match v { - serde_json::Value::String(s) => s.clone(), - serde_json::Value::Null => String::new(), - serde_json::Value::Number(n) => n.to_string(), - serde_json::Value::Bool(b) => b.to_string(), - other => other.to_string(), - } -} - -/// Convert a f64 to a JSON number, preferring integer representation. -/// -/// Returns `Null` for NaN/Infinity. -pub fn to_json_number(n: f64) -> serde_json::Value { - if n.fract() == 0.0 && n.abs() < i64::MAX as f64 { - serde_json::Value::Number(serde_json::Number::from(n as i64)) - } else { - serde_json::Number::from_f64(n) - .map(serde_json::Value::Number) - .unwrap_or(serde_json::Value::Null) - } -} - -#[cfg(test)] -mod tests { - use super::*; - use serde_json::json; - - #[test] - fn coerced_eq_mixed_types() { - assert!(coerced_eq(&json!(5), &json!("5"))); - assert!(coerced_eq(&json!(3.15), &json!("3.15"))); - assert!(!coerced_eq(&json!(5), &json!("6"))); - } - - #[test] - fn coerced_eq_bool_numeric() { - assert!(coerced_eq(&json!(true), &json!(1))); - assert!(coerced_eq(&json!(false), &json!(0))); - assert!(!coerced_eq(&json!(true), &json!(0))); - } - - #[test] - fn compare_numeric_coercion() { - assert_eq!(compare_json(&json!(5), &json!("4")), Ordering::Greater); - assert_eq!(compare_json(&json!("10"), &json!(9)), Ordering::Greater); - } - - #[test] - fn truthiness() { - assert!(is_truthy(&json!(true))); - assert!(!is_truthy(&json!(false))); - assert!(!is_truthy(&json!(null))); - assert!(is_truthy(&json!(1))); - assert!(!is_truthy(&json!(0))); - assert!(is_truthy(&json!("hello"))); - assert!(!is_truthy(&json!(""))); - } - - #[test] - fn to_json_number_nan() { - assert_eq!(to_json_number(f64::NAN), serde_json::Value::Null); - } - - #[test] - fn to_json_number_integer() { - assert_eq!(to_json_number(42.0), json!(42)); - } - - #[test] - fn to_json_number_float() { - assert_eq!(to_json_number(3.15), json!(3.15)); - } -} diff --git a/nodedb/src/control/planner/sql_plan_convert/value/msgpack_write.rs b/nodedb/src/control/planner/sql_plan_convert/value/msgpack_write.rs index 4a0f2f81b..257123a6f 100644 --- a/nodedb/src/control/planner/sql_plan_convert/value/msgpack_write.rs +++ b/nodedb/src/control/planner/sql_plan_convert/value/msgpack_write.rs @@ -100,14 +100,20 @@ pub(crate) fn write_msgpack_value_with(buf: &mut Vec, val: &SqlValue, instan buf.push(0xCB); buf.extend_from_slice(&f.to_be_bytes()); } - SqlValue::Decimal(d) => { - // Write as msgpack string of the decimal's canonical text representation. - // The Data Plane strict encoder accepts Value::String for Decimal columns - // and coerces via rust_decimal::Decimal::from_str. This matches the - // contract of write_msgpack_value: it produces standard msgpack that - // json_from_msgpack / value_from_msgpack can decode without loss. - write_msgpack_str(buf, &d.to_string()); - } + SqlValue::Decimal(d) => match nodedb_types::Value::decimal_as_wide_u64(d) { + // An integer literal past `i64::MAX` that fits `u64` is a msgpack + // `uint64`, the number type that holds it. Readers decode it + // through `Value::from_u64`, back to this same `Decimal`. + Some(u) => { + buf.push(0xCF); + buf.extend_from_slice(&u.to_be_bytes()); + } + // Any other decimal is the msgpack string of its canonical text. + // The Data Plane strict encoder accepts Value::String for Decimal + // columns and coerces via rust_decimal::Decimal::from_str, and + // json_from_msgpack / value_from_msgpack decode it without loss. + None => write_msgpack_str(buf, &d.to_string()), + }, SqlValue::String(s) => write_msgpack_str(buf, s), SqlValue::Array(arr) => { write_msgpack_array_header(buf, arr.len()); @@ -191,6 +197,27 @@ mod tests { assert_eq!(buf, vec![42]); } + /// A literal past `i64::MAX` that fits `u64` is a msgpack `uint64` and + /// reads back as the same exact value. Any other decimal stays text. + #[test] + fn wide_integer_decimal_is_uint64() { + let wide = rust_decimal::Decimal::from(u64::MAX); + let mut buf = Vec::new(); + write_msgpack_value(&mut buf, &SqlValue::Decimal(wide)); + assert_eq!(buf, [0xCF, 0xFF, 0xFF, 0xFF, 0xFF, 0xFF, 0xFF, 0xFF, 0xFF]); + assert_eq!( + nodedb_types::value_from_msgpack(&buf).unwrap(), + nodedb_types::Value::Decimal(wide) + ); + + let mut buf = Vec::new(); + let fractional = rust_decimal::Decimal::from_str_exact("1.5").unwrap(); + write_msgpack_value(&mut buf, &SqlValue::Decimal(fractional)); + let mut expected = Vec::new(); + write_msgpack_str(&mut expected, "1.5"); + assert_eq!(buf, expected); + } + #[test] fn write_msgpack_value_string() { let mut buf = Vec::new(); diff --git a/nodedb/src/control/server/response_shape/cell.rs b/nodedb/src/control/server/response_shape/cell.rs index c3c977d21..a97830a2c 100644 --- a/nodedb/src/control/server/response_shape/cell.rs +++ b/nodedb/src/control/server/response_shape/cell.rs @@ -11,7 +11,6 @@ //! ([`cell_text`]) and, for a timestamp column, the instant it denotes //! ([`instant_of`]). -use base64::Engine; use nodedb_types::error::NodeDbError; use nodedb_types::{NdbDateTime, Value}; @@ -28,11 +27,13 @@ pub fn row_to_wire_json(row: &ShapedRow) -> serde_json::Map Option { @@ -40,16 +41,18 @@ pub fn cell_text(v: &Value) -> Option { Value::Null => None, Value::Bool(b) => Some(if *b { "t" } else { "f" }.to_owned()), Value::Integer(i) => Some(i.to_string()), - Value::Float(f) => serde_json::Number::from_f64(*f).map(|n| n.to_string()), + Value::Float(f) => wire_json_text(&nodedb_types::value::float_to_json(*f)), Value::String(s) => Some(s.clone()), - Value::Bytes(bytes) => Some(base64::engine::general_purpose::STANDARD_NO_PAD.encode(bytes)), + Value::Bytes(bytes) => Some(bytea_hex(bytes)), + // Every digit and the scale, as PostgreSQL renders `numeric`. The + // wire JSON number is an `f64` and drops both. + Value::Decimal(d) => Some(d.to_string()), Value::DateTime(at) | Value::NaiveDateTime(at) => Some(at.to_iso8601()), Value::Array(_) | Value::Object(_) | Value::Uuid(_) | Value::Ulid(_) | Value::Duration(_) - | Value::Decimal(_) | Value::Geometry(_) | Value::Set(_) | Value::Regex(_) @@ -76,6 +79,15 @@ fn wire_json_text(v: &serde_json::Value) -> Option { } } +/// The PostgreSQL text output of a `bytea`: `\x` followed by two lowercase +/// hex digits per byte. +pub fn bytea_hex(bytes: &[u8]) -> String { + let mut text = String::with_capacity(2 + bytes.len() * 2); + text.push_str("\\x"); + text.push_str(&hex::encode(bytes)); + text +} + /// The instant a cell under a timestamp column denotes. /// /// A typed instant is itself. A string is the instant it parses to as @@ -142,8 +154,10 @@ mod tests { /// The one instant these tests use: 2020-03-05T10:00:00Z. const EARLY_MICROS: i64 = 1_583_402_400_000_000; - /// Every scalar renders the same text its wire JSON renders to, so the - /// direct arms and the JSON edge cannot drift apart. + /// Every scalar except bytes and decimals renders the same text its wire + /// JSON renders to, so the direct arms and the JSON edge cannot drift + /// apart. Bytes are base64 in JSON and `\x` hex in PostgreSQL text. A + /// decimal is an `f64` JSON number and exact digits in PostgreSQL text. #[test] fn scalar_text_matches_the_wire_json_text() { let at = NdbDateTime::from_micros(EARLY_MICROS); @@ -159,11 +173,9 @@ mod tests { Value::Float(f64::NAN), Value::Float(f64::INFINITY), Value::String("hello".into()), - Value::Bytes(vec![0, 255, 7]), Value::DateTime(at), Value::NaiveDateTime(at), Value::Uuid("550e8400-e29b-41d4-a716-446655440000".into()), - Value::Decimal(rust_decimal::Decimal::new(110, 2)), Value::Duration(nodedb_types::NdbDuration::from_micros(1_500_000)), Value::Array(vec![Value::Integer(1), Value::Bool(true)]), Value::Range { @@ -190,11 +202,31 @@ mod tests { assert_eq!(cell_text(&Value::Bool(false)).as_deref(), Some("f")); assert_eq!(cell_text(&Value::Integer(42)).as_deref(), Some("42")); assert_eq!(cell_text(&Value::Float(0.0)).as_deref(), Some("0.0")); - assert_eq!(cell_text(&Value::Float(f64::NAN)), None); + assert_eq!(cell_text(&Value::Float(f64::NAN)).as_deref(), Some("NaN")); + assert_eq!( + cell_text(&Value::Float(f64::INFINITY)).as_deref(), + Some("Infinity") + ); + assert_eq!( + cell_text(&Value::Float(f64::NEG_INFINITY)).as_deref(), + Some("-Infinity") + ); assert_eq!(cell_text(&Value::String("x".into())).as_deref(), Some("x")); assert_eq!( cell_text(&Value::Bytes(vec![0, 255, 7])).as_deref(), - Some("AP8H") + Some("\\x00ff07") + ); + assert_eq!(cell_text(&Value::Bytes(Vec::new())).as_deref(), Some("\\x")); + assert_eq!( + cell_text(&Value::Decimal(rust_decimal::Decimal::new(110, 2))).as_deref(), + Some("1.10") + ); + assert_eq!( + cell_text(&Value::Decimal( + "0.1234567890123456789012345678".parse().expect("decimal") + )) + .as_deref(), + Some("0.1234567890123456789012345678") ); assert_eq!( cell_text(&Value::NaiveDateTime(at)).as_deref(), diff --git a/nodedb/src/control/server/shared/ddl/neutral/graph_ops/algo.rs b/nodedb/src/control/server/shared/ddl/neutral/graph_ops/algo.rs index 40056c1f1..a281a0189 100644 --- a/nodedb/src/control/server/shared/ddl/neutral/graph_ops/algo.rs +++ b/nodedb/src/control/server/shared/ddl/neutral/graph_ops/algo.rs @@ -182,9 +182,9 @@ fn clamp_opt( /// Render an algorithm result payload into a protocol-neutral row set. /// /// Every column is emitted as `Text` with its cell pre-rendered to the exact -/// string (all algorithm result columns are text): `Text` → the raw string, `Float64` → `format!("{v}")` or the -/// literal `Infinity` for a non-representable/non-finite score, `Int64` → -/// decimal or `0`. Pre-rendering keeps the wire bytes byte-identical (a native +/// string (all algorithm result columns are text): `Text` → the raw string, +/// `Float64` → `format!("{v}")` or the literal `Infinity` for a +/// non-representable/non-finite score, `Int64` → decimal or `0`. Pre-rendering keeps the wire bytes byte-identical (a native /// float path will change both the column OID and the `Infinity` fallback). fn algo_payload_to_rows( payload: &crate::bridge::envelope::Payload, diff --git a/nodedb/src/data/executor/handlers/columnar_read/sort.rs b/nodedb/src/data/executor/handlers/columnar_read/sort.rs index 21710a061..aa7d880f5 100644 --- a/nodedb/src/data/executor/handlers/columnar_read/sort.rs +++ b/nodedb/src/data/executor/handlers/columnar_read/sort.rs @@ -83,13 +83,14 @@ fn compare_values(a: &Value, b: &Value) -> std::cmp::Ordering { (Value::Null, _) => std::cmp::Ordering::Greater, (_, Value::Null) => std::cmp::Ordering::Less, (Value::Integer(x), Value::Integer(y)) => x.cmp(y), - (Value::Float(x), Value::Float(y)) => x.partial_cmp(y).unwrap_or(std::cmp::Ordering::Equal), - (Value::Integer(x), Value::Float(y)) => (*x as f64) - .partial_cmp(y) - .unwrap_or(std::cmp::Ordering::Equal), - (Value::Float(x), Value::Integer(y)) => x - .partial_cmp(&(*y as f64)) - .unwrap_or(std::cmp::Ordering::Equal), + // NaN sorts above every number and equals NaN, so the order is total. + (Value::Float(x), Value::Float(y)) => nodedb_query::numeric_cmp::cmp_f64(*x, *y), + // Exact: an integer or decimal never rounds through f64. Every + // pair here has a numeric reading, so `numeric_order` is `Some`. + ( + Value::Integer(_) | Value::Float(_) | Value::Decimal(_), + Value::Integer(_) | Value::Float(_) | Value::Decimal(_), + ) => nodedb_query::value_ops::numeric_order(a, b).unwrap_or(std::cmp::Ordering::Equal), (Value::String(x), Value::String(y)) => x.cmp(y), (Value::Bool(x), Value::Bool(y)) => x.cmp(y), (Value::DateTime(x), Value::DateTime(y)) @@ -102,3 +103,30 @@ fn compare_values(a: &Value, b: &Value) -> std::cmp::Ordering { _ => format!("{a:?}").cmp(&format!("{b:?}")), } } + +#[cfg(test)] +mod tests { + use super::*; + use std::cmp::Ordering; + + #[test] + fn numbers_compare_exactly_across_types() { + let above = Value::Integer(9_007_199_254_740_993); + let float = Value::Float(9_007_199_254_740_992.0); + assert_eq!(compare_values(&above, &float), Ordering::Greater); + assert_eq!(compare_values(&float, &above), Ordering::Less); + let u_max = Value::from_u64(u64::MAX); + assert_eq!( + compare_values(&u_max, &Value::Integer(i64::MAX)), + Ordering::Greater + ); + assert_eq!( + compare_values(&Value::from_u64(u64::MAX - 1), &u_max), + Ordering::Less + ); + assert_eq!( + compare_values(&Value::Integer(2), &Value::Float(2.5)), + Ordering::Less + ); + } +} diff --git a/nodedb/src/data/executor/handlers/document/sort/compare.rs b/nodedb/src/data/executor/handlers/document/sort/compare.rs index 02f700758..5959a2e47 100644 --- a/nodedb/src/data/executor/handlers/document/sort/compare.rs +++ b/nodedb/src/data/executor/handlers/document/sort/compare.rs @@ -16,12 +16,92 @@ //! in a sort key fails the statement with SQLSTATE `22012` rather than //! returning the rows in storage order under a sort the client asked for. +use std::str::FromStr; + use nodedb_physical::physical_plan::SortKeySpec; use nodedb_query::{compare_json, eval_expr_on_json, msgpack_scan}; +use nodedb_types::columnar::{ColumnType, StrictSchema}; +use rust_decimal::Decimal; /// One row's evaluated sort keys, positionally aligned with the ORDER BY list. pub(in crate::data::executor) type SortValues = Vec; +/// Per ORDER BY key, whether it is a bare column `schema` declares `DECIMAL`. +/// +/// A `DECIMAL` cell is text, or a `uint64` past `i64::MAX`, so the generic +/// comparators order such a column by type rank and by bytes. A flagged key +/// compares by value instead. A collection with no strict schema flags none. +pub(in crate::data::executor) fn decimal_sort_keys( + sort_keys: &[SortKeySpec], + schema: Option<&StrictSchema>, +) -> Vec { + let Some(schema) = schema else { + return Vec::new(); + }; + sort_keys + .iter() + .map(|key| { + key.as_column().is_some_and(|field| { + schema.columns.iter().any(|column| { + column.name == field && matches!(column.column_type, ColumnType::Decimal { .. }) + }) + }) + }) + .collect() +} + +/// Whether key `idx` is flagged `DECIMAL` in `decimal_keys`. +pub(in crate::data::executor) fn is_decimal_key(decimal_keys: &[bool], idx: usize) -> bool { + decimal_keys.get(idx).copied().unwrap_or(false) +} + +/// The decimal an evaluated sort value stands for: a decimal string or a +/// number. `None` for any other value. +fn json_decimal(value: &serde_json::Value) -> Option { + match value { + serde_json::Value::String(text) => Decimal::from_str(text.trim()) + .or_else(|_| Decimal::from_scientific(text.trim())) + .ok(), + serde_json::Value::Number(number) => { + if let Some(i) = number.as_i64() { + Some(Decimal::from(i)) + } else if let Some(u) = number.as_u64() { + Some(Decimal::from(u)) + } else { + number.as_f64().and_then(|f| Decimal::try_from(f).ok()) + } + } + _ => None, + } +} + +/// Order two evaluated sort values, by value when the key is `DECIMAL`. +fn compare_key_values( + a: &serde_json::Value, + b: &serde_json::Value, + decimal: bool, +) -> std::cmp::Ordering { + if decimal && let (Some(da), Some(db)) = (json_decimal(a), json_decimal(b)) { + return da.cmp(&db); + } + compare_json(a, b) +} + +/// Order two stored cells, by value when the key is `DECIMAL`. +pub(in crate::data::executor) fn compare_cells( + a_bytes: &[u8], + a_range: (usize, usize), + b_bytes: &[u8], + b_range: (usize, usize), + decimal: bool, +) -> std::cmp::Ordering { + if decimal { + msgpack_scan::compare_decimal_field_bytes(a_bytes, a_range, b_bytes, b_range) + } else { + msgpack_scan::compare_field_bytes(a_bytes, a_range, b_bytes, b_range) + } +} + /// True when every key is a bare stored column, so the zero-decode binary /// comparator can be used. pub(in crate::data::executor) fn all_column_keys(sort_keys: &[SortKeySpec]) -> bool { @@ -52,11 +132,13 @@ pub(in crate::data::executor) fn eval_sort_values( .collect() } -/// Compare two rows by their pre-evaluated sort values. +/// Compare two rows by their pre-evaluated sort values. `decimal_keys` flags +/// the keys [`decimal_sort_keys`] found `DECIMAL`. pub(in crate::data::executor) fn compare_sort_values( a: &[serde_json::Value], b: &[serde_json::Value], sort_keys: &[SortKeySpec], + decimal_keys: &[bool], ) -> std::cmp::Ordering { for (idx, key) in sort_keys.iter().enumerate() { let (Some(va), Some(vb)) = (a.get(idx), b.get(idx)) else { @@ -66,7 +148,11 @@ pub(in crate::data::executor) fn compare_sort_values( // is never flipped by the sort direction. let ord = match key.order_nulls(va.is_null(), vb.is_null()) { Some(ord) => ord, - None => key.direct(compare_json(va, vb)), + None => key.direct(compare_key_values( + va, + vb, + is_decimal_key(decimal_keys, idx), + )), }; if ord != std::cmp::Ordering::Equal { return ord; @@ -78,13 +164,15 @@ pub(in crate::data::executor) fn compare_sort_values( /// Compare two raw msgpack documents by column-only sort keys. /// /// Uses binary field extraction — no decode. Shared by the in-memory sort and -/// the external merge so both order rows identically. +/// the external merge so both order rows identically. `decimal_keys` flags +/// the keys [`decimal_sort_keys`] found `DECIMAL`. pub(in crate::data::executor) fn compare_docs_by_keys_binary( a_bytes: &[u8], b_bytes: &[u8], sort_keys: &[SortKeySpec], + decimal_keys: &[bool], ) -> std::cmp::Ordering { - for key in sort_keys { + for (idx, key) in sort_keys.iter().enumerate() { // Column-only callers guarantee `as_column`; a computed key has no // field to extract and compares equal here, which is why the caller // routes those rows through `compare_sort_values` instead. @@ -102,9 +190,13 @@ pub(in crate::data::executor) fn compare_docs_by_keys_binary( ) { Some(ord) => ord, None => match (a_range, b_range) { - (Some(ar), Some(br)) => { - key.direct(msgpack_scan::compare_field_bytes(a_bytes, ar, b_bytes, br)) - } + (Some(ar), Some(br)) => key.direct(compare_cells( + a_bytes, + ar, + b_bytes, + br, + is_decimal_key(decimal_keys, idx), + )), _ => std::cmp::Ordering::Equal, }, }; diff --git a/nodedb/src/data/executor/handlers/document/sort/external.rs b/nodedb/src/data/executor/handlers/document/sort/external.rs index 8e6a0cc64..7e587b0d6 100644 --- a/nodedb/src/data/executor/handlers/document/sort/external.rs +++ b/nodedb/src/data/executor/handlers/document/sort/external.rs @@ -44,10 +44,13 @@ impl CoreLoop { /// not by tempfile auto-delete. The merge reads each run back incrementally /// via [`UringSeqReader`] — one row at a time — so peak read memory is one /// refill buffer per run, not the whole run. + /// + /// `decimal_keys` flags the keys `decimal_sort_keys` found `DECIMAL`. pub(in crate::data::executor) fn external_sort( &self, rows: Vec<(String, Vec)>, sort_keys: &[SortKeySpec], + decimal_keys: &[bool], output_limit: usize, ) -> crate::Result)>> { // Spill directory for the named sort run files. `create_dir_all` is a @@ -73,7 +76,7 @@ impl CoreLoop { for (run_idx, chunk) in rows.chunks(self.query_tuning.sort_run_size).enumerate() { let mut run: Vec<(String, Vec)> = chunk.to_vec(); - sort_rows(&mut run, sort_keys)?; + sort_rows(&mut run, sort_keys, decimal_keys)?; // Build the framed run into one buffer and write it in a single // pass. Writing each tiny frame field separately would be hundreds @@ -121,6 +124,7 @@ impl CoreLoop { record, run_idx: reader.run_idx, sort_keys: sort_keys.to_vec(), + decimal_keys: decimal_keys.to_vec(), })); } } @@ -137,6 +141,7 @@ impl CoreLoop { record: next, run_idx, sort_keys: sort_keys.to_vec(), + decimal_keys: decimal_keys.to_vec(), })); } } @@ -284,6 +289,7 @@ struct MergeEntry { record: SortRecord, run_idx: usize, sort_keys: Vec, + decimal_keys: Vec, } impl PartialEq for MergeEntry { @@ -305,9 +311,19 @@ impl Ord for MergeEntry { // Rows spilled with evaluated keys are merged by those keys; a // column-only sort carries none and compares straight from the bytes. if self.record.keys.is_empty() && other.record.keys.is_empty() { - compare_docs_by_keys_binary(&self.record.doc, &other.record.doc, &self.sort_keys) + compare_docs_by_keys_binary( + &self.record.doc, + &other.record.doc, + &self.sort_keys, + &self.decimal_keys, + ) } else { - compare_sort_values(&self.record.keys, &other.record.keys, &self.sort_keys) + compare_sort_values( + &self.record.keys, + &other.record.keys, + &self.sort_keys, + &self.decimal_keys, + ) } } } @@ -372,6 +388,7 @@ mod tests { record: r, run_idx: reader.run_idx, sort_keys: sort_keys.clone(), + decimal_keys: Vec::new(), })); } } @@ -385,6 +402,7 @@ mod tests { record: next, run_idx, sort_keys: sort_keys.clone(), + decimal_keys: Vec::new(), })); } } diff --git a/nodedb/src/data/executor/handlers/document/sort/in_memory.rs b/nodedb/src/data/executor/handlers/document/sort/in_memory.rs index 1f5cd680f..823797ce2 100644 --- a/nodedb/src/data/executor/handlers/document/sort/in_memory.rs +++ b/nodedb/src/data/executor/handlers/document/sort/in_memory.rs @@ -6,60 +6,30 @@ use nodedb_physical::physical_plan::SortKeySpec; use nodedb_query::msgpack_scan; use super::compare::{ - SortValues, all_column_keys, compare_sort_values, eval_sort_values, is_null_range, + SortValues, all_column_keys, compare_cells, compare_sort_values, eval_sort_values, + is_decimal_key, is_null_range, }; /// Pre-extracted sort key offsets for a single row. /// Each entry is `Option<(usize, usize)>` — byte range of the sort key value. type SortKeyOffsets = Vec>; +/// Sort stored rows by the ORDER BY terms. `decimal_keys` flags the keys +/// `decimal_sort_keys` found `DECIMAL`, and is empty for a collection with +/// no strict schema. pub(in crate::data::executor) fn sort_rows( rows: &mut [(String, Vec)], sort_keys: &[SortKeySpec], + decimal_keys: &[bool], ) -> crate::Result<()> { if sort_keys.is_empty() { return Ok(()); } if all_column_keys(sort_keys) { - return sort_rows_by_column(rows, sort_keys); + return sort_rows_by_column(rows, sort_keys, decimal_keys); } - sort_rows_by_expression(rows, sort_keys) -} - -/// Sort decoded document rows by the ORDER BY terms, with the same ordering -/// rules [`sort_rows`] applies to stored rows. -/// -/// A scan with window functions sorts after the window pass, so ORDER BY can -/// name a window alias; its rows are decoded by then. Each row is encoded -/// once under its position, sorted by [`sort_rows`], and taken back in the -/// sorted order. -pub(in crate::data::executor) fn sort_decoded_rows( - rows: Vec<(String, serde_json::Value)>, - sort_keys: &[SortKeySpec], -) -> crate::Result> { - if sort_keys.is_empty() { - return Ok(rows); - } - let mut keyed: Vec<(String, Vec)> = rows - .iter() - .enumerate() - .map(|(i, (_, doc))| (i.to_string(), nodedb_types::json_to_msgpack_or_empty(doc))) - .collect(); - sort_rows(&mut keyed, sort_keys)?; - let mut slots: Vec> = rows.into_iter().map(Some).collect(); - keyed - .iter() - .map(|(position, _)| { - position - .parse::() - .ok() - .and_then(|i| slots.get_mut(i).and_then(Option::take)) - .ok_or_else(|| crate::Error::Internal { - detail: format!("sort_decoded_rows: row position {position} is not a live row"), - }) - }) - .collect() + sort_rows_by_expression(rows, sort_keys, decimal_keys) } /// Zero-decode path: every key names a stored field, so ordering is decided @@ -67,6 +37,7 @@ pub(in crate::data::executor) fn sort_decoded_rows( fn sort_rows_by_column( rows: &mut [(String, Vec)], sort_keys: &[SortKeySpec], + decimal_keys: &[bool], ) -> crate::Result<()> { // Pre-extract sort key offsets for all rows — one scan per row instead // of O(N log N) scans during comparisons. @@ -91,6 +62,7 @@ fn sort_rows_by_column( &rows[bi].1, &key_offsets[bi], sort_keys, + decimal_keys, ) }); @@ -105,6 +77,7 @@ fn sort_rows_by_column( fn sort_rows_by_expression( rows: &mut [(String, Vec)], sort_keys: &[SortKeySpec], + decimal_keys: &[bool], ) -> crate::Result<()> { let values: Vec = rows .iter() @@ -112,7 +85,8 @@ fn sort_rows_by_expression( .collect::>>()?; let mut indices: Vec = (0..rows.len()).collect(); - indices.sort_by(|&ai, &bi| compare_sort_values(&values[ai], &values[bi], sort_keys)); + indices + .sort_by(|&ai, &bi| compare_sort_values(&values[ai], &values[bi], sort_keys, decimal_keys)); drop(values); apply_permutation(rows, indices) @@ -125,6 +99,7 @@ fn compare_with_preextracted( b_bytes: &[u8], b_offsets: &[Option<(usize, usize)>], sort_keys: &[SortKeySpec], + decimal_keys: &[bool], ) -> std::cmp::Ordering { for (i, key) in sort_keys.iter().enumerate() { let ordered = match key.order_nulls( @@ -133,9 +108,13 @@ fn compare_with_preextracted( ) { Some(ord) => ord, None => match (a_offsets[i], b_offsets[i]) { - (Some(ar), Some(br)) => { - key.direct(msgpack_scan::compare_field_bytes(a_bytes, ar, b_bytes, br)) - } + (Some(ar), Some(br)) => key.direct(compare_cells( + a_bytes, + ar, + b_bytes, + br, + is_decimal_key(decimal_keys, i), + )), _ => std::cmp::Ordering::Equal, }, }; @@ -207,30 +186,51 @@ mod tests { #[test] fn sort_by_int_field_asc() { let mut rows = rows_with_vals(&[("a", 30), ("b", 10), ("c", 20)]); - sort_rows(&mut rows, &[col("val", true)]).expect("sort_rows failed"); + sort_rows(&mut rows, &[col("val", true)], &[]).expect("sort_rows failed"); let order: Vec<&str> = rows.iter().map(|(id, _)| id.as_str()).collect(); assert_eq!(order, vec!["b", "c", "a"], "ASC by val: 10, 20, 30"); } - /// Decoded rows sort on a field the window pass appended, and keep their - /// document ids. + /// A `DECIMAL` key orders its text and `uint64` cells by value, on both + /// the column and the expression comparator. #[test] - fn sort_decoded_rows_orders_by_an_appended_column() { - let rows = vec![ - ("a".to_string(), serde_json::json!({"id": "a", "rn": 3})), - ("b".to_string(), serde_json::json!({"id": "b", "rn": 1})), - ("c".to_string(), serde_json::json!({"id": "c", "rn": 2})), - ]; - let sorted = sort_decoded_rows(rows, &[col("rn", true)]).expect("sort_decoded_rows"); - let order: Vec<&str> = sorted.iter().map(|(id, _)| id.as_str()).collect(); - assert_eq!(order, vec!["b", "c", "a"]); - assert_eq!(sorted[0].1, serde_json::json!({"id": "b", "rn": 1})); + fn decimal_key_orders_text_and_uint64_cells_by_value() { + let rows = || -> Vec<(String, Vec)> { + [ + ("max", serde_json::json!(u64::MAX)), + ("mid", serde_json::json!(9_223_372_036_854_775_808_u64)), + ("small", serde_json::json!("5")), + ("ten", serde_json::json!("10")), + ] + .into_iter() + .map(|(id, v)| (id.to_string(), encode(&serde_json::json!({ "v": v })))) + .collect() + }; + let mut by_column = rows(); + sort_rows(&mut by_column, &[col("v", true)], &[true]).expect("sort_rows failed"); + let order: Vec<&str> = by_column.iter().map(|(id, _)| id.as_str()).collect(); + assert_eq!(order, vec!["small", "ten", "mid", "max"]); + + let key = SortKeySpec { + expr: nodedb_query::SqlExpr::Column("v".into()), + ascending: true, + nulls_first: false, + }; + let computed = SortKeySpec { + expr: nodedb_query::SqlExpr::Literal(nodedb_types::Value::Integer(1)), + ascending: true, + nulls_first: false, + }; + let mut by_expression = rows(); + sort_rows(&mut by_expression, &[key, computed], &[true, false]).expect("sort_rows failed"); + let order: Vec<&str> = by_expression.iter().map(|(id, _)| id.as_str()).collect(); + assert_eq!(order, vec!["small", "ten", "mid", "max"]); } #[test] fn sort_by_int_field_desc() { let mut rows = rows_with_vals(&[("a", 30), ("b", 10), ("c", 20)]); - sort_rows(&mut rows, &[col("val", false)]).expect("sort_rows failed"); + sort_rows(&mut rows, &[col("val", false)], &[]).expect("sort_rows failed"); let order: Vec<&str> = rows.iter().map(|(id, _)| id.as_str()).collect(); assert_eq!(order, vec!["a", "c", "b"], "DESC by val: 30, 20, 10"); } @@ -242,7 +242,7 @@ mod tests { ("b".into(), encode(&serde_json::json!({"name": "alpha"}))), ("c".into(), encode(&serde_json::json!({"name": "bravo"}))), ]; - sort_rows(&mut rows, &[col("name", true)]).expect("sort_rows failed"); + sort_rows(&mut rows, &[col("name", true)], &[]).expect("sort_rows failed"); let order: Vec<&str> = rows.iter().map(|(id, _)| id.as_str()).collect(); assert_eq!(order, vec!["b", "c", "a"]); } @@ -264,7 +264,7 @@ mod tests { ascending: true, nulls_first: false, }; - sort_rows(&mut rows, &[key]).expect("sort_rows failed"); + sort_rows(&mut rows, &[key], &[]).expect("sort_rows failed"); let order: Vec<&str> = rows.iter().map(|(id, _)| id.as_str()).collect(); // 100/20 = 5, 100/5 = 20, 100/2 = 50. assert_eq!(order, vec!["b", "c", "a"]); @@ -286,7 +286,8 @@ mod tests { ascending: true, nulls_first: false, }; - sort_rows(&mut rows, &[key]).expect_err("zero divisor in a sort key must fail the sort"); + sort_rows(&mut rows, &[key], &[]) + .expect_err("zero divisor in a sort key must fail the sort"); } #[test] diff --git a/nodedb/src/data/executor/handlers/document/sort/mod.rs b/nodedb/src/data/executor/handlers/document/sort/mod.rs index 545f85efd..8a82142e5 100644 --- a/nodedb/src/data/executor/handlers/document/sort/mod.rs +++ b/nodedb/src/data/executor/handlers/document/sort/mod.rs @@ -6,4 +6,5 @@ pub(in crate::data::executor) mod compare; pub(in crate::data::executor) mod external; pub(in crate::data::executor) mod in_memory; -pub(in crate::data::executor) use in_memory::{sort_decoded_rows, sort_rows}; +pub(in crate::data::executor) use compare::decimal_sort_keys; +pub(in crate::data::executor) use in_memory::sort_rows; diff --git a/nodedb/src/data/executor/handlers/timeseries/ingest_formats.rs b/nodedb/src/data/executor/handlers/timeseries/ingest_formats.rs index 813a64957..3ee7196d8 100644 --- a/nodedb/src/data/executor/handlers/timeseries/ingest_formats.rs +++ b/nodedb/src/data/executor/handlers/timeseries/ingest_formats.rs @@ -227,7 +227,17 @@ impl CoreLoop { .declared_ts_time_key(task.request.database_id, tid, collection) .map(str::to_string); - let ilp_buf = normalize::json_rows_to_ilp(&rows, measurement, time_key.as_deref()); + let ilp_buf = match normalize::json_rows_to_ilp(&rows, measurement, time_key.as_deref()) { + Ok(buf) => buf, + Err(error) => { + return self.response_error( + task, + ErrorCode::RejectedPrevalidation { + reason: format!("timeseries ingest: {error}"), + }, + ); + } + }; if ilp_buf.is_empty() { return self.response_error( diff --git a/nodedb/src/data/executor/handlers/timeseries/msgpack_decode.rs b/nodedb/src/data/executor/handlers/timeseries/msgpack_decode.rs index 83976ff9f..95900e760 100644 --- a/nodedb/src/data/executor/handlers/timeseries/msgpack_decode.rs +++ b/nodedb/src/data/executor/handlers/timeseries/msgpack_decode.rs @@ -13,6 +13,8 @@ use nodedb_types::read_instant; pub(super) enum MsgpackValue { Int(i64), + /// A `uint64` above `i64::MAX`. Every smaller integer is `Int`. + UInt(u64), Float(f64), Str(String), Bool(bool), @@ -169,10 +171,10 @@ fn read_value(buf: &[u8], pos: &mut usize) -> Result let v = read_be_u32(buf, pos)?; Ok(MsgpackValue::Int(v as i64)) } - // uint64 + // uint64: an `i64` when it fits, else the unsigned value itself. 0xCF => { - let bytes = read_bytes::<8>(buf, pos)?; - Ok(MsgpackValue::Int(u64::from_be_bytes(bytes) as i64)) + let v = u64::from_be_bytes(read_bytes::<8>(buf, pos)?); + Ok(i64::try_from(v).map_or(MsgpackValue::UInt(v), MsgpackValue::Int)) } // fixext8: an instant of either kind. Any other ext type is // unsupported, like every other ext marker. diff --git a/nodedb/src/data/executor/handlers/timeseries/normalize.rs b/nodedb/src/data/executor/handlers/timeseries/normalize.rs index a16ce4d6c..21b06eabb 100644 --- a/nodedb/src/data/executor/handlers/timeseries/normalize.rs +++ b/nodedb/src/data/executor/handlers/timeseries/normalize.rs @@ -10,6 +10,7 @@ use nodedb_types::datetime::NdbDateTime; use sonic_rs::{JsonContainerTrait, JsonValueTrait}; use super::msgpack_decode::MsgpackValue; +use crate::data::executor::strict_format::float_to_i64; use crate::engine::timeseries::ilp::{self, IlpError}; /// Nanoseconds per millisecond — line protocol timestamps are nanoseconds, the @@ -104,6 +105,9 @@ pub(in crate::data::executor) fn msgpack_rows_to_ilp( match val { MsgpackValue::Float(f) => fields.push(format!("{key}={f}")), MsgpackValue::Int(n) => fields.push(format!("{key}={n}i")), + // An unsigned field keeps its exact number. A column whose + // type cannot hold it refuses the line at ingest. + MsgpackValue::UInt(u) => fields.push(format!("{key}={u}u")), MsgpackValue::Str(s) => { // Recover the numeric type `SqlValue::Decimal` encoded as a // string, so schema inference picks Float64/Int64, not Symbol. @@ -157,8 +161,9 @@ fn time_column_nanos( .checked_mul(NANOS_PER_MILLI) .map(Some) .ok_or_else(|| invalid(&n.to_string())), - MsgpackValue::Float(f) => (*f as i64) - .checked_mul(NANOS_PER_MILLI) + // Past `i64::MAX` milliseconds: no timestamp holds it. + MsgpackValue::UInt(u) => Err(invalid(&u.to_string())), + MsgpackValue::Float(f) => float_millis_to_nanos(*f) .map(Some) .ok_or_else(|| invalid(&f.to_string())), // A typed instant of either kind: the stored time column is the @@ -171,15 +176,69 @@ fn time_column_nanos( } } +/// A float count of epoch milliseconds as epoch nanoseconds. A fractional +/// millisecond is kept to the nanosecond, the line timestamp unit. A part +/// below one nanosecond rounds half to even, as PostgreSQL's +/// `to_timestamp(double precision)` rounds below its own unit. `None` for +/// NaN, an infinity, or a value past the nanosecond range. +fn float_millis_to_nanos(ms: f64) -> Option { + float_to_i64((ms * NANOS_PER_MILLI as f64).round_ties_even()) +} + +/// The nanosecond timestamp a JSON time-column cell denotes, under the same +/// rules as [`time_column_nanos`]: `None` for JSON null, an error for a value +/// no timestamp can carry. +fn json_time_column_nanos( + val: &sonic_rs::Value, + column: &str, + line_number: usize, +) -> Result, IlpError> { + let invalid = |detail: &str| { + IlpError::new( + line_number, + &format!("{column}={detail}"), + 0..0, + ilp::IlpErrorKind::InvalidTimestamp, + ) + }; + if val.is_null() { + return Ok(None); + } + if let Some(s) = val.as_str() { + return parse_ts_string_to_nanos(s) + .map(Some) + .ok_or_else(|| invalid(&format!("\"{s}\""))); + } + if let Some(n) = val.as_i64() { + return n + .checked_mul(NANOS_PER_MILLI) + .map(Some) + .ok_or_else(|| invalid(&n.to_string())); + } + if let Some(f) = val.as_f64() { + return float_millis_to_nanos(f) + .map(Some) + .ok_or_else(|| invalid(&f.to_string())); + } + let detail = match val.as_bool() { + Some(b) => b.to_string(), + None if val.is_array() => "".to_string(), + None => "".to_string(), + }; + Err(invalid(&detail)) +} + /// Normalize decoded JSON rows into line protocol. The JSON value model /// carries no decimal-as-string case, so no string is re-parsed as a number. +/// A time column that holds a value the line cannot carry is an error, as it +/// is for [`msgpack_rows_to_ilp`]. pub(in crate::data::executor) fn json_rows_to_ilp( rows: &sonic_rs::Array, measurement: &str, time_key: Option<&str>, -) -> String { +) -> Result { let mut ilp_buf = String::new(); - for row_val in rows.iter() { + for (line_number, row_val) in rows.iter().enumerate() { let Some(obj) = row_val.as_object() else { continue; }; @@ -189,13 +248,7 @@ pub(in crate::data::executor) fn json_rows_to_ilp( for (key, val) in obj.iter() { if is_time_column(key, time_key) { - if let Some(s) = val.as_str() { - timestamp_ns = parse_ts_string_to_nanos(s); - } else if let Some(n) = val.as_i64() { - timestamp_ns = Some(n * NANOS_PER_MILLI); - } else if let Some(f) = val.as_f64() { - timestamp_ns = Some(f as i64 * NANOS_PER_MILLI); - } + timestamp_ns = json_time_column_nanos(val, key, line_number + 1)?; continue; } @@ -212,7 +265,7 @@ pub(in crate::data::executor) fn json_rows_to_ilp( push_line(&mut ilp_buf, measurement, &fields, timestamp_ns); } - ilp_buf + Ok(ilp_buf) } /// Split `batch` into lines, giving every timestamp-less line the batch's @@ -309,6 +362,53 @@ mod tests { } } + /// A float time column keeps its fractional millisecond to the + /// nanosecond; NaN, an infinity, and a value past the nanosecond range + /// are refused instead of being stored as the epoch or a wrapped time. + #[test] + fn a_float_time_key_converts_exactly_or_is_refused() { + let rows = vec![vec![ + ("ts".to_string(), MsgpackValue::Float(1.5)), + ("v".to_string(), MsgpackValue::Int(1)), + ]]; + assert_eq!( + msgpack_rows_to_ilp(&rows, "m", Some("ts")).expect("ilp"), + "m v=1i 1500000\n" + ); + for bad in [f64::NAN, f64::INFINITY, f64::NEG_INFINITY, 1e300] { + let rows = vec![vec![ + ("ts".to_string(), MsgpackValue::Float(bad)), + ("v".to_string(), MsgpackValue::Int(1)), + ]]; + let err = msgpack_rows_to_ilp(&rows, "m", Some("ts")).expect_err("refused"); + assert_eq!(err.kind, ilp::IlpErrorKind::InvalidTimestamp); + } + } + + /// The JSON path applies the msgpack time-column rules: a fractional + /// millisecond is kept, and an overflowing or unparseable value is an + /// error rather than a clock stamp or a wrapped time. + #[test] + fn a_json_time_key_follows_the_msgpack_rules() { + let rows: sonic_rs::Array = + sonic_rs::from_str(r#"[{"ts": 2.25, "v": "x"}]"#).expect("json"); + assert_eq!( + json_rows_to_ilp(&rows, "m", Some("ts")).expect("ilp"), + "m v=\"x\" 2250000\n" + ); + for bad in [ + r#"[{"ts": 9223372036854775807, "v": "x"}]"#, + r#"[{"ts": 1e300, "v": "x"}]"#, + r#"[{"ts": "not a date", "v": "x"}]"#, + r#"[{"ts": true, "v": "x"}]"#, + ] { + let rows: sonic_rs::Array = sonic_rs::from_str(bad).expect("json"); + let err = json_rows_to_ilp(&rows, "m", Some("ts")).expect_err(bad); + assert_eq!(err.kind, ilp::IlpErrorKind::InvalidTimestamp); + assert!(err.raw.starts_with("ts="), "{err:?}"); + } + } + /// A NULL time column carries no timestamp, so the ingest clock stamps /// the row. #[test] diff --git a/nodedb/src/data/executor/handlers/timeseries/resolve_ingest.rs b/nodedb/src/data/executor/handlers/timeseries/resolve_ingest.rs index 7ec65474e..9c64b910f 100644 --- a/nodedb/src/data/executor/handlers/timeseries/resolve_ingest.rs +++ b/nodedb/src/data/executor/handlers/timeseries/resolve_ingest.rs @@ -250,7 +250,11 @@ impl CoreLoop { reason: format!("timeseries resolve: JSON parse error: {error}"), } })?; - Ok(normalize::json_rows_to_ilp(&rows, measurement, time_key)) + normalize::json_rows_to_ilp(&rows, measurement, time_key).map_err(|error| { + ErrorCode::RejectedPrevalidation { + reason: format!("timeseries resolve: {error}"), + } + }) } other => Err(ErrorCode::Internal { detail: format!("timeseries resolve: unknown ingest format: {other}"), diff --git a/nodedb/src/data/executor/handlers/timeseries/sort.rs b/nodedb/src/data/executor/handlers/timeseries/sort.rs index 3284d1268..a23e3c2c7 100644 --- a/nodedb/src/data/executor/handlers/timeseries/sort.rs +++ b/nodedb/src/data/executor/handlers/timeseries/sort.rs @@ -114,8 +114,12 @@ fn compare_values(a: Option<&rmpv::Value>, b: Option<&rmpv::Value>) -> Ordering (None, Some(_)) => return Ordering::Greater, (Some(_), None) => return Ordering::Less, }; - if let (Some(x), Some(y)) = (as_f64(a), as_f64(b)) { - return x.partial_cmp(&y).unwrap_or(Ordering::Equal); + if let (Some(x), Some(y)) = (as_number(a), as_number(b)) { + // Exact: integers never round through f64. NaN sorts above every + // number and equals NaN. Both sides are numbers, so an order exists. + if let Some(order) = nodedb_query::value_ops::numeric_order(&x, &y) { + return order; + } } match (a, b) { (rmpv::Value::String(x), rmpv::Value::String(y)) => { @@ -138,12 +142,17 @@ fn compare_values(a: Option<&rmpv::Value>, b: Option<&rmpv::Value>) -> Ordering } /// Numeric view of a value, so an integer column and a float column compare -/// against each other the way SQL expects. -fn as_f64(value: &rmpv::Value) -> Option { +/// against each other the way SQL expects. An integer stays exact: a `u64` +/// above `i64::MAX` is an integral `Decimal`. +fn as_number(value: &rmpv::Value) -> Option { match value { - rmpv::Value::Integer(n) => n.as_i64().map(|i| i as f64).or_else(|| n.as_f64()), - rmpv::Value::F32(f) => Some(*f as f64), - rmpv::Value::F64(f) => Some(*f), + rmpv::Value::Integer(n) => match (n.as_i64(), n.as_u64()) { + (Some(i), _) => Some(nodedb_types::Value::Integer(i)), + (None, Some(u)) => Some(nodedb_types::Value::from_u64(u)), + (None, None) => None, + }, + rmpv::Value::F32(f) => Some(nodedb_types::Value::Float(f64::from(*f))), + rmpv::Value::F64(f) => Some(nodedb_types::Value::Float(*f)), _ => None, } } @@ -233,11 +242,42 @@ mod tests { row(&[("v", rmpv::Value::F64(1.5))]), ]; sort_rows(&mut rows, &[SortKeySpec::column("v", true)]).expect("sort"); - let vs: Vec = rows + let vs: Vec = rows + .iter() + .map(|r| as_number(field_of(r, "v").unwrap()).unwrap()) + .collect(); + assert_eq!( + vs, + vec![ + nodedb_types::Value::Float(1.5), + nodedb_types::Value::Integer(2), + nodedb_types::Value::Float(2.5) + ] + ); + } + + #[test] + fn integers_past_two_pow_53_and_u64_sort_exactly() { + let mut rows = vec![ + row(&[("v", rmpv::Value::Integer(9_007_199_254_740_993_i64.into()))]), + row(&[("v", rmpv::Value::Integer(u64::MAX.into()))]), + row(&[("v", rmpv::Value::Integer(9_007_199_254_740_992_i64.into()))]), + row(&[("v", rmpv::Value::F64(9_007_199_254_740_992.0))]), + ]; + sort_rows(&mut rows, &[SortKeySpec::column("v", false)]).expect("sort"); + let vs: Vec = rows .iter() - .map(|r| as_f64(field_of(r, "v").unwrap()).unwrap()) + .map(|r| field_of(r, "v").unwrap().clone()) .collect(); - assert_eq!(vs, vec![1.5, 2.0, 2.5]); + assert_eq!( + vs, + vec![ + rmpv::Value::Integer(u64::MAX.into()), + rmpv::Value::Integer(9_007_199_254_740_993_i64.into()), + rmpv::Value::Integer(9_007_199_254_740_992_i64.into()), + rmpv::Value::F64(9_007_199_254_740_992.0), + ] + ); } #[test] diff --git a/nodedb/src/engine/timeseries/ilp_ingest.rs b/nodedb/src/engine/timeseries/ilp_ingest.rs index d29fef553..9a6d6bec2 100644 --- a/nodedb/src/engine/timeseries/ilp_ingest.rs +++ b/nodedb/src/engine/timeseries/ilp_ingest.rs @@ -169,6 +169,8 @@ pub fn ingest_batch_with_lvc(args: IngestBatchArgs<'_, '_>) -> IngestBatchOutcom /// - A `Symbol` column takes a tag or a string field. A numeric or boolean /// field conflicts. /// - A `Timestamp` column takes a numeric field or a datetime string. +/// - An `Int64` or `Timestamp` column refuses an unsigned field past +/// `i64::MAX`. pub fn line_type_conflict( schema: &super::columnar_memtable::ColumnarSchema, line: &IlpLine<'_>, @@ -193,6 +195,16 @@ pub fn line_type_conflict( let Some(column_type) = column_type(key.as_ref()) else { continue; }; + // An `Int64` or `Timestamp` cell is an `i64`: an unsigned field past + // `i64::MAX` would wrap to a negative number. + if let (ColumnType::Int64 | ColumnType::Timestamp(_), FieldValue::UInt(u)) = + (column_type, value) + && i64::try_from(*u).is_err() + { + return Some(format!( + "value {u} of '{key}' is out of range for a {column_type:?} column" + )); + } let fits = match column_type { ColumnType::Float64 | ColumnType::Int64 => !matches!(value, FieldValue::Str(_)), ColumnType::Symbol => matches!(value, FieldValue::Str(_)), @@ -460,6 +472,30 @@ mod tests { assert_eq!(schema.columns[4].1, ColumnType::Int64); // count } + /// An unsigned field past `i64::MAX` is refused by an `Int64` column + /// instead of wrapping to a negative number. One in range is stored. + #[test] + fn unsigned_field_past_i64_max_is_refused_not_wrapped() { + let schema = infer_schema( + &parse_batch("m count=1i 1000000000") + .expect("valid ILP batch") + .into_lines(), + ); + let lines = parse_batch( + "m count=18446744073709551615u 1000000000\n\ + m count=9223372036854775807u 2000000000", + ) + .expect("valid ILP batch") + .into_lines(); + let conflict = line_type_conflict(&schema, &lines[0]).expect("out of range"); + assert!(conflict.contains("18446744073709551615"), "{conflict}"); + assert_eq!(line_type_conflict(&schema, &lines[1]), None); + + let mut mt = ColumnarMemtable::new(schema, default_config()); + let mut catalog = SeriesCatalog::new(); + assert_eq!(ingest_batch(&mut mt, &lines, &mut catalog, 0), (1, 1)); + } + #[test] fn bitemporal_ingest_stamps_reserved_columns() { // Late-arriving IoT backfill: an ILP line with a user-provided diff --git a/nodedb/tests/wire/cases/aggregate_integer_exactness.rs b/nodedb/tests/wire/cases/aggregate_integer_exactness.rs new file mode 100644 index 000000000..4b2588788 --- /dev/null +++ b/nodedb/tests/wire/cases/aggregate_integer_exactness.rs @@ -0,0 +1,211 @@ +// SPDX-License-Identifier: BUSL-1.1 + +//! `MIN` / `MAX` / `SUM` / `AVG` over integers stay exact end to end. +//! +//! `2^53 + 1` and `2^53` collapse to one `f64`, as do nanosecond timestamps +//! one tick apart. An aggregate that rounds through `f64` returns the wrong +//! extreme or a wrong total. A SUM past `i64::MAX` returns the exact total, +//! not a wrapped or rounded one. Each engine runs the same checks. + +use crate::harness::TestServer; + +const ABOVE: i64 = 9_007_199_254_740_993; +const AT: i64 = 9_007_199_254_740_992; +const NANOS: [i64; 3] = [ + 1_700_000_000_000_000_002, + 1_700_000_000_000_000_001, + 1_700_000_000_000_000_003, +]; + +/// Insert `values` into `collection.v`, one row each. `ts` is the row +/// index, so a timeseries collection gets distinct time keys. +async fn insert_values(srv: &TestServer, collection: &str, values: &[i64]) { + for (i, v) in values.iter().enumerate() { + srv.exec(&format!( + "INSERT INTO {collection} (id, ts, v) VALUES ('r{i}', {ts}, {v})", + ts = 1_700_000_000_000_i64 + i as i64 + )) + .await + .unwrap(); + } +} + +/// `SELECT MIN(v), MAX(v), SUM(v), AVG(v)` parsed: the three integer cells +/// exactly, AVG as an `f64`. +async fn min_max_sum_avg(srv: &TestServer, collection: &str) -> (i128, i128, i128, f64) { + let rows = srv + .query_rows(&format!( + "SELECT MIN(v), MAX(v), SUM(v), AVG(v) FROM {collection}" + )) + .await + .unwrap(); + assert_eq!(rows.len(), 1, "one aggregate row, got {rows:?}"); + let int = |i: usize| { + rows[0][i].parse::().unwrap_or_else(|_| { + panic!( + "{collection}: cell {i} must be an exact integer, got `{}`", + rows[0][i] + ) + }) + }; + let avg = rows[0][3] + .parse::() + .unwrap_or_else(|_| panic!("{collection}: AVG must be numeric, got `{}`", rows[0][3])); + (int(0), int(1), int(2), avg) +} + +/// Run every exactness check against three fresh collections made by +/// `create(name)`. +async fn check_engine(srv: &TestServer, create: impl Fn(&str) -> String) { + // Integers one apart above 2^53. + srv.exec(&create("big")).await.unwrap(); + insert_values(srv, "big", &[ABOVE, AT]).await; + let (min, max, sum, avg) = min_max_sum_avg(srv, "big").await; + assert_eq!(min, i128::from(AT), "MIN above 2^53"); + assert_eq!(max, i128::from(ABOVE), "MAX above 2^53"); + assert_eq!(sum, i128::from(ABOVE) + i128::from(AT), "SUM above 2^53"); + assert_eq!(avg, AT as f64, "AVG above 2^53"); + + // Nanosecond timestamps one tick apart. + srv.exec(&create("nanos")).await.unwrap(); + insert_values(srv, "nanos", &NANOS).await; + let (min, max, sum, _) = min_max_sum_avg(srv, "nanos").await; + assert_eq!(min, i128::from(NANOS[1]), "MIN of nanosecond timestamps"); + assert_eq!(max, i128::from(NANOS[2]), "MAX of nanosecond timestamps"); + assert_eq!( + sum, + NANOS.iter().map(|&n| i128::from(n)).sum::(), + "SUM of nanosecond timestamps" + ); + + // A total past i64::MAX is exact, never wrapped or rounded. + srv.exec(&create("over")).await.unwrap(); + insert_values(srv, "over", &[i64::MAX, i64::MAX, 2]).await; + let (min, max, sum, _) = min_max_sum_avg(srv, "over").await; + assert_eq!(min, 2, "MIN beside i64::MAX"); + assert_eq!(max, i128::from(i64::MAX), "MAX at i64::MAX"); + assert_eq!( + sum, + 2 * i128::from(i64::MAX) + 2, + "SUM past i64::MAX must be the exact total" + ); +} + +#[tokio::test] +async fn document_schemaless_integer_aggregates_are_exact() { + let srv = TestServer::start().await; + check_engine(&srv, |name| { + format!("CREATE COLLECTION {name} WITH (engine='document_schemaless')") + }) + .await; +} + +#[tokio::test] +async fn columnar_integer_aggregates_are_exact() { + let srv = TestServer::start().await; + check_engine(&srv, |name| { + format!( + "CREATE COLLECTION {name} \ + COLUMNS (id TEXT, ts BIGINT, v BIGINT) \ + WITH (engine='columnar')" + ) + }) + .await; +} + +#[tokio::test] +async fn timeseries_integer_aggregates_are_exact() { + let srv = TestServer::start().await; + check_engine(&srv, |name| { + format!( + "CREATE COLLECTION {name} \ + COLUMNS (ts BIGINT TIME_KEY, id TEXT, v BIGINT) \ + WITH (engine='timeseries')" + ) + }) + .await; +} + +/// A schemaless column mixing integers and floats sums as a float, and +/// MIN / MAX return the original value of the winning row. +#[tokio::test] +async fn document_schemaless_mixed_int_float() { + let srv = TestServer::start().await; + srv.exec("CREATE COLLECTION mixed WITH (engine='document_schemaless')") + .await + .unwrap(); + srv.exec("INSERT INTO mixed (id, v) VALUES ('a', 2), ('b', 0.5), ('c', 9007199254740993)") + .await + .unwrap(); + let rows = srv + .query_rows("SELECT MIN(v), MAX(v), SUM(v) FROM mixed") + .await + .unwrap(); + assert_eq!(rows.len(), 1, "one aggregate row, got {rows:?}"); + assert_eq!( + rows[0][0].parse::().unwrap(), + 0.5, + "MIN keeps the float" + ); + assert_eq!( + rows[0][1].parse::().unwrap(), + i128::from(ABOVE), + "MAX keeps the exact integer" + ); + assert_eq!( + rows[0][2].parse::().unwrap(), + (ABOVE + 2) as f64 + 0.5, + "SUM with a float input is a float" + ); +} + +/// A SUM whose exact total lies past the decimal range is refused as +/// `numeric_value_out_of_range`, the code a Control-Plane overflow carries. +/// Each input fits the decimal range, so the overflow is in the aggregate +/// itself. +#[tokio::test] +async fn document_schemaless_sum_past_decimal_range_is_22003() { + let srv = TestServer::start().await; + srv.exec("CREATE COLLECTION huge WITH (engine='document_schemaless')") + .await + .unwrap(); + srv.exec( + "INSERT INTO huge (id, v) VALUES \ + ('a', 50000000000000000000000000000), ('b', 50000000000000000000000000000)", + ) + .await + .unwrap(); + srv.expect_error("SELECT SUM(v) FROM huge", "SQLSTATE 22003") + .await; +} + +/// SUM over a `DECIMAL` column with fractions is the exact decimal total. +/// `0.1 + 0.2 + 0.3` through `f64` is `0.6000000000000001`. +#[tokio::test] +async fn decimal_column_sums_exactly() { + let srv = TestServer::start().await; + for (name, create) in [ + ( + "dec_strict", + "CREATE COLLECTION dec_strict (id TEXT PRIMARY KEY, v DECIMAL) \ + WITH (engine='document_strict')", + ), + ( + "dec_columnar", + "CREATE COLLECTION dec_columnar COLUMNS (id TEXT, v DECIMAL) \ + WITH (engine='columnar')", + ), + ] { + srv.exec(create).await.unwrap(); + srv.exec(&format!( + "INSERT INTO {name} (id, v) VALUES ('a', 0.1), ('b', 0.2), ('c', 0.3)" + )) + .await + .unwrap(); + let rows = srv + .query_rows(&format!("SELECT SUM(v) FROM {name}")) + .await + .unwrap(); + assert_eq!(rows, vec![vec!["0.6".to_string()]], "{name}"); + } +} diff --git a/nodedb/tests/wire/cases/aggregate_non_finite_float.rs b/nodedb/tests/wire/cases/aggregate_non_finite_float.rs new file mode 100644 index 000000000..3557c0e94 --- /dev/null +++ b/nodedb/tests/wire/cases/aggregate_non_finite_float.rs @@ -0,0 +1,94 @@ +// SPDX-License-Identifier: BUSL-1.1 + +//! A float aggregate that overflows or meets NaN reaches the client as +//! PostgreSQL renders it: `Infinity`, `-Infinity` or `NaN`, never NULL. +//! +//! `1e308 + 1e308` overflows `f64` to `Infinity`. A SUM over a NaN row is +//! NaN. Each engine runs the same checks. + +use crate::harness::TestServer; + +/// The text of the one cell `SELECT SUM(f) FROM collection` returns. +/// `None` for SQL NULL. +async fn sum_text(srv: &TestServer, collection: &str) -> Option { + let msgs = srv + .client + .simple_query(&format!("SELECT SUM(f) FROM {collection}")) + .await + .unwrap_or_else(|e| panic!("{collection}: SUM must run: {e:?}")); + let cells: Vec> = msgs + .iter() + .filter_map(|m| match m { + tokio_postgres::SimpleQueryMessage::Row(row) => Some(row.get(0).map(str::to_owned)), + _ => None, + }) + .collect(); + assert_eq!( + cells.len(), + 1, + "{collection}: one aggregate row, got {cells:?}" + ); + cells.into_iter().next().flatten() +} + +/// Insert `(id, f)` rows into `collection`. `ts` is the row index, so a +/// collection keyed by time gets distinct keys. +async fn insert_floats(srv: &TestServer, collection: &str, values: &[&str]) { + for (i, f) in values.iter().enumerate() { + srv.exec(&format!( + "INSERT INTO {collection} (id, ts, f) VALUES ('r{i}', {i}, {f})" + )) + .await + .unwrap_or_else(|e| panic!("{collection}: insert {f}: {e}")); + } +} + +/// Run the overflow and NaN checks against collections made by +/// `create(name)`. +async fn check_engine(srv: &TestServer, create: impl Fn(&str) -> String) { + srv.exec(&create("overflow")).await.unwrap(); + insert_floats(srv, "overflow", &["1e308", "1e308"]).await; + assert_eq!( + sum_text(srv, "overflow").await.as_deref(), + Some("Infinity"), + "SUM past f64::MAX must render Infinity" + ); + + srv.exec(&create("negative_overflow")).await.unwrap(); + insert_floats(srv, "negative_overflow", &["-1e308", "-1e308"]).await; + assert_eq!( + sum_text(srv, "negative_overflow").await.as_deref(), + Some("-Infinity"), + "SUM past -f64::MAX must render -Infinity" + ); + + srv.exec(&create("nan")).await.unwrap(); + insert_floats(srv, "nan", &["1.5", "'NaN'::float8"]).await; + assert_eq!( + sum_text(srv, "nan").await.as_deref(), + Some("NaN"), + "SUM over a NaN row must render NaN" + ); +} + +#[tokio::test] +async fn document_schemaless_non_finite_sum_renders_postgres_text() { + let srv = TestServer::start().await; + check_engine(&srv, |name| { + format!("CREATE COLLECTION {name} WITH (engine='document_schemaless')") + }) + .await; +} + +#[tokio::test] +async fn columnar_non_finite_sum_renders_postgres_text() { + let srv = TestServer::start().await; + check_engine(&srv, |name| { + format!( + "CREATE COLLECTION {name} \ + COLUMNS (id TEXT, ts BIGINT, f FLOAT) \ + WITH (engine='columnar')" + ) + }) + .await; +} diff --git a/nodedb/tests/wire/cases/sql_u64_literal.rs b/nodedb/tests/wire/cases/sql_u64_literal.rs new file mode 100644 index 000000000..5c03d5f8e --- /dev/null +++ b/nodedb/tests/wire/cases/sql_u64_literal.rs @@ -0,0 +1,135 @@ +// SPDX-License-Identifier: BUSL-1.1 + +//! An integer literal past `i64::MAX` keeps its exact number end to end. +//! +//! `18446744073709551615` (`u64::MAX`) is stored as a msgpack `uint64` and +//! reads back digit for digit. WHERE equality finds it, and ORDER BY puts it +//! after every smaller integer. A literal past the exact numeric range is +//! refused rather than rounded. + +use crate::harness::TestServer; + +const U64_MAX: &str = "18446744073709551615"; +const JUST_ABOVE_I64: &str = "9223372036854775808"; + +/// Create `collection` with `create` and insert the three test rows. +async fn seed(srv: &TestServer, create: &str, collection: &str) { + srv.exec(create).await.unwrap(); + srv.exec(&format!( + "INSERT INTO {collection} (id, v) VALUES ('max', {U64_MAX}), \ + ('mid', {JUST_ABOVE_I64}), ('small', 5)" + )) + .await + .unwrap(); +} + +/// The one-column result of `sql`, row by row. +async fn column(srv: &TestServer, sql: &str) -> Vec { + srv.query_rows(sql) + .await + .unwrap_or_else(|e| panic!("{sql}: {e}")) + .into_iter() + .map(|row| row[0].clone()) + .collect() +} + +/// Read back, WHERE equality, and ORDER BY on `collection`. +async fn check(srv: &TestServer, collection: &str) { + assert_eq!( + column(srv, &format!("SELECT v FROM {collection} WHERE id = 'max'")).await, + vec![U64_MAX.to_string()], + "{collection}: u64::MAX reads back exactly" + ); + assert_eq!( + column(srv, &format!("SELECT v FROM {collection} WHERE id = 'mid'")).await, + vec![JUST_ABOVE_I64.to_string()], + "{collection}: i64::MAX + 1 reads back exactly" + ); + assert_eq!( + column( + srv, + &format!("SELECT id FROM {collection} WHERE v = {U64_MAX}") + ) + .await, + vec!["max".to_string()], + "{collection}: WHERE equality finds u64::MAX" + ); + assert_eq!( + column( + srv, + &format!("SELECT id FROM {collection} WHERE v = {JUST_ABOVE_I64}") + ) + .await, + vec!["mid".to_string()], + "{collection}: WHERE equality tells i64::MAX + 1 from u64::MAX" + ); + assert_eq!( + column(srv, &format!("SELECT id FROM {collection} ORDER BY v")).await, + vec!["small", "mid", "max"], + "{collection}: ORDER BY v ascending" + ); + assert_eq!( + column(srv, &format!("SELECT id FROM {collection} ORDER BY v DESC")).await, + vec!["max", "mid", "small"], + "{collection}: ORDER BY v descending" + ); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn schemaless_document_keeps_u64_max() { + let srv = TestServer::start().await; + seed( + &srv, + "CREATE COLLECTION u64_doc WITH (engine='document_schemaless')", + "u64_doc", + ) + .await; + check(&srv, "u64_doc").await; +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn strict_decimal_column_keeps_u64_max() { + let srv = TestServer::start().await; + seed( + &srv, + "CREATE COLLECTION u64_strict (id TEXT PRIMARY KEY, v DECIMAL) \ + WITH (engine='document_strict')", + "u64_strict", + ) + .await; + check(&srv, "u64_strict").await; +} + +/// `i64::MIN` written as a negated literal is the integer `i64::MIN`. +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn negated_literal_reaches_i64_min() { + let srv = TestServer::start().await; + srv.exec("CREATE COLLECTION i64_min_doc WITH (engine='document_schemaless')") + .await + .unwrap(); + srv.exec("INSERT INTO i64_min_doc (id, v) VALUES ('min', -9223372036854775808)") + .await + .unwrap(); + assert_eq!( + column( + &srv, + "SELECT v FROM i64_min_doc WHERE v = -9223372036854775808" + ) + .await, + vec!["-9223372036854775808".to_string()] + ); +} + +/// A literal past the 96-bit exact range is an error, never a rounded float. +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn literal_past_exact_range_is_refused() { + let srv = TestServer::start().await; + srv.exec("CREATE COLLECTION u64_huge WITH (engine='document_schemaless')") + .await + .unwrap(); + srv.expect_error( + "INSERT INTO u64_huge (id, v) VALUES ('x', 79228162514264337593543950336)", + "out of range", + ) + .await; +} From f396b9d9050cd195f198c49a9d6595af921588fa Mon Sep 17 00:00:00 2001 From: Farhan Syah Date: Thu, 8 Oct 2026 16:51:29 +0800 Subject: [PATCH 04/24] fix(aggregate): total SUM and AVG exactly and keep MIN/MAX values SUM and AVG accumulate integers exactly and floats with compensated addition, and fail with 22003 when an integer total leaves the Decimal range. MIN and MAX keep the original value and compare it exactly, so an integer column returns an integer. The rule applies to document, columnar and timeseries aggregates, window functions (now evaluated over typed values rather than JSON), HAVING, streaming materialized views, continuous aggregates and the Control-Plane post-aggregate. Spill runs, shuffle state and sketches encode as MessagePack so non-finite values survive. --- nodedb-query/src/msgpack_scan/aggregate.rs | 735 ------------------ .../src/msgpack_scan/aggregate_helpers.rs | 30 +- nodedb-query/src/numeric_sum.rs | 669 ++++++++++++++++ nodedb-query/src/window/aggregate.rs | 386 +++++++-- nodedb-query/src/window/arg.rs | 19 +- nodedb-query/src/window/eval.rs | 282 +++---- nodedb-query/src/window/extremum.rs | 93 +++ nodedb-query/src/window/frame.rs | 37 +- nodedb-query/src/window/helpers.rs | 64 +- nodedb-query/src/window/mod.rs | 1 + nodedb-query/src/window/offset.rs | 10 +- nodedb-query/src/window/ranking.rs | 55 +- nodedb-query/src/window/running.rs | 47 +- nodedb-query/src/window/value_agg.rs | 414 +++++++--- nodedb-query/src/window/value_eval.rs | 6 +- nodedb-sql/src/planner/select/select_stmt.rs | 26 +- nodedb-types/src/approx/hll.rs | 9 +- nodedb-types/src/approx/spacesaving.rs | 8 +- nodedb-types/src/approx/tdigest.rs | 10 +- nodedb/src/bridge/mod.rs | 5 +- nodedb/src/control/server/post_aggregate.rs | 253 +++--- .../src/data/executor/handlers/accum/feed.rs | 77 +- .../data/executor/handlers/accum/finalize.rs | 58 +- .../src/data/executor/handlers/accum/merge.rs | 221 ++++-- .../src/data/executor/handlers/accum/new.rs | 4 +- .../src/data/executor/handlers/accum/state.rs | 40 +- .../data/executor/handlers/aggregate/exec.rs | 94 +-- .../data/executor/handlers/aggregate/mod.rs | 4 +- .../data/executor/handlers/aggregate/rows.rs | 255 ++++-- .../handlers/aggregate/shuffle_merge.rs | 42 +- .../executor/handlers/aggregate/state_emit.rs | 34 +- .../handlers/aggregate/streaming/finalize.rs | 61 +- .../handlers/aggregate/streaming/over_docs.rs | 7 +- .../data/executor/handlers/columnar_agg.rs | 345 +++++--- .../executor/handlers/columnar_agg_support.rs | 107 ++- .../handlers/document/read/projection.rs | 222 +++--- .../executor/handlers/document/read/scan.rs | 171 ++-- .../executor/handlers/grouping_sets_exec.rs | 100 +-- .../handlers/provider_scan_compute.rs | 31 +- .../data/executor/handlers/spill/columnar.rs | 14 +- .../src/data/executor/handlers/spill/core.rs | 80 +- .../data/executor/handlers/spill/groupby.rs | 22 +- .../executor/handlers/timeseries/encode.rs | 158 +++- .../executor/handlers/timeseries_gap_fill.rs | 4 +- nodedb/src/engine/timeseries/columnar_agg.rs | 148 ++-- .../timeseries/continuous_agg/manager.rs | 366 +++++++-- .../engine/timeseries/continuous_agg/mod.rs | 3 +- .../timeseries/continuous_agg/partial.rs | 233 ------ .../continuous_agg/partial/bucket.rs | 211 +++++ .../continuous_agg/partial/column.rs | 213 +++++ .../continuous_agg/partial/layout.rs | 133 ++++ .../timeseries/continuous_agg/partial/mod.rs | 9 + .../timeseries/continuous_agg/refresh.rs | 133 ++-- .../timeseries/continuous_agg/rollup.rs | 59 ++ .../engine/timeseries/grouped_scan/types.rs | 15 +- nodedb/src/event/consumer/run.rs | 24 +- nodedb/src/event/streaming_mv/persist.rs | 101 ++- nodedb/src/event/streaming_mv/processor.rs | 110 ++- nodedb/src/event/streaming_mv/query.rs | 291 +++++-- nodedb/src/event/streaming_mv/state.rs | 285 +++++-- .../tests/inproc/cases/event_streaming_mv.rs | 132 ++-- nodedb/tests/inproc/cases/sql_streaming_mv.rs | 169 +++- 62 files changed, 5105 insertions(+), 2840 deletions(-) delete mode 100644 nodedb-query/src/msgpack_scan/aggregate.rs create mode 100644 nodedb-query/src/numeric_sum.rs create mode 100644 nodedb-query/src/window/extremum.rs delete mode 100644 nodedb/src/engine/timeseries/continuous_agg/partial.rs create mode 100644 nodedb/src/engine/timeseries/continuous_agg/partial/bucket.rs create mode 100644 nodedb/src/engine/timeseries/continuous_agg/partial/column.rs create mode 100644 nodedb/src/engine/timeseries/continuous_agg/partial/layout.rs create mode 100644 nodedb/src/engine/timeseries/continuous_agg/partial/mod.rs create mode 100644 nodedb/src/engine/timeseries/continuous_agg/rollup.rs 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/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/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/planner/select/select_stmt.rs b/nodedb-sql/src/planner/select/select_stmt.rs index 16b9faece..6bf9cf438 100644 --- a/nodedb-sql/src/planner/select/select_stmt.rs +++ b/nodedb-sql/src/planner/select/select_stmt.rs @@ -8,10 +8,10 @@ use super::comma_lateral::try_plan_comma_lateral; use super::derived_from::try_plan_derived_from; use super::helpers::{convert_projection, convert_where_to_filters}; use super::query_tail::QueryTail; -use super::where_search::try_extract_where_search; +use super::where_search::{SearchBodyClauses, refuse_dropped_clauses, try_extract_where_search}; use crate::error::{Result, SqlError}; use crate::functions::registry::FunctionRegistry; -use crate::planner::ast_helpers::strip_single_table_qualifiers; +use crate::planner::ast_helpers::{single_table_qualifiers, strip_single_table_qualifiers}; use crate::resolver::columns::TableScope; use crate::temporal::TemporalScope; use crate::types::*; @@ -146,14 +146,7 @@ pub(super) fn plan_select( // is scoped to the single-table branch only: the JOIN path (handled above // by `try_plan_join`) deliberately keeps qualifiers for merged-document // evaluation and is never reached here. - let valid_qualifiers: Vec<&str> = { - let ref_name = table.ref_name(); - if ref_name == table.name { - vec![table.name.as_str()] - } else { - vec![table.name.as_str(), ref_name] - } - }; + let valid_qualifiers = single_table_qualifiers(table); let normalized_select = strip_single_table_qualifiers(select, &valid_qualifiers)?; let select = &normalized_select; @@ -182,6 +175,19 @@ pub(super) fn plan_select( let where_projection = convert_projection(&select.projection, &scope)?; if let Some(plan) = try_extract_where_search(expr, table, functions, &where_projection)? { + refuse_dropped_clauses( + &plan, + SearchBodyClauses { + temporal: &temporal, + has_subqueries: !subquery_joins.is_empty(), + aggregates: has_aggregation(select, functions), + distinct: select.distinct.is_some(), + windows: !crate::planner::window::extract_window_functions( + select, functions, &scope, + )? + .is_empty(), + }, + )?; return Ok(PlannedSelect { plan, scope }); } cached_projection = Some(where_projection); diff --git a/nodedb-types/src/approx/hll.rs b/nodedb-types/src/approx/hll.rs index 9be0b0e53..cb76324f2 100644 --- a/nodedb-types/src/approx/hll.rs +++ b/nodedb-types/src/approx/hll.rs @@ -7,7 +7,7 @@ /// Uses 2^14 = 16384 registers (12 KB memory). Achieves ~0.8% relative /// error at any cardinality. Mergeable: `hll_a.merge(&hll_b)` produces /// the union cardinality. -#[derive(Debug, serde::Serialize, serde::Deserialize)] +#[derive(Debug, zerompk::ToMessagePack, zerompk::FromMessagePack)] pub struct HyperLogLog { registers: Vec, precision: u8, @@ -214,7 +214,7 @@ mod tests { } #[test] - fn hll_serde_roundtrip_merge_semantics() { + fn hll_msgpack_roundtrip_merge_semantics() { let mut a = HyperLogLog::new(); let mut b = HyperLogLog::new(); for i in 0..1000u64 { @@ -224,9 +224,8 @@ mod tests { b.add(i); } - // Serialize and deserialize a via serde_json (serde derive, not zerompk). - let json = serde_json::to_vec(&a).expect("serialize HLL"); - let mut a_prime: HyperLogLog = serde_json::from_slice(&json).expect("deserialize HLL"); + let bytes = zerompk::to_msgpack_vec(&a).expect("serialize HLL"); + let mut a_prime: HyperLogLog = zerompk::from_msgpack(&bytes).expect("deserialize HLL"); // Merge b into deserialized a'. a_prime.merge(&b); diff --git a/nodedb-types/src/approx/spacesaving.rs b/nodedb-types/src/approx/spacesaving.rs index df1f5fc6c..a963adb67 100644 --- a/nodedb-types/src/approx/spacesaving.rs +++ b/nodedb-types/src/approx/spacesaving.rs @@ -9,7 +9,7 @@ use std::collections::HashMap; /// Tracks the K most frequent items with bounded memory. Items not in /// the top K are approximated — their counts may be over-estimated by /// at most the minimum count in the structure. -#[derive(Debug, serde::Serialize, serde::Deserialize)] +#[derive(Debug, zerompk::ToMessagePack, zerompk::FromMessagePack)] pub struct SpaceSaving { items: HashMap, max_items: usize, @@ -135,7 +135,7 @@ mod tests { } #[test] - fn spacesaving_serde_roundtrip_merge_semantics() { + fn spacesaving_msgpack_roundtrip_merge_semantics() { let mut a = SpaceSaving::new(5); let mut b = SpaceSaving::new(5); for _ in 0..100u64 { @@ -148,9 +148,9 @@ mod tests { b.add(2); } - let bytes = serde_json::to_vec(&a).expect("serialize SpaceSaving"); + let bytes = zerompk::to_msgpack_vec(&a).expect("serialize SpaceSaving"); let mut a_prime: SpaceSaving = - serde_json::from_slice(&bytes).expect("deserialize SpaceSaving"); + zerompk::from_msgpack(&bytes).expect("deserialize SpaceSaving"); a_prime.merge(&b); a.merge(&b); diff --git a/nodedb-types/src/approx/tdigest.rs b/nodedb-types/src/approx/tdigest.rs index db995977c..8cfb26ad4 100644 --- a/nodedb-types/src/approx/tdigest.rs +++ b/nodedb-types/src/approx/tdigest.rs @@ -3,7 +3,7 @@ //! TDigest — approximate percentile estimation (mergeable centroids). /// Centroid in the t-digest: represents a cluster of values. -#[derive(Debug, Clone, Copy, serde::Serialize, serde::Deserialize)] +#[derive(Debug, Clone, Copy, zerompk::ToMessagePack, zerompk::FromMessagePack)] struct Centroid { mean: f64, count: u64, @@ -14,7 +14,7 @@ struct Centroid { /// Maintains a sorted set of centroids that approximate the data distribution. /// Accurate at the extremes (p1, p99) and reasonable in the middle. /// Mergeable across partitions and shards. -#[derive(Debug, serde::Serialize, serde::Deserialize)] +#[derive(Debug, zerompk::ToMessagePack, zerompk::FromMessagePack)] pub struct TDigest { centroids: Vec, max_centroids: usize, @@ -223,7 +223,7 @@ mod tests { } #[test] - fn tdigest_serde_roundtrip_merge_semantics() { + fn tdigest_msgpack_roundtrip_merge_semantics() { let mut a = TDigest::new(); let mut b = TDigest::new(); for i in 0..500 { @@ -233,8 +233,8 @@ mod tests { b.add(i as f64); } - let bytes = serde_json::to_vec(&a).expect("serialize TDigest"); - let mut a_prime: TDigest = serde_json::from_slice(&bytes).expect("deserialize TDigest"); + let bytes = zerompk::to_msgpack_vec(&a).expect("serialize TDigest"); + let mut a_prime: TDigest = zerompk::from_msgpack(&bytes).expect("deserialize TDigest"); a_prime.merge(&b); a.merge(&b); diff --git a/nodedb/src/bridge/mod.rs b/nodedb/src/bridge/mod.rs index 0d18936ac..f1e7b82ee 100644 --- a/nodedb/src/bridge/mod.rs +++ b/nodedb/src/bridge/mod.rs @@ -6,9 +6,8 @@ pub mod envelope; pub mod quiesce; // Shared query engine re-exports. Origin's internal code names -// `crate::bridge::expr_eval`, `crate::bridge::json_ops`, -// `crate::bridge::scan_filter`, and `crate::bridge::window_func`; the types -// behind them come from nodedb-query. +// `crate::bridge::expr_eval`, `crate::bridge::scan_filter`, and +// `crate::bridge::window_func`; the types behind them come from nodedb-query. pub mod expr_eval { pub use nodedb_query::expr::{BinaryOp, CastType, ComputedColumn, SqlExpr}; } diff --git a/nodedb/src/control/server/post_aggregate.rs b/nodedb/src/control/server/post_aggregate.rs index 5a916a0bb..a72185e1e 100644 --- a/nodedb/src/control/server/post_aggregate.rs +++ b/nodedb/src/control/server/post_aggregate.rs @@ -5,11 +5,20 @@ //! When a query has `GROUP BY` over a `JOIN` result, the Data Plane cores //! return raw join rows. This module aggregates them in the Control Plane. //! All processing stays in msgpack — no JSON intermediary. +//! +//! SUM / AVG total exactly per `nodedb_query::ExactSum`. MIN / MAX keep the +//! original number and compare exactly, so integers above 2^53 never round +//! through `f64`. use std::collections::HashMap; -use crate::bridge::envelope::{Payload, Response}; use nodedb_query::agg_key::canonical_agg_key; +use nodedb_query::msgpack_scan::reader; +use nodedb_query::numeric_sum::{ExactSum, sum_input}; +use nodedb_query::window::extremum::value_replaces; +use nodedb_types::Value; + +use crate::bridge::envelope::{Payload, Response}; /// Apply GROUP BY + aggregate functions on a join response payload. /// @@ -52,12 +61,8 @@ pub fn apply_post_aggregation( // Compute and write each aggregate. for (op, field) in aggregates { let agg_key = canonical_agg_key(op, field); - let value = compute_aggregate(op, field, group_rows); - match value { - AggValue::Int(n) => writer::write_kv_i64(&mut buf, &agg_key, n), - AggValue::Float(f) => writer::write_kv_f64(&mut buf, &agg_key, f), - AggValue::Null => writer::write_kv_null(&mut buf, &agg_key), - } + let value = compute_aggregate(op, field, group_rows)?; + write_kv_number(&mut buf, &agg_key, &value); } } @@ -67,16 +72,20 @@ pub fn apply_post_aggregation( }) } -enum AggValue { - Int(i64), - Float(f64), - Null, +/// Write an aggregate result: an integer, a float, NULL, or a `Decimal` +/// as its exact text (the msgpack form of a `Decimal`). +fn write_kv_number(buf: &mut Vec, key: &str, value: &Value) { + use nodedb_query::msgpack_scan::writer; + match value { + Value::Integer(n) => writer::write_kv_i64(buf, key, *n), + Value::Float(f) => writer::write_kv_f64(buf, key, *f), + Value::Decimal(d) => writer::write_kv_str(buf, key, &d.to_string()), + _ => writer::write_kv_null(buf, key), + } } /// Parse a msgpack array payload into individual row slices. fn parse_msgpack_rows(bytes: &[u8]) -> crate::Result> { - use nodedb_query::msgpack_scan::reader; - if bytes.is_empty() { return Ok(Vec::new()); } @@ -99,17 +108,12 @@ fn parse_msgpack_rows(bytes: &[u8]) -> crate::Result> { Ok(rows) } -/// Extract a field value as string from a msgpack map row. -/// Handles "collection.field" suffix matching. -fn extract_field_str(row: &[u8], field: &str) -> Option { - use nodedb_query::msgpack_scan::reader; - - // Try exact match first. - if let Some((start, end)) = nodedb_query::msgpack_scan::extract_field(row, 0, field) { - return Some(read_value_as_string(row, start, end)); +/// The value range of `field` in a msgpack map row: an exact key match +/// first, then a `"collection.field"` suffix match. +fn field_range(row: &[u8], field: &str) -> Option<(usize, usize)> { + if let Some(range) = nodedb_query::msgpack_scan::extract_field(row, 0, field) { + return Some(range); } - - // Suffix match: iterate map keys looking for "*.{field}". let suffix = format!(".{field}"); let (count, mut pos) = reader::map_header(row, 0)?; for _ in 0..count { @@ -117,68 +121,36 @@ fn extract_field_str(row: &[u8], field: &str) -> Option { let key_end = reader::skip_value(row, pos)?; let val_end = reader::skip_value(row, key_end)?; if key.ends_with(&suffix) { - return Some(read_value_as_string(row, key_end, val_end)); + return Some((key_end, val_end)); } pos = val_end; } None } -/// Extract a numeric field value from a msgpack map row. -fn extract_number(row: &[u8], field: &str) -> Option { - use nodedb_query::msgpack_scan::reader; - - let try_field = |name: &str| -> Option { - let (start, _end) = nodedb_query::msgpack_scan::extract_field(row, 0, name)?; - if let Some(i) = reader::read_i64(row, start) { - return Some(i as f64); - } - if let Some(f) = reader::read_f64(row, start) { - return Some(f); - } - // Try parsing string as number. - reader::read_str(row, start).and_then(|s| s.parse().ok()) - }; +/// Extract a field value as string from a msgpack map row. +fn extract_field_str(row: &[u8], field: &str) -> Option { + let (start, end) = field_range(row, field)?; + Some(read_value_as_string(row, start, end)) +} +/// The number a row contributes to SUM / AVG / MIN / MAX of `field`: a +/// number as itself, a numeric string as the number it spells. Integers +/// stay exact. `*` contributes `1`. +fn extract_number(row: &[u8], field: &str) -> Option { if field == "*" { - return Some(1.0); - } - - // Exact match. - if let Some(v) = try_field(field) { - return Some(v); - } - - // Suffix match. - let suffix = format!(".{field}"); - let (count, mut pos) = nodedb_query::msgpack_scan::reader::map_header(row, 0)?; - for _ in 0..count { - let key = nodedb_query::msgpack_scan::reader::read_str(row, pos)?; - let key_end = nodedb_query::msgpack_scan::reader::skip_value(row, pos)?; - let val_end = nodedb_query::msgpack_scan::reader::skip_value(row, key_end)?; - if key.ends_with(&suffix) { - if let Some(i) = nodedb_query::msgpack_scan::reader::read_i64(row, key_end) { - return Some(i as f64); - } - if let Some(f) = nodedb_query::msgpack_scan::reader::read_f64(row, key_end) { - return Some(f); - } - if let Some(s) = nodedb_query::msgpack_scan::reader::read_str(row, key_end) { - return s.parse().ok(); - } - } - pos = val_end; + return Some(Value::Integer(1)); } - None + let (start, _end) = field_range(row, field)?; + sum_input(&reader::read_value(row, start)?) } /// Read a msgpack value at [start..end) as a display string. fn read_value_as_string(bytes: &[u8], start: usize, end: usize) -> String { - use nodedb_query::msgpack_scan::reader; if let Some(s) = reader::read_str(bytes, start) { return s.to_string(); } - if let Some(i) = reader::read_i64(bytes, start) { + if let Some(i) = reader::read_integer(bytes, start) { return i.to_string(); } if let Some(f) = reader::read_f64(bytes, start) { @@ -195,36 +167,123 @@ fn read_value_as_string(bytes: &[u8], start: usize, end: usize) -> String { } /// Compute a single aggregate over a group of msgpack rows. -fn compute_aggregate(op: &str, field: &str, rows: &[&[u8]]) -> AggValue { - match op { - "count" => AggValue::Int(rows.len() as i64), - "sum" => { - let sum: f64 = rows.iter().filter_map(|r| extract_number(r, field)).sum(); - AggValue::Float(sum) +/// +/// Fails with `EvalError::NumericOverflow` when an exact integer SUM / AVG +/// total lies outside the `Decimal` range. +fn compute_aggregate( + op: &str, + field: &str, + rows: &[&[u8]], +) -> Result { + let numbers = || rows.iter().filter_map(|r| extract_number(r, field)); + let exact_sum = || { + let mut acc = ExactSum::new(); + for v in numbers() { + acc.add_value(&v); } - "avg" => { - let values: Vec = rows - .iter() - .filter_map(|r| extract_number(r, field)) - .collect(); - if values.is_empty() { - AggValue::Null - } else { - AggValue::Float(values.iter().sum::() / values.len() as f64) + acc + }; + let extremum = |want_max: bool| { + let mut best: Option = None; + for v in numbers() { + if value_replaces(&v, best.as_ref(), want_max) { + best = Some(v); } } - "min" => rows - .iter() - .filter_map(|r| extract_number(r, field)) - .min_by(|a, b| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal)) - .map(AggValue::Float) - .unwrap_or(AggValue::Null), - "max" => rows - .iter() - .filter_map(|r| extract_number(r, field)) - .max_by(|a, b| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal)) - .map(AggValue::Float) - .unwrap_or(AggValue::Null), - _ => AggValue::Null, + best.unwrap_or(Value::Null) + }; + Ok(match op { + "count" => Value::Integer(rows.len() as i64), + "sum" => exact_sum().sum()?, + "avg" => exact_sum().avg()?, + "min" => extremum(false), + "max" => extremum(true), + _ => Value::Null, + }) +} + +#[cfg(test)] +mod tests { + use super::*; + use serde_json::json; + + const ABOVE: i64 = 9_007_199_254_740_993; + const AT: i64 = 9_007_199_254_740_992; + + fn encode(v: &serde_json::Value) -> Vec { + nodedb_types::json_msgpack::json_to_msgpack(v).expect("encode") + } + + fn agg(op: &str, vals: &[serde_json::Value]) -> Value { + let docs: Vec> = vals.iter().map(|v| encode(&json!({"t.v": v}))).collect(); + let rows: Vec<&[u8]> = docs.iter().map(Vec::as_slice).collect(); + compute_aggregate(op, "v", &rows).unwrap() + } + + #[test] + fn integers_above_2_pow_53_stay_exact() { + let vals = [json!(ABOVE), json!(AT)]; + assert_eq!(agg("sum", &vals), Value::Integer(ABOVE + AT)); + assert_eq!(agg("min", &vals), Value::Integer(AT)); + assert_eq!(agg("max", &vals), Value::Integer(ABOVE)); + assert_eq!(agg("avg", &vals), Value::Float(AT as f64)); + } + + #[test] + fn nanosecond_timestamps_stay_exact() { + let vals = [ + json!(1_700_000_000_000_000_002_i64), + json!(1_700_000_000_000_000_001_i64), + ]; + assert_eq!(agg("min", &vals), Value::Integer(1_700_000_000_000_000_001)); + assert_eq!(agg("sum", &vals), Value::Integer(3_400_000_000_000_000_003)); + } + + #[test] + fn u64_above_i64_max_and_sum_past_i64() { + let vals = [json!(u64::MAX), json!(i64::MAX)]; + assert_eq!( + agg("max", &vals), + Value::Decimal(rust_decimal::Decimal::from(u64::MAX)) + ); + assert_eq!(agg("min", &vals), Value::Integer(i64::MAX)); + assert_eq!( + agg("sum", &vals), + Value::Decimal(rust_decimal::Decimal::from_i128_with_scale( + i128::from(u64::MAX) + i128::from(i64::MAX), + 0 + )) + ); + } + + #[test] + fn mixed_int_float_and_numeric_strings() { + assert_eq!( + agg("sum", &[json!(2), json!(0.5), json!("7")]), + Value::Float(9.5) + ); + let vals = [json!(AT), json!("9007199254740993"), json!(0.5)]; + assert_eq!(agg("max", &vals), Value::Integer(ABOVE)); + assert_eq!(agg("min", &vals), Value::Float(0.5)); + assert_eq!(agg("sum", &[json!("x")]), Value::Null); + } + + #[test] + fn group_key_text_keeps_u64_digits() { + let doc = encode(&json!({"k": u64::MAX})); + assert_eq!( + extract_field_str(&doc, "k"), + Some("18446744073709551615".to_string()) + ); + } + + #[test] + fn decimal_results_write_exact_text() { + // A one-entry map: fixmap header, then the written key/value pair. + let mut doc = vec![0x81]; + let total = Value::Decimal(rust_decimal::Decimal::from(u64::MAX)); + write_kv_number(&mut doc, "sum(v)", &total); + let (start, _) = nodedb_query::msgpack_scan::extract_field(&doc, 0, "sum(v)").unwrap(); + assert_eq!(reader::read_str(&doc, start), Some("18446744073709551615")); } } diff --git a/nodedb/src/data/executor/handlers/accum/feed.rs b/nodedb/src/data/executor/handlers/accum/feed.rs index 8250d2745..1f3e38000 100644 --- a/nodedb/src/data/executor/handlers/accum/feed.rs +++ b/nodedb/src/data/executor/handlers/accum/feed.rs @@ -29,13 +29,9 @@ impl AggAccum { *n += 1; } } - AggAccum::SumAvg { sum, comp, n } => { - if let Some(v) = ah::extract_f64(doc, &agg.field, agg.expr.as_ref())? { - let y = v - *comp; - let t = *sum + y; - *comp = (t - *sum) - y; - *sum = t; - *n += 1; + AggAccum::SumAvg { sum } => { + if let Some(v) = ah::extract_sum_value(doc, &agg.field, agg.expr.as_ref())? { + sum.add_value(&v); } } AggAccum::SumAvgDistinct { seen } => { @@ -43,52 +39,20 @@ impl AggAccum { // SUM(DISTINCT col) / AVG(DISTINCT col) only credit each // distinct value once. NULL bytes (msgpack `0xc0`) are // ignored — they cannot meaningfully participate in a - // numeric aggregate. The parsed f64 is stored alongside - // the key so finalize can derive an order-independent - // sum (which also makes the state mergeable across - // spilled runs). + // numeric aggregate. The contributed number is stored + // alongside the key so finalize can derive an + // order-independent sum (which also makes the state + // mergeable across spilled runs). if let Some(bytes) = ah::extract_bytes(doc, &agg.field, agg.expr.as_ref())? && bytes != [0xc0u8] && let Entry::Vacant(slot) = seen.entry(bytes) - && let Some(v) = ah::extract_f64(doc, &agg.field, agg.expr.as_ref())? + && let Some(v) = ah::extract_sum_value(doc, &agg.field, agg.expr.as_ref())? { slot.insert(v); } } - AggAccum::Min { best } => { - if let Some(v) = ah::extract_value(doc, &agg.field, agg.expr.as_ref())? { - if v.is_null() { - return Ok(()); - } - let replace = match best { - None => true, - Some(cur) => { - nodedb_query::value_ops::compare_values(&v, cur) - == std::cmp::Ordering::Less - } - }; - if replace { - *best = Some(v); - } - } - } - AggAccum::Max { best } => { - if let Some(v) = ah::extract_value(doc, &agg.field, agg.expr.as_ref())? { - if v.is_null() { - return Ok(()); - } - let replace = match best { - None => true, - Some(cur) => { - nodedb_query::value_ops::compare_values(&v, cur) - == std::cmp::Ordering::Greater - } - }; - if replace { - *best = Some(v); - } - } - } + AggAccum::Min { best } => feed_extremum(best, agg, doc, false)?, + AggAccum::Max { best } => feed_extremum(best, agg, doc, true)?, AggAccum::CountDistinct { seen } => { if let Some(bytes) = ah::extract_bytes(doc, &agg.field, agg.expr.as_ref())? && bytes != [0xc0u8] @@ -164,7 +128,26 @@ impl AggAccum { } } -/// FNV-1a hash (matches the implementation in nodedb-query aggregate.rs). +/// Fold one document into a MIN (`want_max` false) or MAX (`want_max` true) +/// state. The original value is kept and compared exactly. NaN sorts above +/// every number, as in PostgreSQL: it wins MAX and loses MIN. +fn feed_extremum( + best: &mut Option, + agg: &AggregateSpec, + doc: &[u8], + want_max: bool, +) -> Result<(), nodedb_query::EvalError> { + use nodedb_query::msgpack_scan::aggregate_helpers as ah; + if let Some(v) = ah::extract_value(doc, &agg.field, agg.expr.as_ref())? + && !v.is_null() + && nodedb_query::window::extremum::value_replaces(&v, best.as_ref(), want_max) + { + *best = Some(v); + } + Ok(()) +} + +/// FNV-1a hash of raw value bytes. #[inline] fn fnv1a(bytes: &[u8]) -> u64 { let mut h: u64 = 0xcbf29ce484222325; diff --git a/nodedb/src/data/executor/handlers/accum/finalize.rs b/nodedb/src/data/executor/handlers/accum/finalize.rs index 01cdd9afb..a302a8528 100644 --- a/nodedb/src/data/executor/handlers/accum/finalize.rs +++ b/nodedb/src/data/executor/handlers/accum/finalize.rs @@ -8,46 +8,32 @@ use nodedb_types::Value; impl AggAccum { /// Consume the accumulator and produce the final `Value`. - pub(crate) fn finalize(self, agg: &AggregateSpec) -> Value { - match self { + /// + /// Fails with `EvalError::NumericOverflow` when an exact integer SUM / + /// AVG total lies outside the `Decimal` range. + pub(crate) fn finalize(self, agg: &AggregateSpec) -> Result { + Ok(match self { AggAccum::Count { n } => Value::Integer(n as i64), - AggAccum::SumAvg { sum, n, .. } => { - // Zero contributing values → NULL for both SUM and AVG. SUM has - // no natural zero identity in SQL: `SUM` over an empty input is - // NULL, not 0.0. `n` counts the values actually fed, so `n == 0` - // is the exact zero-contribution signal (a real SUM that happens - // to total 0.0 still has `n > 0` and returns `0.0`). - if n == 0 { - Value::Null - } else if agg.function == "avg" { - Value::Float(sum / n as f64) + // No contributing value is NULL for both SUM and AVG: SUM has no + // zero identity in SQL. Integer inputs total exactly. + AggAccum::SumAvg { sum } => { + if agg.function == "avg" { + sum.avg()? } else { - Value::Float(sum) + sum.sum()? } } AggAccum::SumAvgDistinct { seen } => { - let n = seen.len(); - // Kahan-compensated sum over the deduped values. Iteration - // order is arbitrary, but a DISTINCT sum is order-independent - // so the result is deterministic regardless. - let mut sum = 0.0f64; - let mut comp = 0.0f64; - for &v in seen.values() { - let y = v - comp; - let t = sum + y; - comp = (t - sum) - y; - sum = t; + // A DISTINCT sum is order-independent, so the arbitrary map + // iteration order gives a deterministic result. + let mut sum = nodedb_query::ExactSum::new(); + for v in seen.values() { + sum.add_value(v); } - // Zero distinct values → NULL for both SUM(DISTINCT) and - // AVG(DISTINCT): same no-zero-identity rule as plain SUM. `n` - // is the distinct-value count (`seen.len()`), so `n == 0` is the - // zero-contribution signal. - if n == 0 { - Value::Null - } else if agg.function == "avg_distinct" { - Value::Float(sum / n as f64) + if agg.function == "avg_distinct" { + sum.avg()? } else { - Value::Float(sum) + sum.sum()? } } AggAccum::Min { best } => best.unwrap_or(Value::Null), @@ -55,7 +41,7 @@ impl AggAccum { AggAccum::CountDistinct { seen } => Value::Integer(seen.len() as i64), AggAccum::Welford { n, mean: _, m2 } => { if n < 2 { - return Value::Null; + return Ok(Value::Null); } let population = matches!( agg.function.as_str(), @@ -107,7 +93,7 @@ impl AggAccum { AggAccum::ArrayAggDistinct { values, .. } => Value::Array(values), AggAccum::PercentileCont { mut values, pct } => { if values.is_empty() { - return Value::Null; + return Ok(Value::Null); } values.sort_by(|a, b| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal)); let idx = (pct * (values.len() - 1) as f64).clamp(0.0, (values.len() - 1) as f64); @@ -117,6 +103,6 @@ impl AggAccum { Value::Float(values[lo] * (1.0 - frac) + values[hi] * frac) } AggAccum::StringAgg { parts } => Value::String(parts.join(",")), - } + }) } } diff --git a/nodedb/src/data/executor/handlers/accum/merge.rs b/nodedb/src/data/executor/handlers/accum/merge.rs index deab34e59..732c8c7dc 100644 --- a/nodedb/src/data/executor/handlers/accum/merge.rs +++ b/nodedb/src/data/executor/handlers/accum/merge.rs @@ -22,60 +22,20 @@ pub(super) fn merge_accum(dst: &mut AggAccum, other: AggAccum) { (AggAccum::Count { n: a }, AggAccum::Count { n: b }) => { *a += b; } - ( - AggAccum::SumAvg { - sum: sa, - comp: ca, - n: na, - }, - AggAccum::SumAvg { - sum: sb, - comp: _cb, - n: nb, - }, - ) => { - // Kahan-compensated addition of the other partial sum into self. - let y = sb - *ca; - let t = *sa + y; - *ca = (t - *sa) - y; - *sa = t; - *na += nb; + (AggAccum::SumAvg { sum: a }, AggAccum::SumAvg { sum: b }) => { + // Integer parts add exactly; float parts add compensated. + a.merge(&b); } (AggAccum::SumAvgDistinct { seen: a }, AggAccum::SumAvgDistinct { seen: b }) => { - // Union the deduped value maps; the first-seen parsed f64 wins - // (all instances of the same key carry the same value, so the - // choice is immaterial). The sum is re-derived at finalize. + // Union the deduped value maps; the first-seen number wins (all + // instances of the same key carry the same value, so the choice + // is immaterial). The sum is re-derived at finalize. for (key, value) in b { a.entry(key).or_insert(value); } } - (AggAccum::Min { best: a }, AggAccum::Min { best: b }) => { - if let Some(bv) = b { - let replace = match a { - None => true, - Some(av) => { - nodedb_query::value_ops::compare_values(&bv, av) == std::cmp::Ordering::Less - } - }; - if replace { - *a = Some(bv); - } - } - } - (AggAccum::Max { best: a }, AggAccum::Max { best: b }) => { - if let Some(bv) = b { - let replace = match a { - None => true, - Some(av) => { - nodedb_query::value_ops::compare_values(&bv, av) - == std::cmp::Ordering::Greater - } - }; - if replace { - *a = Some(bv); - } - } - } + (AggAccum::Min { best: a }, AggAccum::Min { best: b }) => merge_extremum(a, b, false), + (AggAccum::Max { best: a }, AggAccum::Max { best: b }) => merge_extremum(a, b, true), (AggAccum::CountDistinct { seen: a }, AggAccum::CountDistinct { seen: b }) => { a.extend(b); } @@ -156,6 +116,21 @@ pub(super) fn merge_accum(dst: &mut AggAccum, other: AggAccum) { } } +/// Merge a partial MIN (`want_max` false) or MAX (`want_max` true) extreme +/// into `dst`. Exact comparison. NaN sorts above every number, as in +/// PostgreSQL: it wins MAX and loses MIN. +fn merge_extremum( + dst: &mut Option, + other: Option, + want_max: bool, +) { + if let Some(candidate) = other + && nodedb_query::window::extremum::value_replaces(&candidate, dst.as_ref(), want_max) + { + *dst = Some(candidate); + } +} + /// Merge all accumulators from `other` into `dst` element-wise. pub(super) fn merge_group_state(dst: &mut GroupState, other: GroupState) { assert_eq!( @@ -274,18 +249,18 @@ mod tests { a_sum.merge_from(b_sum); a_avg.merge_from(b_avg); - let Value::Float(cs) = combined_sum.finalize(&sum_spec) else { + let Value::Float(cs) = combined_sum.finalize(&sum_spec).unwrap() else { panic!("expected float"); }; - let Value::Float(ms) = a_sum.finalize(&sum_spec) else { + let Value::Float(ms) = a_sum.finalize(&sum_spec).unwrap() else { panic!("expected float"); }; assert!((cs - ms).abs() < 1e-9, "sum mismatch: {cs} vs {ms}"); - let Value::Float(ca) = combined_avg.finalize(&avg_spec) else { + let Value::Float(ca) = combined_avg.finalize(&avg_spec).unwrap() else { panic!("expected float"); }; - let Value::Float(ma) = a_avg.finalize(&avg_spec) else { + let Value::Float(ma) = a_avg.finalize(&avg_spec).unwrap() else { panic!("expected float"); }; assert!((ca - ma).abs() < 1e-9, "avg mismatch: {ca} vs {ma}"); @@ -330,20 +305,142 @@ mod tests { a_sum.merge_from(b_sum); a_avg.merge_from(b_avg); - assert_eq!(combined_sum.finalize(&sum_spec), Value::Float(15.0)); - assert_eq!(combined_avg.finalize(&avg_spec), Value::Float(3.0)); + assert_eq!(combined_sum.finalize(&sum_spec), Ok(Value::Integer(15))); + assert_eq!(combined_avg.finalize(&avg_spec), Ok(Value::Float(3.0))); assert_eq!( a_sum.finalize(&sum_spec), - Value::Float(15.0), + Ok(Value::Integer(15)), "sum_distinct merge" ); assert_eq!( a_avg.finalize(&avg_spec), - Value::Float(3.0), + Ok(Value::Float(3.0)), "avg_distinct merge" ); } + /// Feed `docs_a` and `docs_b` into one accumulator, and separately into + /// two partials merged after; return both finalized results. + fn single_and_merged( + spec: &AggregateSpec, + docs_a: &[Vec], + docs_b: &[Vec], + ) -> [Value; 2] { + let mut combined = AggAccum::new(spec); + for d in docs_a.iter().chain(docs_b) { + combined.feed(spec, d).unwrap(); + } + let mut a = AggAccum::new(spec); + for d in docs_a { + a.feed(spec, d).unwrap(); + } + let mut b = AggAccum::new(spec); + for d in docs_b { + b.feed(spec, d).unwrap(); + } + a.merge_from(b); + [combined.finalize(spec).unwrap(), a.finalize(spec).unwrap()] + } + + const ABOVE: i64 = 9_007_199_254_740_993; + const AT: i64 = 9_007_199_254_740_992; + + /// A one-field doc whose value is a raw msgpack `uint64`. + fn make_doc_u64(field: &str, value: u64) -> Vec { + let mut doc = vec![0x81, 0xa0 | field.len() as u8]; + doc.extend_from_slice(field.as_bytes()); + doc.push(0xcf); + doc.extend_from_slice(&value.to_be_bytes()); + doc + } + + #[test] + fn sum_min_max_keep_integers_above_2_pow_53_exact_through_merge() { + let a = [make_doc_i64("v", ABOVE)]; + let b = [make_doc_i64("v", AT)]; + for got in single_and_merged(&make_spec("sum", "v"), &a, &b) { + assert_eq!(got, Value::Integer(ABOVE + AT)); + } + for got in single_and_merged(&make_spec("min", "v"), &a, &b) { + assert_eq!(got, Value::Integer(AT)); + } + for got in single_and_merged(&make_spec("max", "v"), &a, &b) { + assert_eq!(got, Value::Integer(ABOVE)); + } + for got in single_and_merged(&make_spec("avg", "v"), &a, &b) { + assert_eq!(got, Value::Float(AT as f64)); + } + } + + #[test] + fn nanosecond_timestamps_sum_and_extremes_exactly() { + let a = [make_doc_i64("v", 1_700_000_000_000_000_002)]; + let b = [ + make_doc_i64("v", 1_700_000_000_000_000_001), + make_doc_i64("v", 1_700_000_000_000_000_003), + ]; + for got in single_and_merged(&make_spec("sum", "v"), &a, &b) { + assert_eq!(got, Value::Integer(5_100_000_000_000_000_006)); + } + for got in single_and_merged(&make_spec("min", "v"), &a, &b) { + assert_eq!(got, Value::Integer(1_700_000_000_000_000_001)); + } + for got in single_and_merged(&make_spec("max", "v"), &a, &b) { + assert_eq!(got, Value::Integer(1_700_000_000_000_000_003)); + } + } + + #[test] + fn sum_past_i64_is_decimal_through_merge() { + let a = [make_doc_i64("v", i64::MAX)]; + let b = [make_doc_i64("v", i64::MAX), make_doc_u64("v", u64::MAX)]; + let want = rust_decimal::Decimal::from_i128_with_scale( + 2 * i128::from(i64::MAX) + i128::from(u64::MAX), + 0, + ); + for got in single_and_merged(&make_spec("sum", "v"), &a, &b) { + assert_eq!(got, Value::Decimal(want)); + } + for got in single_and_merged(&make_spec("max", "v"), &a, &b) { + assert_eq!(got, Value::Decimal(rust_decimal::Decimal::from(u64::MAX))); + } + } + + #[test] + fn sum_mixed_int_float_is_float() { + let a = [make_doc_i64("v", 2)]; + let b = [make_doc_f64("v", 0.5)]; + for got in single_and_merged(&make_spec("sum", "v"), &a, &b) { + assert_eq!(got, Value::Float(2.5)); + } + for got in single_and_merged(&make_spec("max", "v"), &a, &b) { + assert_eq!(got, Value::Integer(2)); + } + } + + /// NaN sorts above every number, as in PostgreSQL: MAX is NaN and MIN + /// skips it, in a single pass and across a merge. + #[test] + fn nan_is_the_largest_extreme_in_feed_and_merge() { + let is_nan = |v: &Value| matches!(v, Value::Float(f) if f.is_nan()); + let a = [make_doc_f64("v", f64::NAN), make_doc_i64("v", 5)]; + let b = [make_doc_i64("v", 3)]; + for got in single_and_merged(&make_spec("min", "v"), &a, &b) { + assert_eq!(got, Value::Integer(3)); + } + for got in single_and_merged(&make_spec("max", "v"), &a, &b) { + assert!(is_nan(&got), "got {got:?}"); + } + let nan_only = [make_doc_f64("v", f64::NAN)]; + let nums = [make_doc_i64("v", 7)]; + for got in single_and_merged(&make_spec("max", "v"), &nan_only, &nums) { + assert!(is_nan(&got), "got {got:?}"); + } + for got in single_and_merged(&make_spec("min", "v"), &nums, &nan_only) { + assert_eq!(got, Value::Integer(7)); + } + } + #[test] fn merge_from_min_max() { let min_spec = make_spec("min", "v"); @@ -410,10 +507,10 @@ mod tests { } a.merge_from(b); - let Value::Float(cv) = combined.finalize(&spec) else { + let Value::Float(cv) = combined.finalize(&spec).expect("finalize") else { panic!("expected float"); }; - let Value::Float(mv) = a.finalize(&spec) else { + let Value::Float(mv) = a.finalize(&spec).expect("finalize") else { panic!("expected float"); }; let rel = (cv - mv).abs() / cv.abs().max(1e-12); @@ -481,10 +578,10 @@ mod tests { } a.merge_from(b); - let Value::Integer(cv) = combined.finalize(&spec) else { + let Value::Integer(cv) = combined.finalize(&spec).expect("finalize") else { panic!("expected int"); }; - let Value::Integer(mv) = a.finalize(&spec) else { + let Value::Integer(mv) = a.finalize(&spec).expect("finalize") else { panic!("expected int"); }; // HLL is approximate; require within 5% of expected 1000. @@ -510,7 +607,7 @@ mod tests { a.merge_from(b); // p50 of 0..200 should be close to 100. - let Value::Float(p50) = a.finalize(&spec) else { + let Value::Float(p50) = a.finalize(&spec).expect("finalize") else { panic!("expected float"); }; assert!((50.0..150.0).contains(&p50), "TDigest merge p50={p50}"); @@ -533,7 +630,7 @@ mod tests { } a.merge_from(b); - let Value::Array(arr) = a.finalize(&spec) else { + let Value::Array(arr) = a.finalize(&spec).expect("finalize") else { panic!("expected array"); }; assert_eq!(arr.len(), 3, "TopK should return k=3 items"); diff --git a/nodedb/src/data/executor/handlers/accum/new.rs b/nodedb/src/data/executor/handlers/accum/new.rs index 65a966bf8..7dd0c3f28 100644 --- a/nodedb/src/data/executor/handlers/accum/new.rs +++ b/nodedb/src/data/executor/handlers/accum/new.rs @@ -12,9 +12,7 @@ impl AggAccum { match agg.function.as_str() { "count" => AggAccum::Count { n: 0 }, "sum" | "avg" => AggAccum::SumAvg { - sum: 0.0, - comp: 0.0, - n: 0, + sum: nodedb_query::ExactSum::new(), }, "sum_distinct" | "avg_distinct" => AggAccum::SumAvgDistinct { seen: HashMap::new(), diff --git a/nodedb/src/data/executor/handlers/accum/state.rs b/nodedb/src/data/executor/handlers/accum/state.rs index 969d5a9ec..81f668265 100644 --- a/nodedb/src/data/executor/handlers/accum/state.rs +++ b/nodedb/src/data/executor/handlers/accum/state.rs @@ -18,23 +18,24 @@ pub(super) const ARRAY_AGG_CAP: usize = 10_000; /// Per-(group, aggregate-spec) running accumulator. /// -/// Derives `Serialize` / `Deserialize` so that partial states can be spilled -/// to disk by `GroupBySpiller` and merged back during finalize. -#[derive(serde::Serialize, serde::Deserialize)] +/// Encodes as MessagePack so that partial states can be spilled to disk by +/// `GroupBySpiller` and shipped between shards, then merged back during +/// finalize. MessagePack keeps NaN and ±Infinity exact. +#[derive(zerompk::ToMessagePack, zerompk::FromMessagePack)] pub(crate) enum AggAccum { /// count(*) or count(field). Count { n: u64 }, - /// sum / avg: Kahan-compensated running sum + count. - SumAvg { sum: f64, comp: f64, n: u64 }, + /// sum / avg: exact integer total, compensated float total, count. + SumAvg { sum: nodedb_query::ExactSum }, /// sum(DISTINCT col) / avg(DISTINCT col): map each distinct input - /// value (keyed by its raw msgpack bytes) to its parsed numeric - /// value. The sum and count are derived at finalize time, so the - /// state is order-independent and therefore mergeable across - /// spilled runs. Memory: O(num_distinct). - SumAvgDistinct { seen: HashMap, f64> }, - /// min. + /// value (keyed by its raw msgpack bytes) to the number it contributes. + /// The sum and count are derived at finalize time, so the state is + /// order-independent and therefore mergeable across spilled runs. + /// Memory: O(num_distinct). + SumAvgDistinct { seen: HashMap, Value> }, + /// min: the original value, compared exactly. Min { best: Option }, - /// max. + /// max: the original value, compared exactly. Max { best: Option }, /// count_distinct: set of raw msgpack bytes. CountDistinct { seen: HashSet> }, @@ -68,8 +69,9 @@ pub(crate) enum AggAccum { /// Per-group running state: one `AggAccum` per aggregate spec. /// -/// Serializable so that `GroupBySpiller` can spill partial states to disk. -#[derive(serde::Serialize, serde::Deserialize)] +/// Encodes as MessagePack so that `GroupBySpiller` can spill partial states +/// to disk and the shuffle path can ship them between shards. +#[derive(zerompk::ToMessagePack, zerompk::FromMessagePack)] pub(crate) struct GroupState { pub(super) accums: Vec, } @@ -99,11 +101,17 @@ impl GroupState { super::merge::merge_group_state(self, other); } - pub(crate) fn finalize(self, aggregates: &[AggregateSpec]) -> Vec<(String, Value)> { + /// Produce `(alias, value)` per aggregate. Fails with + /// `EvalError::NumericOverflow` when an exact SUM / AVG total lies + /// outside the `Decimal` range. + pub(crate) fn finalize( + self, + aggregates: &[AggregateSpec], + ) -> Result, nodedb_query::EvalError> { self.accums .into_iter() .zip(aggregates) - .map(|(accum, agg)| (agg.alias.clone(), accum.finalize(agg))) + .map(|(accum, agg)| Ok((agg.alias.clone(), accum.finalize(agg)?))) .collect() } } diff --git a/nodedb/src/data/executor/handlers/aggregate/exec.rs b/nodedb/src/data/executor/handlers/aggregate/exec.rs index 9ad11e02d..cb05daed7 100644 --- a/nodedb/src/data/executor/handlers/aggregate/exec.rs +++ b/nodedb/src/data/executor/handlers/aggregate/exec.rs @@ -8,7 +8,7 @@ use tracing::debug; use super::cache_entry::AggregateCacheEntry; use super::cache_key::{AggregateCacheKeyInputs, aggregate_cache_key, legacy_aggregate_pairs}; -use super::rows::{apply_user_aliases_to_rows, sort_aggregated_rows}; +use super::rows::{apply_user_aliases_to_rows, retain_having, sort_aggregated_rows}; use crate::bridge::envelope::{ErrorCode, Response, Status}; use crate::bridge::scan_filter::ScanFilter; use crate::data::executor::core_loop::CoreLoop; @@ -204,12 +204,7 @@ impl CoreLoop { } return match Ok::, crate::Error>(payload_buf) { Ok(payload) => self.response_with_payload(task, payload), - Err(e) => self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ), + Err(e) => self.response_error(task, ErrorCode::from(e)), }; } } @@ -260,23 +255,30 @@ impl CoreLoop { .join("groupby-spill") .join(format!("core-{}-columnar", self.core_id)); let columnar_spill_cap = self.query_tuning.groupby_max_groups_in_mem; - if let Some(mut agg_result) = legacy_aggs.and_then(|pairs| { - super::super::columnar_agg::try_columnar_aggregate( - &super::super::columnar_agg::ColumnarAggParams { - mt, - group_by: &group_fields, - aggregates: &pairs, - filters: &filter_predicates, - limit, - scan_limit, - spill_dir: &columnar_spill_dir, - spill_cap: columnar_spill_cap, - governor: self.governor.clone(), - db: task.request.database_id, - tenant: task.request.tenant_id, - }, - ) - }) { + let columnar = legacy_aggs + .map(|pairs| { + super::super::columnar_agg::try_columnar_aggregate( + &super::super::columnar_agg::ColumnarAggParams { + mt, + group_by: &group_fields, + aggregates: &pairs, + filters: &filter_predicates, + limit, + scan_limit, + spill_dir: &columnar_spill_dir, + spill_cap: columnar_spill_cap, + governor: self.governor.clone(), + db: task.request.database_id, + tenant: task.request.tenant_id, + }, + ) + }) + .transpose(); + let columnar = match columnar { + Ok(result) => result.flatten(), + Err(e) => return self.response_error(task, ErrorCode::from(e)), + }; + if let Some(mut agg_result) = columnar { if !having.is_empty() { let having_predicates: Vec = match zerompk::from_msgpack(having) { Ok(h) => h, @@ -285,30 +287,10 @@ impl CoreLoop { Vec::new() } }; - if !having_predicates.is_empty() { - // `Vec::retain`'s closure must return `bool`, so an - // evaluation error in a HAVING predicate is captured - // via this side-channel and checked once the retain - // finishes. HAVING is WHERE-shaped, so it gets the - // full error treatment. - let predicate_err: std::cell::RefCell> = - std::cell::RefCell::new(None); - agg_result.rows.retain(|row| { - if predicate_err.borrow().is_some() { - return true; - } - let mp = nodedb_types::json_to_msgpack_or_empty(row); - match ScanFilter::all_match_binary(&having_predicates, &mp) { - Ok(keep) => keep, - Err(e) => { - predicate_err.replace(Some(e)); - true - } - } - }); - if let Some(e) = predicate_err.take() { - return self.response_error(task, ErrorCode::from(e)); - } + // HAVING is WHERE-shaped, so an evaluation error in a + // predicate fails the statement. + if let Err(e) = retain_having(&mut agg_result.rows, &having_predicates) { + return self.response_error(task, ErrorCode::from(e)); } } @@ -322,7 +304,7 @@ impl CoreLoop { } agg_result.rows.truncate(limit); - return match crate::data::executor::response_codec::encode_json_vec_as_msgpack( + return match crate::data::executor::response_codec::encode_value_vec( &agg_result.rows, ) { Ok(payload) => { @@ -355,12 +337,7 @@ impl CoreLoop { } self.response_with_payload(task, payload) } - Err(e) => self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ), + Err(e) => self.response_error(task, ErrorCode::from(e)), }; } } @@ -392,12 +369,7 @@ impl CoreLoop { let docs = match scan_result { Ok(d) => d, Err(e) => { - return self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ); + return self.response_error(task, ErrorCode::from(e)); } }; self.aggregate_over_docs(super::streaming::over_docs::AggregateOverDocsParams { diff --git a/nodedb/src/data/executor/handlers/aggregate/mod.rs b/nodedb/src/data/executor/handlers/aggregate/mod.rs index 0d45ea620..279b2b954 100644 --- a/nodedb/src/data/executor/handlers/aggregate/mod.rs +++ b/nodedb/src/data/executor/handlers/aggregate/mod.rs @@ -12,14 +12,14 @@ //! `exec` (dispatch + fast paths), `streaming` (the spill-backed group-by //! accumulation, itself split into accumulate / finalize / over_docs phases), //! `cache_key` (result-cache key derivation), `rows` (post-aggregate alias -//! renaming and ORDER BY sorting), `state_emit` (the distributed-shuffle +//! renaming, HAVING and ORDER BY sorting), `state_emit` (the distributed-shuffle //! partial-state producer), and `shuffle_merge` (the partial-state consumer). mod cache_entry; mod cache_key; pub(in crate::data::executor) mod exec; mod invalidate; -mod rows; +pub(in crate::data::executor::handlers) mod rows; pub(in crate::data::executor) mod shuffle_merge; pub(in crate::data::executor) mod state_emit; mod streaming; diff --git a/nodedb/src/data/executor/handlers/aggregate/rows.rs b/nodedb/src/data/executor/handlers/aggregate/rows.rs index f4c015379..b1c9e3c5d 100644 --- a/nodedb/src/data/executor/handlers/aggregate/rows.rs +++ b/nodedb/src/data/executor/handlers/aggregate/rows.rs @@ -1,11 +1,23 @@ // SPDX-License-Identifier: BUSL-1.1 -//! Post-aggregate row helpers: user-alias renaming and ORDER BY sorting. +//! Post-aggregate row helpers: user-alias renaming, HAVING, and ORDER BY +//! sorting. +//! +//! Rows are `nodedb_types::Value::Object` maps keyed by output column name. +//! `Value` holds every aggregate result as computed, NaN and ±Infinity +//! included, which a JSON number cannot. + +use std::cmp::Ordering; use nodedb_physical::physical_plan::AggregateSpec; +use nodedb_types::Value; + +use crate::bridge::scan_filter::ScanFilter; -pub(super) fn apply_user_aliases_to_rows( - rows: &mut [serde_json::Value], +/// Rename each aggregate output column from its canonical alias to the +/// user alias, where the two differ. +pub(in crate::data::executor::handlers) fn apply_user_aliases_to_rows( + rows: &mut [Value], aggregates: &[AggregateSpec], ) { let renames: Vec<(&str, &str)> = aggregates @@ -23,7 +35,7 @@ pub(super) fn apply_user_aliases_to_rows( } for row in rows { - if let Some(obj) = row.as_object_mut() { + if let Value::Object(obj) = row { for (from, to) in &renames { if let Some(value) = obj.remove(*from) { obj.insert((*to).to_string(), value); @@ -33,10 +45,31 @@ pub(super) fn apply_user_aliases_to_rows( } } +/// Keep the rows that match every HAVING predicate. An evaluation error in +/// a predicate fails the statement. +pub(in crate::data::executor::handlers) fn retain_having( + rows: &mut Vec, + predicates: &[ScanFilter], +) -> crate::Result<()> { + if predicates.is_empty() { + return Ok(()); + } + let mut kept = Vec::with_capacity(rows.len()); + for row in rows.drain(..) { + let mp = nodedb_types::value_to_msgpack(&row).map_err(|e| crate::Error::Codec { + detail: format!("HAVING row encode: {e}"), + })?; + if ScanFilter::all_match_binary(predicates, &mp)? { + kept.push(row); + } + } + *rows = kept; + Ok(()) +} + /// Sort finalized group rows by the post-aggregate ORDER BY terms. /// -/// Each row is a `serde_json::Value::Object` keyed by output column name. A -/// key naming one of those columns reads straight out of the row; a computed +/// A key naming an output column reads straight out of the row; a computed /// key (`ORDER BY 1000 / SUM(amount)`) is evaluated against it, with the /// planner having already bound each aggregate call to the column it lands in. /// @@ -45,20 +78,20 @@ pub(super) fn apply_user_aliases_to_rows( /// the error can propagate, rather than inside the comparator. Keys missing /// from a row sort as NULL, placed by the key's NULLS FIRST/LAST setting. The /// sort is stable to preserve relative order of equal-key rows. -pub(super) fn sort_aggregated_rows( - rows: &mut [serde_json::Value], +pub(in crate::data::executor::handlers) fn sort_aggregated_rows( + rows: &mut [Value], sort_keys: &[nodedb_physical::physical_plan::SortKeySpec], ) -> crate::Result<()> { if sort_keys.is_empty() { return Ok(()); } - let keyed: Vec> = rows + let keyed: Vec> = rows .iter() .map(|row| { sort_keys .iter() - .map(|k| nodedb_query::eval_expr_on_json(&k.expr, row).map_err(crate::Error::from)) + .map(|k| k.expr.eval(row).map_err(crate::Error::from)) .collect::>>() }) .collect::>>()?; @@ -66,20 +99,22 @@ pub(super) fn sort_aggregated_rows( let mut order: Vec = (0..rows.len()).collect(); order.sort_by(|&a, &b| { for (idx, key) in sort_keys.iter().enumerate() { - let av = keyed[a].get(idx); - let bv = keyed[b].get(idx); - let ord = match key.order_nulls( - matches!(av, None | Some(serde_json::Value::Null)), - matches!(bv, None | Some(serde_json::Value::Null)), - ) { - Some(ord) => ord, - None => key.direct(compare_json_values(av, bv)), + let av = keyed[a].get(idx).filter(|v| !v.is_null()); + let bv = keyed[b].get(idx).filter(|v| !v.is_null()); + let ord = match (av, bv) { + (Some(x), Some(y)) => { + key.direct(nodedb_query::value_ops::compare_sort_values(x, y)) + } + // At least one side is NULL, so `order_nulls` returns `Some`. + _ => key + .order_nulls(av.is_none(), bv.is_none()) + .unwrap_or(Ordering::Equal), }; - if ord != std::cmp::Ordering::Equal { + if ord != Ordering::Equal { return ord; } } - std::cmp::Ordering::Equal + Ordering::Equal }); let original = rows.to_vec(); @@ -89,34 +124,156 @@ pub(super) fn sort_aggregated_rows( Ok(()) } -/// Compare two `Option<&serde_json::Value>` for sort. Nulls / absent -/// keys sort last; numbers compare numerically; everything else falls -/// back to string comparison. -fn compare_json_values( - a: Option<&serde_json::Value>, - b: Option<&serde_json::Value>, -) -> std::cmp::Ordering { - use serde_json::Value as V; - use std::cmp::Ordering; - let a_is_null = matches!(a, None | Some(V::Null)); - let b_is_null = matches!(b, None | Some(V::Null)); - if a_is_null && b_is_null { - return Ordering::Equal; - } - if a_is_null { - return Ordering::Greater; - } - if b_is_null { - return Ordering::Less; - } - match (a.unwrap(), b.unwrap()) { - (V::Number(x), V::Number(y)) => { - let xf = x.as_f64().unwrap_or(0.0); - let yf = y.as_f64().unwrap_or(0.0); - xf.partial_cmp(&yf).unwrap_or(Ordering::Equal) - } - (V::String(x), V::String(y)) => x.cmp(y), - (V::Bool(x), V::Bool(y)) => x.cmp(y), - (x, y) => x.to_string().cmp(&y.to_string()), +#[cfg(test)] +mod tests { + use super::*; + use nodedb_physical::physical_plan::SortKeySpec; + use std::collections::HashMap; + + fn row(pairs: &[(&str, Value)]) -> Value { + Value::Object( + pairs + .iter() + .map(|(k, v)| ((*k).to_string(), v.clone())) + .collect::>(), + ) + } + + fn sorted(mut rows: Vec, key: SortKeySpec) -> Vec { + sort_aggregated_rows(&mut rows, &[key]).expect("sort"); + rows + } + + fn column(rows: &[Value], name: &str) -> Vec { + rows.iter() + .map(|r| match r { + Value::Object(m) => m.get(name).cloned().unwrap_or(Value::Null), + _ => Value::Null, + }) + .collect() + } + + fn s(text: &str) -> Value { + Value::String(text.into()) + } + + /// `2^53 + 1` and `2^53` collapse to one `f64`. MAX outputs order exactly. + #[test] + fn max_outputs_past_two_pow_53_order_exactly() { + let rows = vec![ + row(&[ + ("g", s("a")), + ("max_v", Value::Integer(9_007_199_254_740_993)), + ]), + row(&[ + ("g", s("b")), + ("max_v", Value::Integer(9_007_199_254_740_992)), + ]), + ]; + let asc = sorted(rows.clone(), SortKeySpec::column("max_v", true)); + assert_eq!(column(&asc, "g"), vec![s("b"), s("a")]); + let desc = sorted(rows, SortKeySpec::column("max_v", false)); + assert_eq!(column(&desc, "g"), vec![s("a"), s("b")]); + } + + #[test] + fn min_outputs_nanosecond_timestamps_order_exactly() { + let rows = vec![ + row(&[ + ("g", s("late")), + ("min_ts", Value::Integer(1_700_000_000_000_000_002)), + ]), + row(&[ + ("g", s("early")), + ("min_ts", Value::Integer(1_700_000_000_000_000_001)), + ]), + row(&[ + ("g", s("mid")), + ("min_ts", Value::Float(1_700_000_000_000_000_001.5)), + ]), + ]; + let asc = sorted(rows, SortKeySpec::column("min_ts", true)); + // The float rounds to `1_700_000_000_000_000_000`, below both integers. + assert_eq!(column(&asc, "g"), vec![s("mid"), s("early"), s("late")]); + } + + #[test] + fn u64_outputs_order_above_i64() { + let rows = vec![ + row(&[("g", s("u")), ("v", Value::from_u64(u64::MAX))]), + row(&[("g", s("i")), ("v", Value::Integer(i64::MAX))]), + row(&[("g", s("neg")), ("v", Value::Integer(i64::MIN))]), + ]; + let asc = sorted(rows, SortKeySpec::column("v", true)); + assert_eq!(column(&asc, "g"), vec![s("neg"), s("i"), s("u")]); + } + + #[test] + fn small_values_and_nulls_keep_their_order() { + let rows = vec![ + row(&[("g", s("two")), ("v", Value::Integer(2))]), + row(&[("g", s("null")), ("v", Value::Null)]), + row(&[("g", s("half")), ("v", Value::Float(1.5))]), + row(&[("g", s("absent"))]), + row(&[("g", s("one")), ("v", Value::Integer(1))]), + ]; + let asc = sorted(rows.clone(), SortKeySpec::column("v", true)); + assert_eq!( + column(&asc, "g"), + vec![s("one"), s("half"), s("two"), s("null"), s("absent")] + ); + let desc = sorted(rows, SortKeySpec::column("v", false)); + assert_eq!( + column(&desc, "g"), + vec![s("null"), s("absent"), s("two"), s("half"), s("one")] + ); + } + + /// NaN and ±Infinity outputs keep their float values and order as in + /// PostgreSQL: NaN above every number. + #[test] + fn non_finite_outputs_order_as_postgres() { + let rows = vec![ + row(&[("g", s("nan")), ("v", Value::Float(f64::NAN))]), + row(&[("g", s("inf")), ("v", Value::Float(f64::INFINITY))]), + row(&[("g", s("one")), ("v", Value::Integer(1))]), + row(&[("g", s("neg")), ("v", Value::Float(f64::NEG_INFINITY))]), + ]; + let asc = sorted(rows.clone(), SortKeySpec::column("v", true)); + assert_eq!( + column(&asc, "g"), + vec![s("neg"), s("one"), s("inf"), s("nan")] + ); + let desc = sorted(rows, SortKeySpec::column("v", false)); + assert_eq!( + column(&desc, "g"), + vec![s("nan"), s("inf"), s("one"), s("neg")] + ); + } + + #[test] + fn text_outputs_order_by_bytes() { + let rows = vec![ + row(&[("g", s("10"))]), + row(&[("g", s("9"))]), + row(&[("g", s("abc"))]), + ]; + let asc = sorted(rows, SortKeySpec::column("g", true)); + assert_eq!(column(&asc, "g"), vec![s("10"), s("9"), s("abc")]); + } + + #[test] + fn user_aliases_rename_output_columns() { + let mut rows = vec![row(&[("sum(v)", Value::Float(f64::INFINITY))])]; + let spec = AggregateSpec { + function: "sum".into(), + alias: "sum(v)".into(), + user_alias: Some("total".into()), + field: "v".into(), + expr: None, + }; + apply_user_aliases_to_rows(&mut rows, std::slice::from_ref(&spec)); + assert_eq!(column(&rows, "total"), vec![Value::Float(f64::INFINITY)]); + assert_eq!(column(&rows, "sum(v)"), vec![Value::Null]); } } diff --git a/nodedb/src/data/executor/handlers/aggregate/shuffle_merge.rs b/nodedb/src/data/executor/handlers/aggregate/shuffle_merge.rs index 397ae69f2..151c8b3e6 100644 --- a/nodedb/src/data/executor/handlers/aggregate/shuffle_merge.rs +++ b/nodedb/src/data/executor/handlers/aggregate/shuffle_merge.rs @@ -84,9 +84,9 @@ pub(in crate::data::executor) fn merge_state_frames( detail: format!("shuffle-aggregate `{AGG_STATE_FIELD}` is not binary"), })?; - // GroupState was serialized via serde (sonic_rs JSON) on the producer. + // The producer encodes GroupState as MessagePack. let state: GroupState = - sonic_rs::from_slice(state_bytes).map_err(|e| crate::Error::Codec { + zerompk::from_msgpack(state_bytes).map_err(|e| crate::Error::Codec { detail: format!("shuffle-aggregate partial-state decode: {e}"), })?; @@ -136,12 +136,7 @@ impl CoreLoop { let merged = match merge_state_frames(Path::new(state_path), group_by, aggregates) { Ok(m) => m, Err(e) => { - return self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ); + return self.response_error(task, ErrorCode::from(e)); } }; @@ -158,12 +153,7 @@ impl CoreLoop { sort_keys, }) { Ok(payload) => self.response_with_payload(task, payload), - Err(e) => self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ), + Err(e) => self.response_error(task, ErrorCode::from(e)), } } } @@ -238,7 +228,7 @@ mod tests { row.insert(spec.output_name.clone(), Value::from(jv)); part_idx += 1; } - let state_bytes = sonic_rs::to_vec(&state).expect("state json"); + let state_bytes = zerompk::to_msgpack_vec(&state).expect("state msgpack"); row.insert( super::AGG_STATE_FIELD.to_string(), Value::Bytes(state_bytes), @@ -270,7 +260,16 @@ mod tests { .expect("feed"); } map.into_iter() - .map(|(k, s)| (k, s.finalize(specs).into_iter().map(|(_, v)| v).collect())) + .map(|(k, s)| { + ( + k, + s.finalize(specs) + .expect("finalize") + .into_iter() + .map(|(_, v)| v) + .collect(), + ) + }) .collect() } @@ -332,7 +331,16 @@ mod tests { let merged = merge_state_frames(&combined, &group_by, &specs).expect("merge"); let got: HashMap> = merged .into_iter() - .map(|(k, s)| (k, s.finalize(&specs).into_iter().map(|(_, v)| v).collect())) + .map(|(k, s)| { + ( + k, + s.finalize(&specs) + .expect("finalize") + .into_iter() + .map(|(_, v)| v) + .collect(), + ) + }) .collect(); // Reference: single pass over the union of both doc sets. diff --git a/nodedb/src/data/executor/handlers/aggregate/state_emit.rs b/nodedb/src/data/executor/handlers/aggregate/state_emit.rs index 304a4c652..17235b619 100644 --- a/nodedb/src/data/executor/handlers/aggregate/state_emit.rs +++ b/nodedb/src/data/executor/handlers/aggregate/state_emit.rs @@ -80,12 +80,7 @@ impl CoreLoop { ) { Ok(d) => d, Err(e) => { - return self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ); + return self.response_error(task, ErrorCode::from(e)); } } }; @@ -104,35 +99,20 @@ impl CoreLoop { }) { Ok(g) => g, Err(e) => { - return self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ); + return self.response_error(task, ErrorCode::from(e)); } }; let rows = match Self::partial_state_rows(groups, group_by) { Ok(r) => r, Err(e) => { - return self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ); + return self.response_error(task, ErrorCode::from(e)); } }; match crate::data::executor::response_codec::encode_value_vec(&rows) { Ok(payload) => self.response_with_payload(task, payload), - Err(e) => self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ), + Err(e) => self.response_error(task, ErrorCode::from(e)), } } @@ -175,9 +155,9 @@ impl CoreLoop { } } - // GroupState serializes via serde (sonic_rs JSON) — the same - // canonical encoding `GroupBySpiller` already uses to persist it. - let state_bytes = sonic_rs::to_vec(&state).map_err(|e| crate::Error::Codec { + // GroupState encodes as MessagePack, the same encoding + // `GroupBySpiller` uses to persist it. NaN and ±Infinity stay exact. + let state_bytes = zerompk::to_msgpack_vec(&state).map_err(|e| crate::Error::Codec { detail: format!("partial-state serialize: {e}"), })?; map.insert(AGG_STATE_FIELD.to_string(), Value::Bytes(state_bytes)); diff --git a/nodedb/src/data/executor/handlers/aggregate/streaming/finalize.rs b/nodedb/src/data/executor/handlers/aggregate/streaming/finalize.rs index ec52a2730..fad5a4f0d 100644 --- a/nodedb/src/data/executor/handlers/aggregate/streaming/finalize.rs +++ b/nodedb/src/data/executor/handlers/aggregate/streaming/finalize.rs @@ -10,7 +10,9 @@ use std::collections::HashMap; -use super::super::rows::{apply_user_aliases_to_rows, sort_aggregated_rows}; +use nodedb_types::Value; + +use super::super::rows::{apply_user_aliases_to_rows, retain_having, sort_aggregated_rows}; use crate::bridge::scan_filter::ScanFilter; use crate::data::executor::core_loop::CoreLoop; use crate::data::executor::handlers::accum::GroupState; @@ -78,10 +80,10 @@ impl CoreLoop { groups.insert("__all__".to_string(), GroupState::new(aggregates)); } - let mut results: Vec = Vec::new(); + let mut results: Vec = Vec::new(); for (group_key, state) in groups { - let mut row = serde_json::Map::new(); + let mut row: HashMap = HashMap::new(); if !group_by.is_empty() && let Ok(parts) = sonic_rs::from_str::>(&group_key) @@ -98,46 +100,41 @@ impl CoreLoop { let val = parts .get(part_idx) .cloned() - .unwrap_or(serde_json::Value::Null); + .map_or(Value::Null, Value::from); row.insert(spec.output_name.clone(), val); part_idx += 1; } } - for (alias, val) in state.finalize(aggregates) { - let json_val: serde_json::Value = val.into(); - row.insert(alias, json_val); + for (alias, val) in state.finalize(aggregates)? { + row.insert(alias, val); } if need_sub { let sub_map = sub_groups.remove(&group_key).unwrap_or_default(); - let mut sub_results: Vec = Vec::new(); + let mut sub_results: Vec = Vec::new(); for (sub_key, sub_state) in sub_map { - let mut sub_row = serde_json::Map::new(); + let mut sub_row: HashMap = HashMap::new(); if let Ok(parts) = sonic_rs::from_str::>(&sub_key) { for (i, field) in sub_group_by.iter().enumerate() { - let val = parts.get(i).cloned().unwrap_or(serde_json::Value::Null); + let val = parts.get(i).cloned().map_or(Value::Null, Value::from); sub_row.insert(field.clone(), val); } } - for (alias, val) in sub_state.finalize(sub_aggregates) { - let json_val: serde_json::Value = val.into(); - sub_row.insert(alias, json_val); + for (alias, val) in sub_state.finalize(sub_aggregates)? { + sub_row.insert(alias, val); } - let mut sub_value = serde_json::Value::Object(sub_row); + let mut sub_value = Value::Object(sub_row); apply_user_aliases_to_rows( std::slice::from_mut(&mut sub_value), sub_aggregates, ); sub_results.push(sub_value); } - row.insert( - "sub_groups".to_string(), - serde_json::Value::Array(sub_results), - ); + row.insert("sub_groups".to_string(), Value::Array(sub_results)); } - results.push(serde_json::Value::Object(row)); + results.push(Value::Object(row)); } if !having.is_empty() { @@ -152,29 +149,7 @@ impl CoreLoop { Vec::new() } }; - if !having_predicates.is_empty() { - // `Vec::retain`'s closure must return `bool`, so an evaluation - // error in a HAVING predicate is captured via this side-channel - // and checked once the retain finishes. - let predicate_err: std::cell::RefCell> = - std::cell::RefCell::new(None); - results.retain(|row| { - if predicate_err.borrow().is_some() { - return true; - } - let mp = nodedb_types::json_to_msgpack_or_empty(row); - match ScanFilter::all_match_binary(&having_predicates, &mp) { - Ok(keep) => keep, - Err(e) => { - predicate_err.replace(Some(e)); - true - } - } - }); - if let Some(e) = predicate_err.take() { - return Err(crate::Error::from(e)); - } - } + retain_having(&mut results, &having_predicates)?; } apply_user_aliases_to_rows(&mut results, aggregates); @@ -183,6 +158,6 @@ impl CoreLoop { sort_aggregated_rows(&mut results, sort_keys)?; results.truncate(limit); - crate::data::executor::response_codec::encode_json_vec_as_msgpack(&results) + crate::data::executor::response_codec::encode_value_vec(&results) } } diff --git a/nodedb/src/data/executor/handlers/aggregate/streaming/over_docs.rs b/nodedb/src/data/executor/handlers/aggregate/streaming/over_docs.rs index e1e282bab..df6bbb771 100644 --- a/nodedb/src/data/executor/handlers/aggregate/streaming/over_docs.rs +++ b/nodedb/src/data/executor/handlers/aggregate/streaming/over_docs.rs @@ -133,12 +133,7 @@ impl CoreLoop { } self.response_with_payload(task, payload) } - Err(e) => self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ), + Err(e) => self.response_error(task, ErrorCode::from(e)), } } } diff --git a/nodedb/src/data/executor/handlers/columnar_agg.rs b/nodedb/src/data/executor/handlers/columnar_agg.rs index a9f32fcfd..947cdb42e 100644 --- a/nodedb/src/data/executor/handlers/columnar_agg.rs +++ b/nodedb/src/data/executor/handlers/columnar_agg.rs @@ -17,6 +17,8 @@ use std::collections::HashMap; +use nodedb_types::Value; + use crate::engine::timeseries::columnar_memtable::{ColumnData, ColumnType, ColumnarMemtable}; use nodedb_query::agg_key::canonical_agg_key; @@ -27,9 +29,10 @@ use super::columnar_agg_support::{ use super::columnar_filter; use super::spill::columnar::ColumnarGroupBySpiller; -/// Result of native columnar aggregation. +/// Result of native columnar aggregation: one `Value::Object` per group, +/// keyed by output column name. pub(super) struct ColumnarAggResult { - pub rows: Vec, + pub rows: Vec, } /// Parameters for [`try_columnar_aggregate`]. @@ -58,16 +61,46 @@ pub(super) struct ColumnarAggParams<'a> { /// `p.spill_dir` and `p.spill_cap` control GROUP BY spill-to-disk: if the /// number of distinct groups exceeds `spill_cap`, partial accumulators are /// spilled to `spill_dir` and k-way merged at finalize time. -pub(super) fn try_columnar_aggregate(p: &ColumnarAggParams<'_>) -> Option { +/// +/// Fails with `EvalError::NumericOverflow` when an exact integer SUM total +/// lies outside the `Decimal` range. +pub(super) fn try_columnar_aggregate( + p: &ColumnarAggParams<'_>, +) -> Result, nodedb_query::EvalError> { + let Some(acc) = accumulate_groups(p) else { + return Ok(None); + }; + build_results_from_groups( + &acc.groups, + p.group_by, + &acc.group_col_info, + p.aggregates, + p.mt, + p.limit, + ) + .map(Some) +} + +/// Per-group accumulators and the resolved GROUP BY column info. +struct AccumulatedGroups { + groups: HashMap>, + group_col_info: Vec<(usize, ColumnType)>, +} + +/// Filter, group, and accumulate the memtable rows. `None` when the query +/// cannot run natively. +fn accumulate_groups(p: &ColumnarAggParams<'_>) -> Option { let (mt, group_by, aggregates, filters) = (p.mt, p.group_by, p.aggregates, p.filters); - let (limit, scan_limit, spill_dir, spill_cap) = - (p.limit, p.scan_limit, p.spill_dir, p.spill_cap); + let (scan_limit, spill_dir, spill_cap) = (p.scan_limit, p.spill_dir, p.spill_cap); let (db, tenant) = (p.db, p.tenant); let schema = mt.schema(); let row_count = (mt.row_count() as usize).min(scan_limit); if row_count == 0 { - return Some(ColumnarAggResult { rows: Vec::new() }); + return Some(AccumulatedGroups { + groups: HashMap::new(), + group_col_info: Vec::new(), + }); } // --- Phase 1: Resolve column indices for group-by and aggregate fields --- @@ -221,16 +254,8 @@ pub(super) fn try_columnar_aggregate(p: &ColumnarAggParams<'_>) -> Option accums[agg_idx].feed_count_only(), Some((_, col_data)) => { - let val = match col_data { - ColumnData::Float64(vals) => vals[row_idx], - ColumnData::Int64(vals) => vals[row_idx] as f64, - ColumnData::Timestamp(vals) => vals[row_idx] as f64, - _ => return, - }; - if op == "count" { - accums[agg_idx].feed_count_only(); - } else { - accums[agg_idx].feed(val); + if !accums[agg_idx].feed_op(op, col_data, row_idx) { + return; } } } @@ -253,14 +278,10 @@ pub(super) fn try_columnar_aggregate(p: &ColumnarAggParams<'_>) -> Option) -> Option accums[agg_idx].feed_count_only(), Some((_, col_data)) => { - let val = match col_data { - ColumnData::Float64(vals) => vals[row_idx], - ColumnData::Int64(vals) => vals[row_idx] as f64, - ColumnData::Timestamp(vals) => vals[row_idx] as f64, - _ => return true, - }; - if op == "count" { - accums[agg_idx].feed_count_only(); - } else { - accums[agg_idx].feed(val); + if !accums[agg_idx].feed_op(op, col_data, row_idx) { + return true; } } } @@ -343,44 +356,37 @@ pub(super) fn try_columnar_aggregate(p: &ColumnarAggParams<'_>) -> Option>, group_by: &[String], - group_col_info: &[( - usize, - crate::engine::timeseries::columnar_memtable::ColumnType, - )], + group_col_info: &[(usize, ColumnType)], aggregates: &[(String, String)], - mt: &crate::engine::timeseries::columnar_memtable::ColumnarMemtable, + mt: &ColumnarMemtable, limit: usize, -) -> ColumnarAggResult { - let mut results: Vec = Vec::with_capacity(groups.len().min(limit)); +) -> Result { + let mut results: Vec = Vec::with_capacity(groups.len().min(limit)); for (group_key, accums) in groups { - let mut row = serde_json::Map::new(); + let mut row: HashMap = HashMap::new(); for (i, field) in group_by.iter().enumerate() { let (col_idx, _) = group_col_info[i]; let val = if i < group_key.len() { resolve_key_part(mt, col_idx, &group_key[i]) } else { - serde_json::Value::Null + Value::Null }; row.insert(field.clone(), val); } @@ -389,47 +395,23 @@ fn build_results_from_groups( let agg_key = canonical_agg_key(op, field); let accum = &accums[agg_idx]; let val = match op.as_str() { - "count" => serde_json::json!(accum.count), - "sum" => { - if accum.count == 0 { - serde_json::Value::Null - } else { - serde_json::json!(accum.sum) - } - } - "avg" => { - if accum.count == 0 { - serde_json::Value::Null - } else { - serde_json::json!(accum.sum / accum.count as f64) - } - } - "min" => { - if accum.count == 0 { - serde_json::Value::Null - } else { - serde_json::json!(accum.min) - } - } - "max" => { - if accum.count == 0 { - serde_json::Value::Null - } else { - serde_json::json!(accum.max) - } - } - _ => serde_json::Value::Null, + "count" => Value::from_u64(accum.count), + "sum" => accum.sum.sum()?, + "avg" => accum.sum.avg()?, + "min" => accum.min.clone().unwrap_or(Value::Null), + "max" => accum.max.clone().unwrap_or(Value::Null), + _ => Value::Null, }; row.insert(agg_key, val); } - results.push(serde_json::Value::Object(row)); + results.push(Value::Object(row)); if results.len() >= limit { break; } } - ColumnarAggResult { rows: results } + Ok(ColumnarAggResult { rows: results }) } #[cfg(test)] @@ -506,11 +488,12 @@ mod tests { db: crate::types::DatabaseId::DEFAULT, tenant: crate::types::TenantId::new(1), }) + .unwrap() .unwrap(); assert_eq!(result.rows.len(), 2); // A and AAAA for row in &result.rows { - let count = row.get("count(*)").and_then(|v| v.as_u64()).unwrap(); + let count = row.get("count(*)").and_then(|v| v.as_i64()).unwrap(); assert_eq!(count, 50); // 100 rows / 2 types } } @@ -532,6 +515,7 @@ mod tests { db: crate::types::DatabaseId::DEFAULT, tenant: crate::types::TenantId::new(1), }) + .unwrap() .unwrap(); assert_eq!(result.rows.len(), 5); // 5 unique qnames @@ -561,6 +545,7 @@ mod tests { db: crate::types::DatabaseId::DEFAULT, tenant: crate::types::TenantId::new(1), }) + .unwrap() .unwrap(); // Only rows with value > 5000 (i >= 51, value >= 5100) @@ -588,13 +573,193 @@ mod tests { db: crate::types::DatabaseId::DEFAULT, tenant: crate::types::TenantId::new(1), }) + .unwrap() .unwrap(); assert_eq!(result.rows.len(), 1); let count = result.rows[0] .get("count(*)") - .and_then(|v| v.as_u64()) + .and_then(|v| v.as_i64()) .unwrap(); assert_eq!(count, 100); } + + const ABOVE: i64 = 9_007_199_254_740_993; + const AT: i64 = 9_007_199_254_740_992; + + /// A memtable with a nanosecond `Int64` column `n`, a `Float64` column + /// `f`, and a symbol `g`, one row per `(ts, n, f, g)`. + fn int_memtable(rows: &[(i64, i64, f64, &str)]) -> ColumnarMemtable { + use crate::engine::timeseries::columnar_memtable::ColumnValue; + let schema = ColumnarSchema { + columns: vec![ + ("timestamp".into(), ColumnType::Timestamp(TimeKind::Millis)), + ("n".into(), ColumnType::Int64), + ("f".into(), ColumnType::Float64), + ("g".into(), ColumnType::Symbol), + ], + timestamp_idx: 0, + codecs: vec![], + }; + let mut mt = ColumnarMemtable::new(schema, ColumnarMemtableConfig::default()); + for (i, &(ts, n, f, g)) in rows.iter().enumerate() { + let values = [ + ColumnValue::Timestamp(ts), + ColumnValue::Int64(n), + ColumnValue::Float64(f), + ColumnValue::Symbol(g.to_string()), + ]; + mt.ingest_row(i as SeriesId, &values).unwrap(); + } + mt + } + + /// Aggregate `aggregates` over `mt`, grouped by `group_by`, with a spill + /// cap of `spill_cap` groups. + fn run( + mt: &ColumnarMemtable, + group_by: &[String], + aggregates: &[(String, String)], + spill_cap: usize, + suffix: &str, + ) -> Result, nodedb_query::EvalError> { + let sd = test_spill_dir(suffix); + let result = try_columnar_aggregate(&ColumnarAggParams { + mt, + group_by, + aggregates, + filters: &[], + limit: 100, + scan_limit: 100_000, + spill_dir: &sd, + spill_cap, + governor: crate::data::executor::core_loop::test_governor(), + db: crate::types::DatabaseId::DEFAULT, + tenant: crate::types::TenantId::new(1), + })?; + Ok(result.unwrap().rows) + } + + fn aggs(pairs: &[(&str, &str)]) -> Vec<(String, String)> { + pairs + .iter() + .map(|(op, f)| (op.to_string(), f.to_string())) + .collect() + } + + /// The output value of column `name` in `row`; NULL when absent. + fn col(row: &Value, name: &str) -> Value { + row.get(name).cloned().unwrap_or(Value::Null) + } + + #[test] + fn int64_min_max_sum_stay_exact_above_2_pow_53() { + let mt = int_memtable(&[(1, ABOVE, 0.0, "a"), (2, AT, 0.0, "a")]); + let a = aggs(&[("min", "n"), ("max", "n"), ("sum", "n"), ("avg", "n")]); + // Dense-symbol path and the no-GROUP-BY hash path. + for group_by in [vec!["g".to_string()], vec![]] { + let rows = run(&mt, &group_by, &a, 1_000_000, "exact").unwrap(); + assert_eq!(rows.len(), 1); + let row = &rows[0]; + assert_eq!(col(row, "min(n)"), Value::Integer(AT)); + assert_eq!(col(row, "max(n)"), Value::Integer(ABOVE)); + assert_eq!(col(row, "sum(n)"), Value::Integer(ABOVE + AT)); + assert_eq!(col(row, "avg(n)"), Value::Float(AT as f64)); + } + } + + #[test] + fn nanosecond_timestamps_min_max_exact() { + let t = [ + 1_700_000_000_000_000_002_i64, + 1_700_000_000_000_000_001, + 1_700_000_000_000_000_003, + ]; + let mt = int_memtable(&[ + (t[0], 0, 0.0, "a"), + (t[1], 0, 0.0, "a"), + (t[2], 0, 0.0, "a"), + ]); + let a = aggs(&[("min", "timestamp"), ("max", "timestamp")]); + let rows = run(&mt, &[], &a, 1_000_000, "nanos").unwrap(); + assert_eq!(col(&rows[0], "min(timestamp)"), Value::Integer(t[1])); + assert_eq!(col(&rows[0], "max(timestamp)"), Value::Integer(t[2])); + } + + #[test] + fn int64_sum_past_i64_is_decimal_exact() { + let mt = int_memtable(&[(1, i64::MAX, 0.0, "a"), (2, i64::MAX, 0.0, "a")]); + let rows = run(&mt, &[], &aggs(&[("sum", "n")]), 1_000_000, "dec").unwrap(); + let want = 2 * i128::from(i64::MAX); + assert_eq!( + col(&rows[0], "sum(n)"), + Value::Decimal(rust_decimal::Decimal::from_i128_with_scale(want, 0)) + ); + } + + /// NaN sorts above every number, so MAX is NaN and MIN skips it, and a + /// SUM over NaN is NaN, as in PostgreSQL. The float results stay floats. + #[test] + fn float_min_max_sum_keep_nan() { + let mt = int_memtable(&[(1, 0, f64::NAN, "a"), (2, 0, 1.5, "a"), (3, 0, -2.5, "a")]); + let a = aggs(&[("min", "f"), ("max", "f"), ("sum", "f")]); + let rows = run(&mt, &[], &a, 1_000_000, "nan").unwrap(); + assert_eq!(col(&rows[0], "min(f)"), Value::Float(-2.5)); + assert!(matches!(col(&rows[0], "max(f)"), Value::Float(f) if f.is_nan())); + assert!(matches!(col(&rows[0], "sum(f)"), Value::Float(f) if f.is_nan())); + } + + /// Float overflow is `Infinity` and `Infinity` plus `-Infinity` is NaN; + /// both reach the result row as floats, also through spill runs. + #[test] + fn float_sum_keeps_infinity_through_spill() { + let mt = int_memtable(&[ + (1, 0, 1e308, "a"), + (2, 0, 1e308, "a"), + (3, 0, f64::INFINITY, "b"), + (4, 0, f64::NEG_INFINITY, "b"), + ]); + let a = aggs(&[("sum", "f"), ("avg", "f"), ("max", "f")]); + for spill_cap in [1_000_000, 1] { + let rows = run(&mt, &["g".to_string()], &a, spill_cap, "inf").unwrap(); + assert_eq!(rows.len(), 2); + for row in &rows { + match col(row, "g") { + Value::String(g) if g == "a" => { + assert_eq!(col(row, "sum(f)"), Value::Float(f64::INFINITY)); + assert_eq!(col(row, "avg(f)"), Value::Float(f64::INFINITY)); + assert_eq!(col(row, "max(f)"), Value::Float(1e308)); + } + Value::String(g) if g == "b" => { + assert!(matches!(col(row, "sum(f)"), Value::Float(f) if f.is_nan())); + assert_eq!(col(row, "max(f)"), Value::Float(f64::INFINITY)); + } + other => panic!("unexpected group {other:?}"), + } + } + } + } + + /// A spill cap of one group forces every group through spill runs, so + /// the partial merge must keep the integer totals exact. + #[test] + fn spill_merge_keeps_integers_exact() { + let mt = int_memtable(&[ + (1, ABOVE, 0.0, "a"), + (2, AT, 0.0, "b"), + (3, AT, 0.0, "a"), + (4, ABOVE, 0.0, "b"), + ]); + let a = aggs(&[("min", "n"), ("max", "n"), ("sum", "n")]); + let rows = run(&mt, &["timestamp".to_string()], &a, 1, "spill").unwrap(); + assert_eq!(rows.len(), 4); + let rows = run(&mt, &["n".to_string()], &a, 1, "spill_n").unwrap(); + assert_eq!(rows.len(), 2); + for row in &rows { + let n = col(row, "n").as_i64().unwrap(); + assert_eq!(col(row, "min(n)"), Value::Integer(n)); + assert_eq!(col(row, "max(n)"), Value::Integer(n)); + assert_eq!(col(row, "sum(n)"), Value::Integer(2 * n)); + } + } } diff --git a/nodedb/src/data/executor/handlers/columnar_agg_support.rs b/nodedb/src/data/executor/handlers/columnar_agg_support.rs index ec71cb21d..d82738c13 100644 --- a/nodedb/src/data/executor/handlers/columnar_agg_support.rs +++ b/nodedb/src/data/executor/handlers/columnar_agg_support.rs @@ -2,6 +2,8 @@ //! Supporting types and low-level routines for columnar aggregation. +use nodedb_query::window::extremum::value_replaces; + use crate::engine::timeseries::columnar_memtable::{ColumnData, ColumnType, ColumnarMemtable}; /// Iterate over every set bit in a packed `u64` bitmask, calling `f(row_idx)`. @@ -32,43 +34,83 @@ pub(in crate::data::executor::handlers) fn for_each_set_bit( } /// Accumulator for running aggregate computation per group. -#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)] +/// +/// SUM / AVG total exactly per `ExactSum`: an `Int64` or `Timestamp` cell +/// never rounds through `f64`. MIN / MAX keep the original cell, compared +/// exactly, so an integer column returns an integer. Spill runs encode it +/// as MessagePack, which keeps NaN and ±Infinity exact. +#[derive(Debug, Clone, Default, zerompk::ToMessagePack, zerompk::FromMessagePack)] pub(in crate::data::executor::handlers) struct AggAccum { pub count: u64, - pub sum: f64, - pub min: f64, - pub max: f64, + pub sum: nodedb_query::ExactSum, + pub min: Option, + pub max: Option, } impl AggAccum { pub(in crate::data::executor::handlers) fn new() -> Self { - Self { - count: 0, - sum: 0.0, - min: f64::INFINITY, - max: f64::NEG_INFINITY, + Self::default() + } + + /// Feed aggregate `op` the cell at `row_idx` of a numeric column: + /// `count` counts it, any other op folds its value. Returns `false` for + /// a non-numeric column, which feeds nothing; the caller stops feeding + /// the row. + pub(in crate::data::executor::handlers) fn feed_op( + &mut self, + op: &str, + col_data: &ColumnData, + row_idx: usize, + ) -> bool { + let cell = match col_data { + ColumnData::Float64(vals) => nodedb_types::Value::Float(vals[row_idx]), + ColumnData::Int64(vals) => nodedb_types::Value::Integer(vals[row_idx]), + ColumnData::Timestamp(vals) => nodedb_types::Value::Integer(vals[row_idx]), + _ => return false, + }; + if op == "count" { + self.feed_count_only(); + } else { + self.feed(cell); } + true } - pub(in crate::data::executor::handlers) fn feed(&mut self, val: f64) { + fn feed(&mut self, cell: nodedb_types::Value) { self.count += 1; - self.sum += val; - if val < self.min { - self.min = val; + self.sum.add_value(&cell); + if value_replaces(&cell, self.min.as_ref(), false) { + self.min = Some(cell.clone()); } - if val > self.max { - self.max = val; + if value_replaces(&cell, self.max.as_ref(), true) { + self.max = Some(cell); } } pub(in crate::data::executor::handlers) fn feed_count_only(&mut self) { self.count += 1; } + + /// Fold a partial accumulator (a spilled run) into this one without loss. + pub(in crate::data::executor::handlers) fn merge(&mut self, other: AggAccum) { + self.count += other.count; + self.sum.merge(&other.sum); + if let Some(min) = other.min + && value_replaces(&min, self.min.as_ref(), false) + { + self.min = Some(min); + } + if let Some(max) = other.max + && value_replaces(&max, self.max.as_ref(), true) + { + self.max = Some(max); + } + } } /// A group key composed of symbol IDs (for Symbol columns) or raw i64/f64 /// values (for numeric group-by columns). Avoids string allocation entirely. -#[derive(Debug, Clone, PartialEq, Eq, Hash, serde::Serialize, serde::Deserialize)] +#[derive(Debug, Clone, PartialEq, Eq, Hash, zerompk::ToMessagePack, zerompk::FromMessagePack)] pub(in crate::data::executor::handlers) enum GroupKeyPart { SymbolId(u32), Int64(i64), @@ -118,26 +160,23 @@ pub(in crate::data::executor::handlers) fn extract_group_key_part( } } -/// Resolve a group key part to a serde_json::Value for output. +/// Resolve a group key part to its output value. A float key keeps NaN and +/// ±Infinity. pub(in crate::data::executor::handlers) fn resolve_key_part( mt: &ColumnarMemtable, col_idx: usize, part: &GroupKeyPart, -) -> serde_json::Value { +) -> nodedb_types::Value { match part { GroupKeyPart::SymbolId(id) => mt .symbol_dict(col_idx) .and_then(|dict| dict.get(*id)) - .map(|s| serde_json::Value::String(s.to_string())) - .unwrap_or(serde_json::Value::Null), - GroupKeyPart::Int64(v) => serde_json::Value::Number(serde_json::Number::from(*v)), - GroupKeyPart::Float64Bits(bits) => { - let v = f64::from_bits(*bits); - serde_json::Number::from_f64(v) - .map(serde_json::Value::Number) - .unwrap_or(serde_json::Value::Null) - } - GroupKeyPart::Null => serde_json::Value::Null, + .map_or(nodedb_types::Value::Null, |s| { + nodedb_types::Value::String(s.to_string()) + }), + GroupKeyPart::Int64(v) => nodedb_types::Value::Integer(*v), + GroupKeyPart::Float64Bits(bits) => nodedb_types::Value::Float(f64::from_bits(*bits)), + GroupKeyPart::Null => nodedb_types::Value::Null, } } @@ -181,16 +220,8 @@ pub(in crate::data::executor::handlers) fn aggregate_dense_symbol( match &p.agg_col_data[agg_idx] { None => accums[agg_idx].feed_count_only(), Some((_, col_data)) => { - let val = match col_data { - ColumnData::Float64(vals) => vals[row_idx], - ColumnData::Int64(vals) => vals[row_idx] as f64, - ColumnData::Timestamp(vals) => vals[row_idx] as f64, - _ => return, - }; - if op == "count" { - accums[agg_idx].feed_count_only(); - } else { - accums[agg_idx].feed(val); + if !accums[agg_idx].feed_op(op, col_data, row_idx) { + return; } } } diff --git a/nodedb/src/data/executor/handlers/document/read/projection.rs b/nodedb/src/data/executor/handlers/document/read/projection.rs index 947c27987..5906a6814 100644 --- a/nodedb/src/data/executor/handlers/document/read/projection.rs +++ b/nodedb/src/data/executor/handlers/document/read/projection.rs @@ -2,66 +2,19 @@ //! Projection and computed-column application for document scans. //! -//! Two parallel paths must agree on missing-key semantics: -//! a projection key absent from the row maps to SQL NULL on **both** the JSON -//! and the msgpack code paths. Silently skipping a missing key drops the -//! column from the response and breaks the pgwire RowDescription contract. +//! A projection key absent from the row maps to SQL NULL. Skipping a missing +//! key would drop the column from the response and break the pgwire +//! RowDescription contract. use crate::bridge::expr_eval::ComputedColumn; -/// Apply projection or computed columns to a decoded document. -/// -/// Missing projection keys are emitted as `Value::Null`, mirroring -/// [`apply_projection_msgpack`]. Earlier revisions silently skipped them, -/// which produced responses where window-function aliases (or any computed -/// alias not yet present on the row) disappeared from the output without an -/// error. -pub(in crate::data::executor) fn apply_projection( - data: serde_json::Value, - computed_cols: &[ComputedColumn], - projection: &[String], -) -> crate::Result { - Ok(match data { - serde_json::Value::Object(obj) => { - if computed_cols.is_empty() && projection.is_empty() { - return Ok(serde_json::Value::Object(obj)); - } - - let doc_val = nodedb_types::Value::from(serde_json::Value::Object(obj.clone())); - let mut out = if projection.is_empty() { - serde_json::Map::with_capacity(computed_cols.len()) - } else { - let mut projected = - serde_json::Map::with_capacity(projection.len() + computed_cols.len()); - for col in projection { - let val = obj.get(col).cloned().unwrap_or(serde_json::Value::Null); - projected.insert(col.clone(), val); - } - projected - }; - - for cc in computed_cols { - let existing = out.get(&cc.alias); - if matches!(existing, Some(v) if !v.is_null()) { - continue; - } - // A division/modulo-by-zero in a computed column fails the - // whole query instead of silently materializing NULL into - // the response. - let v = cc.expr.eval(&doc_val)?; - out.insert(cc.alias.clone(), serde_json::Value::from(v)); - } - - serde_json::Value::Object(out) - } - other => other, - }) -} - /// Apply projection and computed columns on raw msgpack bytes. /// /// For projection-only (no computed columns), uses zero-decode binary field extraction. -/// For computed columns, decodes fields on-demand from msgpack. +/// For computed columns, decodes fields on-demand from msgpack. A missing +/// projection key is written as NULL. A computed column whose alias is also +/// projected keeps the projected value, so a window alias already on the +/// row is not overwritten. pub(in crate::data::executor) fn apply_projection_msgpack( data: &[u8], computed_cols: &[ComputedColumn], @@ -71,11 +24,13 @@ pub(in crate::data::executor) fn apply_projection_msgpack( return Ok(data.to_vec()); } - let field_count = if projection.is_empty() { - computed_cols.len() - } else { - projection.len() + computed_cols.len() - }; + // A computed column whose alias is projected writes no entry of its own, + // so the map header counts only the computed columns that do. + let written_computed = computed_cols + .iter() + .filter(|cc| !projection.iter().any(|p| p == &cc.alias)) + .count(); + let field_count = projection.len() + written_computed; let mut buf = Vec::with_capacity(data.len()); nodedb_query::msgpack_scan::write_map_header(&mut buf, field_count); @@ -92,7 +47,13 @@ pub(in crate::data::executor) fn apply_projection_msgpack( } if !computed_cols.is_empty() { - let doc_val = nodedb_types::value_from_msgpack(data).unwrap_or(nodedb_types::Value::Null); + // A body that does not decode fails the query: evaluating the + // columns against `Null` would ship NULL for every one of them. + let doc_val = + nodedb_types::value_from_msgpack(data).map_err(|e| crate::Error::Serialization { + format: "msgpack".into(), + detail: format!("computed columns: stored row does not decode: {e}"), + })?; for cc in computed_cols { let already_present = projection.iter().any(|p| p == &cc.alias); if already_present { @@ -103,11 +64,13 @@ pub(in crate::data::executor) fn apply_projection_msgpack( // query instead of silently materializing NULL into the // response. let result = cc.expr.eval(&doc_val)?; - if let Ok(mp) = nodedb_types::value_to_msgpack(&result) { - buf.extend_from_slice(&mp); - } else { - nodedb_query::msgpack_scan::write_null(&mut buf); - } + let encoded = nodedb_types::value_to_msgpack(&result).map_err(|e| { + crate::Error::Serialization { + format: "msgpack".into(), + detail: format!("computed column '{}' does not encode: {e}", cc.alias), + } + })?; + buf.extend_from_slice(&encoded); } } @@ -118,39 +81,62 @@ pub(in crate::data::executor) fn apply_projection_msgpack( mod tests { use super::*; use crate::bridge::expr_eval::SqlExpr; + use nodedb_types::Value; + + fn row(doc: Value) -> Vec { + nodedb_types::value_to_msgpack(&doc).expect("encode row") + } + + fn object(pairs: &[(&str, Value)]) -> Value { + Value::Object( + pairs + .iter() + .map(|(k, v)| ((*k).to_string(), v.clone())) + .collect(), + ) + } + + fn project( + doc: Value, + computed: &[ComputedColumn], + projection: &[String], + ) -> crate::Result { + let bytes = apply_projection_msgpack(&row(doc), computed, projection)?; + Ok(nodedb_types::value_from_msgpack(&bytes).expect("decode projected row")) + } #[test] - fn apply_projection_keeps_base_fields_when_computed_columns_exist() { - let data = serde_json::json!({ - "id": "u1", - "name": "Ada", - "age": 42 - }); + fn projection_keeps_base_fields_when_computed_columns_exist() { + let data = object(&[ + ("id", Value::from("u1")), + ("name", Value::from("Ada")), + ("age", Value::Integer(42)), + ]); let computed = vec![ComputedColumn { alias: "label".into(), expr: SqlExpr::Column("name".into()), }]; let projection = vec!["name".to_string(), "age".to_string()]; - let projected = apply_projection(data, &computed, &projection).unwrap(); + let projected = project(data, &computed, &projection).unwrap(); assert_eq!( projected, - serde_json::json!({ - "name": "Ada", - "age": 42, - "label": "Ada" - }) + object(&[ + ("name", Value::from("Ada")), + ("age", Value::Integer(42)), + ("label", Value::from("Ada")), + ]) ); } #[test] - fn apply_projection_does_not_overwrite_existing_window_alias() { - let data = serde_json::json!({ - "name": "Ada", - "age": 42, - "rn": 1 - }); + fn projection_does_not_overwrite_existing_window_alias() { + let data = object(&[ + ("name", Value::from("Ada")), + ("age", Value::Integer(42)), + ("rn", Value::Integer(1)), + ]); let computed = vec![ComputedColumn { alias: "rn".into(), expr: SqlExpr::Function { @@ -160,57 +146,85 @@ mod tests { }]; let projection = vec!["name".to_string(), "age".to_string(), "rn".to_string()]; - let projected = apply_projection(data, &computed, &projection).unwrap(); + let projected = project(data, &computed, &projection).unwrap(); assert_eq!( projected, - serde_json::json!({ - "name": "Ada", - "age": 42, - "rn": 1 - }) + object(&[ + ("name", Value::from("Ada")), + ("age", Value::Integer(42)), + ("rn", Value::Integer(1)), + ]) ); } + /// `SELECT to_jsonb(*) AS document`: the whole stored row, every field + /// with its own type, under one alias. + #[test] + fn whole_row_computed_column_carries_every_field() { + let data = object(&[ + ("id", Value::from("u1")), + ("n", Value::Integer(5)), + ("ratio", Value::Float(1.5)), + ]); + let computed = vec![ComputedColumn { + alias: "document".into(), + expr: SqlExpr::Function { + name: "to_jsonb".into(), + args: vec![SqlExpr::Column(nodedb_query::expr::WHOLE_ROW_COLUMN.into())], + }, + }]; + + let projected = project(data.clone(), &computed, &[]).unwrap(); + + assert_eq!(projected, object(&[("document", data)])); + } + #[test] - fn apply_projection_emits_null_for_missing_keys() { - let data = serde_json::json!({ - "id": "u1", - "score": 1.0 - }); + fn projection_emits_null_for_missing_keys() { + let data = object(&[("id", Value::from("u1")), ("score", Value::Float(1.0))]); let projection = vec![ "id".to_string(), "score".to_string(), "pr_score".to_string(), ]; - let projected = apply_projection(data, &[], &projection).unwrap(); + let projected = project(data, &[], &projection).unwrap(); assert_eq!( projected, - serde_json::json!({ - "id": "u1", - "score": 1.0, - "pr_score": serde_json::Value::Null, - }) + object(&[ + ("id", Value::from("u1")), + ("score", Value::Float(1.0)), + ("pr_score", Value::Null), + ]) ); } + /// A non-finite window result on the row reaches the projected row as a + /// float. + #[test] + fn projection_keeps_non_finite_floats() { + let data = object(&[("total", Value::Float(f64::INFINITY))]); + let projected = project(data, &[], &["total".to_string()]).unwrap(); + assert_eq!(projected, object(&[("total", Value::Float(f64::INFINITY))])); + } + /// A computed column that divides by zero fails the projection instead /// of silently materializing `NULL`. #[test] - fn apply_projection_computed_column_division_by_zero_errors() { + fn projection_computed_column_division_by_zero_errors() { use crate::bridge::expr_eval::BinaryOp; - let data = serde_json::json!({"denom": 0}); + let data = object(&[("denom", Value::Integer(0))]); let computed = vec![ComputedColumn { alias: "bad".into(), expr: SqlExpr::BinaryOp { - left: Box::new(SqlExpr::Literal(nodedb_types::Value::Integer(10))), + left: Box::new(SqlExpr::Literal(Value::Integer(10))), op: BinaryOp::Div, right: Box::new(SqlExpr::Column("denom".into())), }, }]; - let err = apply_projection(data, &computed, &[]).unwrap_err(); + let err = project(data, &computed, &[]).unwrap_err(); assert!(matches!(err, crate::Error::DivisionByZero), "got {err:?}"); } } diff --git a/nodedb/src/data/executor/handlers/document/read/scan.rs b/nodedb/src/data/executor/handlers/document/read/scan.rs index bfbbbc17a..8f0d101ee 100644 --- a/nodedb/src/data/executor/handlers/document/read/scan.rs +++ b/nodedb/src/data/executor/handlers/document/read/scan.rs @@ -5,14 +5,14 @@ use nodedb_types::StorageKey; use tracing::{debug, warn}; -use super::projection::{apply_projection, apply_projection_msgpack}; +use super::projection::apply_projection_msgpack; use super::{DocFetchParams, DocScanMode}; use crate::bridge::envelope::{ErrorCode, Response}; use crate::bridge::scan_filter::ScanFilter; use crate::data::executor::core_loop::CoreLoop; use crate::data::executor::handlers::document::sort; +use crate::data::executor::handlers::provider_scan_compute::apply_windows_and_computed; use crate::data::executor::handlers::transaction::overlay::SidecarRowShape; -use crate::data::executor::response_codec::DocumentRow; use crate::data::executor::scan_normalize::sparse_row_to_doc; use crate::data::executor::sparse_body_format::{SparseBodyFormat, SparseBodyFormatRef}; use crate::data::executor::task::ExecutionTask; @@ -71,14 +71,34 @@ impl CoreLoop { if window_functions_bytes.is_empty() { Vec::new() } else { - zerompk::from_msgpack(window_functions_bytes).unwrap_or_default() + match zerompk::from_msgpack(window_functions_bytes) { + Ok(specs) => specs, + Err(e) => { + return self.response_error( + task, + ErrorCode::Internal { + detail: format!("malformed scan window functions: {e}"), + }, + ); + } + } }; let computed_cols: Vec = if computed_columns_bytes.is_empty() { Vec::new() } else { - zerompk::from_msgpack(computed_columns_bytes).unwrap_or_default() + match zerompk::from_msgpack(computed_columns_bytes) { + Ok(specs) => specs, + Err(e) => { + return self.response_error( + task, + ErrorCode::Internal { + detail: format!("malformed scan computed columns: {e}"), + }, + ); + } + } }; let scan_budget_bytes = self.query_tuning.max_scan_result_bytes; @@ -105,6 +125,8 @@ impl CoreLoop { crate::types::TenantId::new(tid), collection.to_string(), ); + let identity_column = + self.identity_column(task.request.database_id.as_u64(), tid, collection); let strict_schema = self.doc_configs.get(&config_key).and_then(|c| { if let nodedb_physical::physical_plan::StorageMode::Strict { ref schema } = c.storage_mode @@ -114,6 +136,9 @@ impl CoreLoop { None } }); + // A `DECIMAL` sort key orders by value: its cells are text and wide + // integers side by side. + let decimal_keys = sort::decimal_sort_keys(sort_keys, strict_schema.as_ref()); // Fetch stage: the ONLY part that differs between a current-time read // and a bitemporal `AS OF` / all-versions audit read. It returns the @@ -194,6 +219,7 @@ impl CoreLoop { &filter_predicates, row_key, value, + &identity_column, ) { Ok(b) => b, Err(e) => { @@ -275,7 +301,9 @@ impl CoreLoop { let filtered: Vec<(String, Vec)> = if normalizes { filtered .into_iter() - .map(|(key, bytes)| sparse_row_to_doc(&key, &bytes, body_format)) + .map(|(key, bytes)| { + sparse_row_to_doc(&key, &bytes, body_format, &identity_column) + }) .collect() } else { filtered @@ -284,21 +312,40 @@ impl CoreLoop { .collect() }; - // With window functions the sort runs after the window pass - // (below), so ORDER BY can name a window alias, as the - // provider scan orders it. - let sorted = if sort_keys.is_empty() || !window_specs.is_empty() { + // With window functions the window pass runs first and the + // sort after it, so ORDER BY can name a window alias, as the + // provider scan orders it. The windowed rows stay msgpack, so + // a NaN or ±Infinity window result reaches the client. + let sorted = if !window_specs.is_empty() { + let (ids, bodies): (Vec, Vec>) = filtered.into_iter().unzip(); + let bodies = + match apply_windows_and_computed(bodies, window_functions_bytes, &[]) { + Ok(bodies) => bodies, + Err(e) => return self.response_error(task, e), + }; + let mut windowed: Vec<(String, Vec)> = + ids.into_iter().zip(bodies).collect(); + if let Err(e) = sort::sort_rows(&mut windowed, sort_keys, &decimal_keys) { + return self.response_error(task, e); + } + windowed + } else if sort_keys.is_empty() { filtered } else if filtered.len() <= self.query_tuning.sort_run_size { let mut v = filtered; // Propagate the typed error: a zero divisor in a sort key // is a `22012` statement failure, not an internal fault. - if let Err(e) = sort::sort_rows(&mut v, sort_keys) { + if let Err(e) = sort::sort_rows(&mut v, sort_keys, &decimal_keys) { return self.response_error(task, e); } v } else { - match self.external_sort(filtered, sort_keys, limit.saturating_add(offset)) { + match self.external_sort( + filtered, + sort_keys, + &decimal_keys, + limit.saturating_add(offset), + ) { Ok(merged) => merged, Err(e) => { warn!(core = self.core_id, error = %e, "external sort failed"); @@ -351,103 +398,49 @@ impl CoreLoop { return self.send_document_rows_raw(task, &result, stream_chunk_size); } - if !window_specs.is_empty() { - let mut decoded_rows: Vec<(String, serde_json::Value)> = match sorted - .into_iter() - .map(|(doc_id, mp)| { - crate::data::executor::doc_format::decode_document(&mp) - .map(|doc| (doc_id, doc)) - }) - .collect::>>() - { - Ok(rows) => rows, - Err(e) => return self.response_error(task, e), - }; - if let Err(e) = crate::bridge::window_func::evaluate_window_functions( - &mut decoded_rows, - &window_specs, - ) { - return self.response_error(task, crate::Error::from(e)); - } - let decoded_rows = match sort::sort_decoded_rows(decoded_rows, sort_keys) { - Ok(rows) => rows, - Err(e) => return self.response_error(task, e), - }; + let needs_transform = !computed_cols.is_empty() || !projection.is_empty(); - // Project first, then dedupe on the projected JSON value - // so `SELECT DISTINCT col` honours SQL semantics. - let projected_rows: Vec<_> = match decoded_rows + if needs_transform { + // Project first so DISTINCT acts on the projected + // row, not the raw document. + let projected_rows: Vec<_> = match sorted .into_iter() - .map(|(doc_id, data)| { - let projected = apply_projection(data, &computed_cols, projection)?; - Ok(DocumentRow { - id: doc_id, - data: projected, - }) + .map(|(doc_id, mp)| { + let projected = + apply_projection_msgpack(&mp, &computed_cols, projection)?; + Ok((doc_id, projected)) }) .collect::>>() { Ok(rows) => rows, Err(e) => return self.response_error(task, e), }; - - let deduped: Vec<_> = if distinct { + let deduped = if distinct { let mut seen = std::collections::HashSet::new(); projected_rows .into_iter() - .filter(|row| seen.insert(row.data.to_string())) + .filter(|(_, value)| seen.insert(value.clone())) .collect() } else { projected_rows }; - let result: Vec<_> = deduped.into_iter().skip(offset).take(limit).collect(); - self.send_document_rows_transformed(task, &result, stream_chunk_size) + self.send_document_rows_raw(task, &result, stream_chunk_size) } else { - let needs_transform = !computed_cols.is_empty() || !projection.is_empty(); - - if needs_transform { - // Project first so DISTINCT acts on the projected - // row, not the raw document. - let projected_rows: Vec<_> = match sorted + // No projection — `SELECT DISTINCT *` semantics dedupe + // on the entire raw value, which is what the + // pre-existing path does. + let deduped = if distinct { + let mut seen = std::collections::HashSet::new(); + sorted .into_iter() - .map(|(doc_id, mp)| { - let projected = - apply_projection_msgpack(&mp, &computed_cols, projection)?; - Ok((doc_id, projected)) - }) - .collect::>>() - { - Ok(rows) => rows, - Err(e) => return self.response_error(task, e), - }; - let deduped = if distinct { - let mut seen = std::collections::HashSet::new(); - projected_rows - .into_iter() - .filter(|(_, value)| seen.insert(value.clone())) - .collect() - } else { - projected_rows - }; - let result: Vec<_> = deduped.into_iter().skip(offset).take(limit).collect(); - self.send_document_rows_raw(task, &result, stream_chunk_size) + .filter(|(_, value)| seen.insert(value.clone())) + .collect() } else { - // No projection — `SELECT DISTINCT *` semantics dedupe - // on the entire raw value, which is what the - // pre-existing path does. - let deduped = if distinct { - let mut seen = std::collections::HashSet::new(); - sorted - .into_iter() - .filter(|(_, value)| seen.insert(value.clone())) - .collect() - } else { - sorted - }; - let rows: Vec<_> = deduped.into_iter().skip(offset).take(limit).collect(); - self.send_document_rows_raw(task, &rows, stream_chunk_size) - } + sorted + }; + let rows: Vec<_> = deduped.into_iter().skip(offset).take(limit).collect(); + self.send_document_rows_raw(task, &rows, stream_chunk_size) } } Err(e) => self.response_error(task, e), diff --git a/nodedb/src/data/executor/handlers/grouping_sets_exec.rs b/nodedb/src/data/executor/handlers/grouping_sets_exec.rs index f209f9062..e6f80ad42 100644 --- a/nodedb/src/data/executor/handlers/grouping_sets_exec.rs +++ b/nodedb/src/data/executor/handlers/grouping_sets_exec.rs @@ -12,9 +12,11 @@ use std::collections::HashMap; +use nodedb_types::Value; use sonic_rs; use super::accum::GroupState; +use super::aggregate::rows::{apply_user_aliases_to_rows, retain_having}; use crate::bridge::envelope::{ErrorCode, Response}; use crate::bridge::scan_filter::ScanFilter; use crate::data::executor::core_loop::CoreLoop; @@ -99,12 +101,7 @@ pub(super) fn execute_grouping_sets( let owned_docs: Vec> = match docs_result { Ok(docs) => docs.into_iter().map(|(_, v)| v.to_vec()).collect(), Err(e) => { - return core.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ); + return core.response_error(task, ErrorCode::from(e)); } }; @@ -121,7 +118,7 @@ pub(super) fn execute_grouping_sets( .filter(|a| a.function == "grouping") .collect(); - let mut all_rows: Vec = Vec::new(); + let mut all_rows: Vec = Vec::new(); for set in grouping_sets { // Build the active key list for this set. @@ -201,7 +198,7 @@ pub(super) fn execute_grouping_sets( } for (group_key, state) in groups { - let mut row = serde_json::Map::new(); + let mut row: HashMap = HashMap::new(); // Parse the group key (JSON array of active key values). let active_values: Vec = @@ -223,16 +220,20 @@ pub(super) fn execute_grouping_sets( let val = active_values .get(pos) .cloned() - .unwrap_or(serde_json::Value::Null); + .map_or(Value::Null, Value::from); row.insert(key.clone(), val); } else { - row.insert(key.clone(), serde_json::Value::Null); + row.insert(key.clone(), Value::Null); } } // Real aggregate results. - for (alias, val) in state.finalize(&real_agg_slice) { - row.insert(alias, val.into()); + let finalized = match state.finalize(&real_agg_slice) { + Ok(finalized) => finalized, + Err(e) => return core.response_error(task, ErrorCode::from(e)), + }; + for (alias, val) in finalized { + row.insert(alias, val); } // GROUPING(col) pseudo-aggregates: compute from the bitmask. @@ -243,85 +244,28 @@ pub(super) fn execute_grouping_sets( let bit_present = (grouping_id >> col_idx) & 1; let grouping_val = 1u64 ^ bit_present; let output_alias = agg.user_alias.as_deref().unwrap_or(&agg.alias); - row.insert( - output_alias.to_string(), - serde_json::Value::Number(serde_json::Number::from(grouping_val)), - ); + row.insert(output_alias.to_string(), Value::from_u64(grouping_val)); } // Hidden grouping bitmask column — available for downstream use. - row.insert( - GROUPING_ID_COL.to_string(), - serde_json::Value::Number(serde_json::Number::from(grouping_id)), - ); + row.insert(GROUPING_ID_COL.to_string(), Value::from_u64(grouping_id)); - all_rows.push(serde_json::Value::Object(row)); + all_rows.push(Value::Object(row)); } } - // Apply HAVING. - if !having_predicates.is_empty() { - // `Vec::retain`'s closure must return `bool`, so an evaluation error - // in a HAVING predicate is captured via this side-channel and checked - // once the retain finishes. - let predicate_err: std::cell::RefCell> = - std::cell::RefCell::new(None); - all_rows.retain(|row| { - if predicate_err.borrow().is_some() { - return true; - } - let mp = nodedb_types::json_to_msgpack_or_empty(row); - match ScanFilter::all_match_binary(&having_predicates, &mp) { - Ok(keep) => keep, - Err(e) => { - predicate_err.replace(Some(e)); - true - } - } - }); - if let Some(e) = predicate_err.take() { - return core.response_error(task, ErrorCode::from(e)); - } + // Apply HAVING. An evaluation error in a predicate fails the statement. + if let Err(e) = retain_having(&mut all_rows, &having_predicates) { + return core.response_error(task, ErrorCode::from(e)); } // Apply user aliases for real aggregates (grouping aliases were applied above). - apply_user_aliases(&mut all_rows, &real_agg_slice); + apply_user_aliases_to_rows(&mut all_rows, &real_agg_slice); all_rows.truncate(limit); - match super::super::response_codec::encode_json_vec_as_msgpack(&all_rows) { + match super::super::response_codec::encode_value_vec(&all_rows) { Ok(payload) => core.response_with_payload(task, payload), - Err(e) => core.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ), - } -} - -fn apply_user_aliases(rows: &mut [serde_json::Value], aggregates: &[AggregateSpec]) { - let renames: Vec<(&str, &str)> = aggregates - .iter() - .filter_map(|agg| { - agg.user_alias - .as_deref() - .filter(|alias| *alias != agg.alias) - .map(|alias| (agg.alias.as_str(), alias)) - }) - .collect(); - - if renames.is_empty() { - return; - } - - for row in rows { - if let Some(obj) = row.as_object_mut() { - for (from, to) in &renames { - if let Some(value) = obj.remove(*from) { - obj.insert((*to).to_string(), value); - } - } - } + Err(e) => core.response_error(task, ErrorCode::from(e)), } } diff --git a/nodedb/src/data/executor/handlers/provider_scan_compute.rs b/nodedb/src/data/executor/handlers/provider_scan_compute.rs index 5dc27c063..18470c44f 100644 --- a/nodedb/src/data/executor/handlers/provider_scan_compute.rs +++ b/nodedb/src/data/executor/handlers/provider_scan_compute.rs @@ -3,9 +3,10 @@ //! Window-function and computed-column evaluation for `QueryOp::ProviderScan`. //! //! Runs after filter and before sort/distinct/offset/project/limit in the -//! `ProviderScan` pipeline. Each msgpack row decodes to a `serde_json::Value`, +//! `ProviderScan` pipeline. Each msgpack row decodes to a `nodedb_types::Value`, //! windows evaluate over the full row set, computed columns evaluate -//! per-row, and the result re-encodes to msgpack. Skipped entirely when both +//! per-row, and the result re-encodes to msgpack. `Value` keeps NaN and +//! ±Infinity results, which a JSON number cannot. Skipped entirely when both //! byte slices are empty, so the zero-decode msgpack path stays untouched for //! a plain relational scan. @@ -50,43 +51,43 @@ pub(in crate::data::executor) fn apply_windows_and_computed( decode_bytes(computed_bytes, "computed")? }; - let mut json_rows: Vec<(String, serde_json::Value)> = Vec::with_capacity(rows.len()); + let mut value_rows: Vec<(String, nodedb_types::Value)> = Vec::with_capacity(rows.len()); for (idx, row) in rows.iter().enumerate() { let value = nodedb_types::value_from_msgpack(row).map_err(|e| { crate::Error::DataPlane(ErrorCode::Internal { detail: format!("ProviderScan: malformed row for window/computed evaluation: {e}"), }) })?; - json_rows.push((idx.to_string(), serde_json::Value::from(value))); + value_rows.push((idx.to_string(), value)); } if !window_specs.is_empty() { - evaluate_window_functions(&mut json_rows, &window_specs).map_err(crate::Error::from)?; + evaluate_window_functions(&mut value_rows, &window_specs).map_err(crate::Error::from)?; } - for (_, row_json) in &mut json_rows { + for (_, row) in &mut value_rows { if computed_cols.is_empty() { continue; } // Every computed column evaluates against the row as it stood before - // this loop, matching `apply_projection`'s semantics: later computed + // this loop, matching `apply_projection_msgpack`'s semantics: later computed // columns never observe earlier ones' results. - let doc_val = nodedb_types::Value::from(row_json.clone()); + let before = row.clone(); for cc in &computed_cols { - let already_present = matches!(row_json.get(&cc.alias), Some(v) if !v.is_null()); + let already_present = matches!(row.get(&cc.alias), Some(v) if !v.is_null()); if already_present { continue; } - let v = cc.expr.eval(&doc_val)?; - if let serde_json::Value::Object(obj) = row_json { - obj.insert(cc.alias.clone(), serde_json::Value::from(v)); + let v = cc.expr.eval(&before)?; + if let nodedb_types::Value::Object(obj) = row { + obj.insert(cc.alias.clone(), v); } } } - let mut out = Vec::with_capacity(json_rows.len()); - for (_, row_json) in json_rows { - let bytes = nodedb_types::json_to_msgpack(&row_json).map_err(|e| { + let mut out = Vec::with_capacity(value_rows.len()); + for (_, row) in value_rows { + let bytes = nodedb_types::value_to_msgpack(&row).map_err(|e| { crate::Error::DataPlane(ErrorCode::Internal { detail: format!( "ProviderScan: failed to re-encode row after window/computed evaluation: {e}" diff --git a/nodedb/src/data/executor/handlers/spill/columnar.rs b/nodedb/src/data/executor/handlers/spill/columnar.rs index 1a1961bc8..f03c53817 100644 --- a/nodedb/src/data/executor/handlers/spill/columnar.rs +++ b/nodedb/src/data/executor/handlers/spill/columnar.rs @@ -4,9 +4,8 @@ //! //! The columnar path uses `GroupKey = Vec` (integer/symbol IDs) //! rather than the JSON-encoded string keys used by the schemaless path. -//! Both `GroupKey` and `Vec` derive `serde::Serialize + -//! Deserialize`, so they serialize directly without intermediate string -//! encoding — no unwrap fallbacks, no double-serialization. +//! Both `GroupKey` and `Vec` derive the `zerompk` MessagePack +//! traits, so they encode directly without intermediate string encoding. use std::collections::HashMap; use std::path::PathBuf; @@ -99,14 +98,7 @@ impl ColumnarGroupBySpiller { ) -> crate::Result>> { self.core.merge(&mut self.in_mem, self.cap, |dst, src| { for (d, s) in dst.iter_mut().zip(src) { - d.count += s.count; - d.sum += s.sum; - if s.min < d.min { - d.min = s.min; - } - if s.max > d.max { - d.max = s.max; - } + d.merge(s); } }) } diff --git a/nodedb/src/data/executor/handlers/spill/core.rs b/nodedb/src/data/executor/handlers/spill/core.rs index a00bc36ba..2bf1c5811 100644 --- a/nodedb/src/data/executor/handlers/spill/core.rs +++ b/nodedb/src/data/executor/handlers/spill/core.rs @@ -39,8 +39,9 @@ const FINALIZE_CAP_FACTOR: usize = 10; /// Generic spill-to-disk manager for a `HashMap`. /// -/// Each spill run is serialized (JSON, via `sonic_rs`) into its own file -/// inside `spill_dir`. On `merge()`, all runs plus any remaining in-memory +/// Each spill run is encoded as MessagePack (via `zerompk`) into its own +/// file inside `spill_dir`. MessagePack keeps every float exact, NaN and +/// ±Infinity included. On `merge()`, all runs plus any remaining in-memory /// entries are folded together using a caller-supplied merge function. pub(super) struct SpillCore { spill_dir: PathBuf, @@ -53,8 +54,8 @@ pub(super) struct SpillCore { impl SpillCore where - K: serde::Serialize + serde::de::DeserializeOwned + Eq + Hash, - V: serde::Serialize + serde::de::DeserializeOwned, + K: zerompk::ToMessagePack + for<'de> zerompk::FromMessagePack<'de> + Eq + Hash, + V: zerompk::ToMessagePack + for<'de> zerompk::FromMessagePack<'de>, { pub(super) fn new(spill_dir: PathBuf) -> crate::Result { std::fs::create_dir_all(&spill_dir).map_err(|e| crate::Error::Storage { @@ -78,7 +79,7 @@ where return Ok(()); } - let encoded = sonic_rs::to_vec(&entries).map_err(|e| crate::Error::Storage { + let encoded = zerompk::to_msgpack_vec(&entries).map_err(|e| crate::Error::Storage { engine: "groupby_spill".into(), detail: format!("spill serialize error: {e}"), })?; @@ -129,7 +130,7 @@ where for run_path in &self.runs { let buf = read_run_file(&mut reader, run_path)?; let entries: Vec<(K, V)> = - sonic_rs::from_slice(&buf).map_err(|e| crate::Error::Storage { + zerompk::from_msgpack(&buf).map_err(|e| crate::Error::Storage { engine: "groupby_spill".into(), detail: format!("spill run deserialize error: {e}"), })?; @@ -317,6 +318,73 @@ mod tests { assert_eq!(out.len(), 2); } + /// A group state holding NaN and ±Infinity spills, reads back, and + /// merges with every non-finite value intact. + #[test] + fn spill_keeps_non_finite_group_state() { + use crate::data::executor::handlers::columnar_agg_support::AggAccum; + use nodedb_types::Value; + + fn state(v: f64) -> Vec { + let mut acc = AggAccum::new(); + acc.count = 1; + acc.sum.add_f64(v); + acc.min = Some(Value::Float(v)); + acc.max = Some(Value::Float(v)); + vec![acc] + } + fn is_float(v: Option<&Value>, want: f64) -> bool { + matches!(v, Some(Value::Float(f)) if f == &want || (f.is_nan() && want.is_nan())) + } + + let dir = tempfile::tempdir().unwrap(); + let mut core: SpillCore> = + SpillCore::new(dir.path().join("sc")).unwrap(); + core.flush_run( + vec![ + ("g".to_string(), state(f64::INFINITY)), + ("n".to_string(), state(f64::NAN)), + ("o".to_string(), state(1e308)), + ] + .into_iter(), + ) + .unwrap(); + core.flush_run( + vec![ + ("g".to_string(), state(f64::NEG_INFINITY)), + ("n".to_string(), state(1.0)), + ] + .into_iter(), + ) + .unwrap(); + let mut in_mem: HashMap> = HashMap::new(); + in_mem.insert("o".to_string(), state(1e308)); + + let out = core + .merge(&mut in_mem, 100, |dst, src| { + for (d, s) in dst.iter_mut().zip(src) { + d.merge(s); + } + }) + .unwrap(); + + // Infinity plus -Infinity is NaN; the extremes keep both infinities. + let g = &out["g"][0]; + assert_eq!(g.count, 2); + assert!(is_float(Some(&g.sum.sum().unwrap()), f64::NAN)); + assert!(is_float(g.min.as_ref(), f64::NEG_INFINITY)); + assert!(is_float(g.max.as_ref(), f64::INFINITY)); + // NaN survives the run and is the largest extreme. + let n = &out["n"][0]; + assert!(is_float(Some(&n.sum.sum().unwrap()), f64::NAN)); + assert!(is_float(n.min.as_ref(), 1.0)); + assert!(is_float(n.max.as_ref(), f64::NAN)); + // A run total plus the in-memory total overflows to Infinity. + let o = &out["o"][0]; + assert_eq!(o.count, 2); + assert!(is_float(Some(&o.sum.sum().unwrap()), f64::INFINITY)); + } + /// Exceeding `cap × FINALIZE_CAP_FACTOR` distinct keys returns a /// deterministic error rather than growing unbounded. #[test] diff --git a/nodedb/src/data/executor/handlers/spill/groupby.rs b/nodedb/src/data/executor/handlers/spill/groupby.rs index b9fd23ac9..81d799854 100644 --- a/nodedb/src/data/executor/handlers/spill/groupby.rs +++ b/nodedb/src/data/executor/handlers/spill/groupby.rs @@ -191,7 +191,16 @@ mod tests { .unwrap(); } map.into_iter() - .map(|(k, s)| (k, s.finalize(specs).into_iter().map(|(_, v)| v).collect())) + .map(|(k, s)| { + ( + k, + s.finalize(specs) + .expect("finalize") + .into_iter() + .map(|(_, v)| v) + .collect(), + ) + }) .collect() } @@ -201,7 +210,16 @@ mod tests { ) -> HashMap> { result .into_iter() - .map(|(k, s)| (k, s.finalize(specs).into_iter().map(|(_, v)| v).collect())) + .map(|(k, s)| { + ( + k, + s.finalize(specs) + .expect("finalize") + .into_iter() + .map(|(_, v)| v) + .collect(), + ) + }) .collect() } diff --git a/nodedb/src/data/executor/handlers/timeseries/encode.rs b/nodedb/src/data/executor/handlers/timeseries/encode.rs index 5d0d9a269..d73df93ca 100644 --- a/nodedb/src/data/executor/handlers/timeseries/encode.rs +++ b/nodedb/src/data/executor/handlers/timeseries/encode.rs @@ -7,6 +7,7 @@ use nodedb_query::agg_key::canonical_agg_key; use crate::data::executor::core_loop::TsGroupKeyKind; use crate::data::executor::handlers::columnar_read::rmpv_time_cell; use crate::engine::timeseries::columnar_memtable::TimeKind; +use crate::util::rmpv_value::value_to_rmpv; /// The wire types of a grouped result's key columns. pub(in crate::data::executor) struct GroupedKeyTypes<'a> { @@ -136,14 +137,17 @@ pub(in crate::data::executor) fn encode_grouped_results( for (agg_idx, agg_key) in agg_keys.iter().enumerate() { let accum = &accums[agg_idx]; let op = &aggregates[agg_idx].0; + // SUM is exact; MIN / MAX / FIRST / LAST keep the cell's own + // type, so an integer column stays an integer. + let cell = |v: Option<&nodedb_types::Value>| v.map_or(rmpv::Value::Nil, value_to_rmpv); let val = match op.as_str() { "count" => rmpv::Value::Integer((accum.count as i64).into()), - "sum" if accum.count > 0 => rmpv::Value::F64(accum.sum()), - "avg" if accum.count > 0 => rmpv::Value::F64(accum.sum() / accum.count as f64), - "min" if accum.count > 0 => rmpv::Value::F64(accum.min), - "max" if accum.count > 0 => rmpv::Value::F64(accum.max), - "first" if accum.count > 0 => rmpv::Value::F64(accum.first()), - "last" if accum.count > 0 => rmpv::Value::F64(accum.last()), + "sum" => value_to_rmpv(&accum.sum_value()?), + "avg" => accum.avg_f64()?.map_or(rmpv::Value::Nil, rmpv::Value::F64), + "min" => cell(accum.min()), + "max" => cell(accum.max()), + "first" => cell(accum.first()), + "last" => cell(accum.last()), "stddev" | "ts_stddev" if accum.count >= 2 => { rmpv::Value::F64(accum.stddev_population()) } @@ -228,6 +232,148 @@ mod tests { ); } + /// Encode one ungrouped row whose accumulators were fed `feed`, and + /// return the aggregate cells by key. + fn aggregate_cells( + ops: &[&str], + feed: impl Fn(&mut crate::engine::timeseries::columnar_agg::AggAccum), + ) -> Vec<(String, rmpv::Value)> { + let mut result = GroupedAggResult::new(ops.len()); + let accums = ops + .iter() + .map(|_| { + let mut a = crate::engine::timeseries::columnar_agg::AggAccum::default(); + feed(&mut a); + a + }) + .collect(); + result.groups.insert(String::new(), accums); + let aggregates: Vec<(String, String)> = ops + .iter() + .map(|op| (op.to_string(), "v".to_string())) + .collect(); + let bytes = encode_grouped_results( + &result, + &[], + &aggregates, + usize::MAX, + 0, + &[], + GroupedKeyTypes { + group_key_kinds: &[], + bucket_kind: TimeKind::Millis, + }, + ) + .expect("encode"); + let rmpv::Value::Array(rows) = + crate::util::bounded_msgpack::read_value(&bytes).expect("decode") + else { + panic!("not an array"); + }; + let rmpv::Value::Map(fields) = &rows[0] else { + panic!("not a map"); + }; + fields + .iter() + .map(|(k, v)| (k.as_str().unwrap_or_default().to_string(), v.clone())) + .collect() + } + + fn cell<'a>(cells: &'a [(String, rmpv::Value)], key: &str) -> &'a rmpv::Value { + &cells + .iter() + .find(|(k, _)| k == key) + .expect("aggregate cell") + .1 + } + + #[test] + fn integer_sum_min_max_first_last_stay_exact() { + const ABOVE: i64 = 9_007_199_254_740_993; + const AT: i64 = 9_007_199_254_740_992; + let ops = ["sum", "min", "max", "first", "last", "avg"]; + let cells = aggregate_cells(&ops, |a| { + a.feed_int(ABOVE); + a.feed_int(AT); + }); + assert_eq!( + *cell(&cells, "sum(v)"), + rmpv::Value::Integer((ABOVE + AT).into()) + ); + assert_eq!(*cell(&cells, "min(v)"), rmpv::Value::Integer(AT.into())); + assert_eq!(*cell(&cells, "max(v)"), rmpv::Value::Integer(ABOVE.into())); + assert_eq!( + *cell(&cells, "first(v)"), + rmpv::Value::Integer(ABOVE.into()) + ); + assert_eq!(*cell(&cells, "last(v)"), rmpv::Value::Integer(AT.into())); + assert_eq!(*cell(&cells, "avg(v)"), rmpv::Value::F64(AT as f64)); + } + + #[test] + fn nanosecond_timestamps_and_sum_past_i64() { + let ops = ["sum", "min", "max"]; + let cells = aggregate_cells(&ops, |a| { + a.feed_int(1_700_000_000_000_000_002); + a.feed_int(1_700_000_000_000_000_001); + a.feed_int(i64::MAX); + }); + let want = 3_400_000_000_000_000_003_i128 + i128::from(i64::MAX); + assert_eq!( + *cell(&cells, "sum(v)"), + rmpv::Value::String(want.to_string().into()) + ); + assert_eq!( + *cell(&cells, "min(v)"), + rmpv::Value::Integer(1_700_000_000_000_000_001_i64.into()) + ); + assert_eq!( + *cell(&cells, "max(v)"), + rmpv::Value::Integer(i64::MAX.into()) + ); + } + + #[test] + fn mixed_int_float_and_nan_extremes() { + let ops = ["sum", "min", "max"]; + let cells = aggregate_cells(&ops, |a| { + a.feed(f64::NAN); + a.feed_int(2); + a.feed(0.5); + }); + assert!(matches!(cell(&cells, "sum(v)"), rmpv::Value::F64(f) if f.is_nan())); + // NaN sorts above every number, as in PostgreSQL, so it is the max. + assert_eq!(*cell(&cells, "min(v)"), rmpv::Value::F64(0.5)); + assert!(matches!(cell(&cells, "max(v)"), rmpv::Value::F64(f) if f.is_nan())); + + let cells = aggregate_cells(&ops, |a| { + a.feed_int(2); + a.feed(0.5); + }); + assert_eq!(*cell(&cells, "sum(v)"), rmpv::Value::F64(2.5)); + } + + #[test] + fn partial_merge_keeps_integer_sum_exact() { + use crate::engine::timeseries::columnar_agg::AggAccum; + let mut a = AggAccum::default(); + a.feed_int(9_007_199_254_740_993); + let mut b = AggAccum::default(); + b.feed_int(9_007_199_254_740_992); + b.feed_int(1); + a.merge(&b); + assert_eq!( + a.sum_value().unwrap(), + nodedb_types::Value::Integer(18_014_398_509_481_986) + ); + assert_eq!(a.min(), Some(&nodedb_types::Value::Integer(1))); + assert_eq!( + a.max(), + Some(&nodedb_types::Value::Integer(9_007_199_254_740_993)) + ); + assert_eq!(a.last(), Some(&nodedb_types::Value::Integer(1))); + } + #[test] fn an_instant_group_key_is_a_typed_instant() { let cell = group_key_value( diff --git a/nodedb/src/data/executor/handlers/timeseries_gap_fill.rs b/nodedb/src/data/executor/handlers/timeseries_gap_fill.rs index 591659da5..e79bc85d3 100644 --- a/nodedb/src/data/executor/handlers/timeseries_gap_fill.rs +++ b/nodedb/src/data/executor/handlers/timeseries_gap_fill.rs @@ -129,7 +129,7 @@ fn find_prev_bucket_avg( && let Some(a) = accums.first() && a.count > 0 { - return Some((ts, a.sum() / a.count as f64)); + return Some((ts, a.mean_f64())); } ts -= interval; } @@ -151,7 +151,7 @@ fn find_next_bucket_avg( && let Some(a) = accums.first() && a.count > 0 { - return Some((ts, a.sum() / a.count as f64)); + return Some((ts, a.mean_f64())); } ts += interval; } diff --git a/nodedb/src/engine/timeseries/columnar_agg.rs b/nodedb/src/engine/timeseries/columnar_agg.rs index 2f6d52fed..94121f0b2 100644 --- a/nodedb/src/engine/timeseries/columnar_agg.rs +++ b/nodedb/src/engine/timeseries/columnar_agg.rs @@ -7,6 +7,10 @@ //! dispatch to SIMD kernels via `simd_agg::ts_runtime()` which //! auto-detects AVX-512 / AVX2 / NEON at startup. +use nodedb_query::ExactSum; +use nodedb_query::window::extremum::value_replaces; +use nodedb_types::Value; + use super::simd_agg::ts_runtime; /// Aggregation result for a group of rows. @@ -254,61 +258,52 @@ pub fn count_by_time_bucket(timestamps: &[i64], bucket_interval_ms: i64) -> Vec< /// Streaming accumulator for single-pass aggregation. /// -/// Maintains running sum (with Kahan compensation), min, max, count, -/// first, and last. Zero intermediate allocation. -#[derive(Debug, Clone)] +/// SUM / AVG total exactly per `ExactSum`: an integer cell never rounds +/// through `f64`. MIN / MAX / FIRST / LAST keep the original cell, so an +/// integer column returns an integer; MIN / MAX compare exactly, with NaN +/// above every number. STDDEV stays in `f64` (Welford). +#[derive(Debug, Clone, Default)] pub struct AggAccum { pub count: u64, - sum: f64, - compensation: f64, - pub min: f64, - pub max: f64, - first: f64, - last: f64, + sum: ExactSum, + min: Option, + max: Option, + first: Option, + last: Option, /// Welford's M2 for online variance/stddev computation. mean: f64, m2: f64, } -impl Default for AggAccum { - fn default() -> Self { - Self { - count: 0, - sum: 0.0, - compensation: 0.0, - min: f64::INFINITY, - max: f64::NEG_INFINITY, - first: f64::NAN, - last: f64::NAN, - mean: 0.0, - m2: 0.0, - } - } -} - impl AggAccum { - /// Feed a single value into the accumulator. + /// Feed a single float value into the accumulator. pub fn feed(&mut self, v: f64) { + self.feed_cell(Value::Float(v), v); + } + + /// Feed a single integer value into the accumulator, kept exact. + pub fn feed_int(&mut self, v: i64) { + self.feed_cell(Value::Integer(v), v as f64); + } + + /// Feed `cell`; `as_f64` is its reading for the `f64` STDDEV state. + fn feed_cell(&mut self, cell: Value, as_f64: f64) { if self.count == 0 { - self.first = v; + self.first = Some(cell.clone()); } - self.last = v; self.count += 1; - // Kahan compensated summation. - let y = v - self.compensation; - let t = self.sum + y; - self.compensation = (t - self.sum) - y; - self.sum = t; - if v < self.min { - self.min = v; + self.sum.add_value(&cell); + if value_replaces(&cell, self.min.as_ref(), false) { + self.min = Some(cell.clone()); } - if v > self.max { - self.max = v; + if value_replaces(&cell, self.max.as_ref(), true) { + self.max = Some(cell.clone()); } + self.last = Some(cell); // Welford's online variance (for stddev). - let delta = v - self.mean; + let delta = as_f64 - self.mean; self.mean += delta / self.count as f64; - let delta2 = v - self.mean; + let delta2 = as_f64 - self.mean; self.m2 += delta * delta2; } @@ -317,15 +312,17 @@ impl AggAccum { self.count += 1; } - /// Merge another accumulator into this one. + /// Merge another accumulator into this one. Integer totals stay exact. pub fn merge(&mut self, other: &AggAccum) { if other.count == 0 { return; } - if self.count == 0 { - self.first = other.first; + if self.first.is_none() { + self.first = other.first.clone(); + } + if other.last.is_some() { + self.last = other.last.clone(); } - self.last = other.last; // Chan's parallel algorithm for merging Welford variance. let n_a = self.count as f64; @@ -335,12 +332,16 @@ impl AggAccum { self.mean = (n_a * self.mean + n_b * other.mean) / (n_a + n_b); self.count += other.count; - self.sum += other.sum; - if other.min < self.min { - self.min = other.min; + self.sum.merge(&other.sum); + if let Some(min) = &other.min + && value_replaces(min, self.min.as_ref(), false) + { + self.min = Some(min.clone()); } - if other.max > self.max { - self.max = other.max; + if let Some(max) = &other.max + && value_replaces(max, self.max.as_ref(), true) + { + self.max = Some(max.clone()); } } @@ -352,28 +353,55 @@ impl AggAccum { (self.m2 / self.count as f64).max(0.0).sqrt() } - /// Convert to final `AggResult`. + /// Convert to the `f64` `AggResult` of a float series. pub fn into_agg_result(self) -> AggResult { + let reading = |v: &Option| v.as_ref().and_then(Value::as_f64).unwrap_or(f64::NAN); AggResult { count: self.count, - sum: self.sum, - min: if self.count == 0 { f64::NAN } else { self.min }, - max: if self.count == 0 { f64::NAN } else { self.max }, - first: self.first, - last: self.last, + sum: self.sum.sum_f64(), + min: reading(&self.min), + max: reading(&self.max), + first: reading(&self.first), + last: reading(&self.last), } } - pub fn sum(&self) -> f64 { - self.sum + /// Exact SUM: `Integer`, `Decimal`, or `Float` per `ExactSum`; NULL for + /// no value. + pub fn sum_value(&self) -> Result { + self.sum.sum() + } + + /// AVG from the exact total. `None` for no value. + pub fn avg_f64(&self) -> Result, nodedb_query::EvalError> { + self.sum.avg_f64() + } + + /// The `f64` running mean of the values fed (Welford), `0.0` for none. + /// For float-valued consumers such as gap-fill interpolation; AVG + /// results use [`Self::avg_f64`]. + pub fn mean_f64(&self) -> f64 { + self.mean + } + + /// The smallest value fed, as fed. + pub fn min(&self) -> Option<&Value> { + self.min.as_ref() + } + + /// The largest value fed, as fed. + pub fn max(&self) -> Option<&Value> { + self.max.as_ref() } - pub fn first(&self) -> f64 { - self.first + /// The first value fed, as fed. + pub fn first(&self) -> Option<&Value> { + self.first.as_ref() } - pub fn last(&self) -> f64 { - self.last + /// The last value fed, as fed. + pub fn last(&self) -> Option<&Value> { + self.last.as_ref() } } diff --git a/nodedb/src/engine/timeseries/continuous_agg/manager.rs b/nodedb/src/engine/timeseries/continuous_agg/manager.rs index a112f68e1..15fa538d1 100644 --- a/nodedb/src/engine/timeseries/continuous_agg/manager.rs +++ b/nodedb/src/engine/timeseries/continuous_agg/manager.rs @@ -9,7 +9,8 @@ use std::collections::HashMap; use super::definition::{ContinuousAggregateDef, RefreshPolicy}; use super::partial::PartialAggregate; -use super::refresh; +use super::refresh::{self, RefreshResult}; +use super::rollup; use super::watermark::WatermarkState; use crate::engine::timeseries::columnar_memtable::ColumnarDrainResult; @@ -19,7 +20,24 @@ type AggKey = (u64, String); /// Materialized partials for one aggregate: /// `(bucket_ts, group_key) → PartialAggregate`. -type MaterializedBuckets = HashMap<(i64, Vec), PartialAggregate>; +type MaterializedBuckets = refresh::Buckets; + +/// Whether an aggregate sourced from another aggregate takes that +/// aggregate's refreshes. A `Manual` or `Periodic` aggregate refreshes only +/// on its own trigger. +fn takes_upstream_refresh(policy: &RefreshPolicy) -> bool { + matches!(policy, RefreshPolicy::OnFlush | RefreshPolicy::OnSeal) +} + +/// Whether two definitions of one aggregate build the same partial state. +/// A change to any of these makes the materialized buckets of the old +/// definition meaningless under the new one. +fn same_shape(a: &ContinuousAggregateDef, b: &ContinuousAggregateDef) -> bool { + a.source == b.source + && a.bucket_interval_ms == b.bucket_interval_ms + && a.group_by == b.group_by + && a.aggregates == b.aggregates +} /// Manages all continuous aggregates for a timeseries engine instance. /// @@ -58,18 +76,29 @@ impl ContinuousAggregateManager { /// Idempotent: boot re-registration and a replayed post-apply both /// register a definition the core can already hold. A repeated dependency /// edge refreshes the aggregate twice per flush. + /// + /// A definition of another shape (source, bucket, GROUP BY, or + /// aggregates) replaces the old one with no materialized buckets and a + /// default watermark: buckets built for the old shape do not hold the + /// new one's columns. pub fn register(&mut self, def: ContinuousAggregateDef) { let database_id = def.database_id; let source = def.source.clone(); let name = def.name.clone(); + let key = (database_id, name.clone()); - if let Some(previous) = self.definitions.get(&(database_id, name.clone())) - && previous.source != source - && let Some(deps) = self - .dependencies - .get_mut(&(database_id, previous.source.clone())) - { - deps.retain(|n| n != &name); + if let Some(previous) = self.definitions.get(&key) { + if previous.source != source + && let Some(deps) = self + .dependencies + .get_mut(&(database_id, previous.source.clone())) + { + deps.retain(|n| n != &name); + } + if !same_shape(previous, &def) { + self.materialized.remove(&key); + self.watermarks.remove(&key); + } } self.watermarks .entry((database_id, name.clone())) @@ -119,7 +148,8 @@ impl ContinuousAggregateManager { /// Process a flush event from a source collection. /// /// Finds all aggregates that depend on `source_collection` with - /// `RefreshPolicy::OnFlush` and refreshes them incrementally. + /// `RefreshPolicy::OnFlush` and refreshes them incrementally, then rolls + /// each refresh up into the aggregates sourced from it, transitively. /// /// Returns the names of aggregates that were refreshed. pub fn on_flush( @@ -136,7 +166,6 @@ impl ContinuousAggregateManager { .unwrap_or_default(); let mut refreshed = Vec::new(); - for agg_name in &agg_names { let key = (database_id, agg_name.clone()); let Some(def) = self.definitions.get(&key) else { @@ -145,44 +174,15 @@ impl ContinuousAggregateManager { if def.refresh_policy != RefreshPolicy::OnFlush || def.stale { continue; } - - let def_clone = def.clone(); let watermark = self.watermarks.get(&key).cloned().unwrap_or_default(); - let mat = self.materialized.entry(key.clone()).or_default(); - - let result = refresh::refresh_from_drain(&def_clone, drain, &watermark, mat); - - // Update watermark. - if let Some(wm) = self.watermarks.get_mut(&key) { - wm.advance(result.max_ts, result.rows_processed, now_ms); - if let Some(o3_ts) = result.o3_min_ts { - wm.record_o3(o3_ts); - } - } - - refreshed.push(agg_name.clone()); + let result = refresh::refresh_from_drain(def, drain, &watermark); + self.apply_refresh(database_id, agg_name, result, now_ms, &mut refreshed); } - - // Multi-tier chaining: check if refreshed aggregates have downstream - // dependents within the same database. - let mut chain_refreshed = Vec::new(); - for name in &refreshed { - if let Some(downstream) = self.dependencies.get(&(database_id, name.clone())).cloned() { - for ds_name in &downstream { - if let Some(ds_def) = self.definitions.get(&(database_id, ds_name.clone())) - && ds_def.refresh_policy == RefreshPolicy::OnFlush - && !ds_def.stale - { - chain_refreshed.push(ds_name.clone()); - } - } - } - } - refreshed.extend(chain_refreshed); refreshed } - /// Manually refresh an aggregate (for Manual or Periodic policies). + /// Manually refresh an aggregate (for Manual or Periodic policies), and + /// roll the refresh up into the aggregates sourced from it. pub fn manual_refresh( &mut self, database_id: u64, @@ -191,19 +191,70 @@ impl ContinuousAggregateManager { now_ms: i64, ) { let key = (database_id, agg_name.to_string()); - let Some(def) = self.definitions.get(&key).cloned() else { + let Some(def) = self.definitions.get(&key) else { return; }; let watermark = self.watermarks.get(&key).cloned().unwrap_or_default(); - let mat = self.materialized.entry(key.clone()).or_default(); + let result = refresh::refresh_from_drain(def, drain, &watermark); + let mut refreshed = Vec::new(); + self.apply_refresh(database_id, agg_name, result, now_ms, &mut refreshed); + } - let result = refresh::refresh_from_drain(&def, drain, &watermark, mat); + /// Merge `result` into `name`'s materialized buckets and advance its + /// watermark. Then roll the refresh delta up into every non-stale + /// aggregate sourced from `name` that takes upstream refreshes, and so on + /// down the chain. Each aggregate takes one refresh per call, so a chain + /// that loops back stops at the first repeat. Appends every refreshed + /// name to `refreshed`. + fn apply_refresh( + &mut self, + database_id: u64, + name: &str, + result: RefreshResult, + now_ms: i64, + refreshed: &mut Vec, + ) { + let mut pending = vec![(name.to_string(), result)]; + while let Some((name, result)) = pending.pop() { + if refreshed.contains(&name) { + continue; + } + let key = (database_id, name); + let Some(def) = self.definitions.get(&key) else { + continue; + }; + let RefreshResult { + rows_processed, + max_ts, + o3_min_ts, + delta, + } = result; + + if let Some(downstream) = self.dependencies.get(&key) { + for ds_name in downstream { + let Some(ds_def) = self.definitions.get(&(database_id, ds_name.clone())) else { + continue; + }; + if ds_def.stale || !takes_upstream_refresh(&ds_def.refresh_policy) { + continue; + } + let rolled = RefreshResult { + rows_processed, + max_ts, + o3_min_ts, + delta: rollup::rollup_delta(def, ds_def, &delta), + }; + pending.push((ds_name.clone(), rolled)); + } + } - if let Some(wm) = self.watermarks.get_mut(&key) { - wm.advance(result.max_ts, result.rows_processed, now_ms); - if let Some(o3_ts) = result.o3_min_ts { + refresh::merge_delta(self.materialized.entry(key.clone()).or_default(), delta); + let wm = self.watermarks.entry(key.clone()).or_default(); + wm.advance(max_ts, rows_processed, now_ms); + if let Some(o3_ts) = o3_min_ts { wm.record_o3(o3_ts); } + refreshed.push(key.1); } } @@ -374,6 +425,7 @@ mod tests { AggFunction, AggregateExpr, RefreshPolicy, }; use crate::engine::timeseries::time_bucket; + use nodedb_types::Value; use nodedb_types::timeseries::MetricSample; fn test_memtable_config() -> ColumnarMemtableConfig { @@ -630,4 +682,218 @@ mod tests { assert!(!results.is_empty()); assert!(results.len() <= 11); } + + // ── Exactness: materialized buckets equal the ad-hoc aggregate ── + + const ABOVE: i64 = 9_007_199_254_740_993; + const AT: i64 = 9_007_199_254_740_992; + const T0: i64 = 1_700_000_040_000; + + /// `(timestamp_ms, v)` rows: integers one apart above 2^53, nanosecond + /// timestamps one tick apart, and `i64::MAX` twice so a bucket total + /// leaves `i64`. + fn exact_rows() -> Vec<(i64, i64)> { + vec![ + (T0, ABOVE), + (T0 + 1_000, AT), + (T0 + 2_000, 1_700_000_000_000_000_002), + (T0 + 61_000, 1_700_000_000_000_000_001), + (T0 + 62_000, i64::MAX), + (T0 + 63_000, i64::MAX), + (T0 + 3_600_000, 1_700_000_000_000_000_003), + ] + } + + fn exact_def(name: &str, source: &str, bucket: &str) -> ContinuousAggregateDef { + let expr = |function, column: &str| AggregateExpr { + function, + source_column: column.into(), + output_column: String::new(), + }; + ContinuousAggregateDef { + aggregates: vec![ + expr(AggFunction::Count, "*"), + expr(AggFunction::Sum, "v"), + expr(AggFunction::Min, "v"), + expr(AggFunction::Max, "v"), + expr(AggFunction::Avg, "v"), + expr(AggFunction::First, "v"), + expr(AggFunction::Last, "v"), + ], + ..make_agg_def(name, source, bucket) + } + } + + /// One drain of `rows` over a `(timestamp, v BIGINT)` schema. + fn int_drain(rows: &[(i64, i64)]) -> ColumnarDrainResult { + let schema = ColumnarSchema { + columns: vec![ + ("timestamp".into(), ColumnType::Timestamp(TimeKind::Millis)), + ("v".into(), ColumnType::Int64), + ], + timestamp_idx: 0, + codecs: vec![nodedb_codec::ColumnCodec::Auto; 2], + }; + let mut mt = ColumnarMemtable::new(schema, test_memtable_config()); + for &(ts, v) in rows { + mt.ingest_row(1, &[ColumnValue::Timestamp(ts), ColumnValue::Int64(v)]) + .unwrap(); + } + mt.drain() + } + + /// The ad-hoc timeseries aggregate of `rows` per bucket of + /// `bucket_ms`, in `exact_def` order: COUNT, SUM, MIN, MAX, AVG, FIRST, + /// LAST. + fn ad_hoc(rows: &[(i64, i64)], bucket_ms: i64) -> Vec<(i64, Vec)> { + use crate::engine::timeseries::columnar_agg::AggAccum; + let mut sorted = rows.to_vec(); + sorted.sort_by_key(|&(ts, _)| ts); + let mut buckets: std::collections::BTreeMap = Default::default(); + for (ts, v) in sorted { + buckets + .entry(time_bucket::time_bucket(bucket_ms, ts)) + .or_default() + .feed_int(v); + } + buckets + .into_iter() + .map(|(bucket, a)| { + let cell = |v: Option<&Value>| v.cloned().unwrap_or(Value::Null); + let values = vec![ + Value::Integer(a.count as i64), + a.sum_value().unwrap(), + cell(a.min()), + cell(a.max()), + a.avg_f64().unwrap().map_or(Value::Null, Value::Float), + cell(a.first()), + cell(a.last()), + ]; + (bucket, values) + }) + .collect() + } + + /// Every materialized bucket of `name`, finalized. + fn materialized(mgr: &ContinuousAggregateManager, name: &str) -> Vec<(i64, Vec)> { + let def = mgr.get_definition(0, name).unwrap(); + let layout = crate::engine::timeseries::continuous_agg::ColumnLayout::of(def); + mgr.get_materialized(0, name) + .unwrap() + .into_iter() + .map(|p| { + let values = def + .aggregates + .iter() + .map(|e| p.finalize(e, &layout).unwrap()) + .collect(); + (p.bucket_ts, values) + }) + .collect() + } + + #[test] + fn materialized_integers_equal_ad_hoc_across_refreshes_and_o3() { + let mut mgr = ContinuousAggregateManager::new(); + mgr.register(exact_def("v_1m", "metrics", "1m")); + let rows = exact_rows(); + + // Three refreshes; the last carries rows below the watermark. + mgr.on_flush(0, "metrics", &int_drain(&rows[3..5]), T0); + mgr.on_flush(0, "metrics", &int_drain(&rows[5..]), T0); + mgr.on_flush(0, "metrics", &int_drain(&rows[..3]), T0); + + assert!( + mgr.get_watermark(0, "v_1m") + .unwrap() + .o3_watermark_ts + .is_some() + ); + let got = materialized(&mgr, "v_1m"); + assert_eq!(got, ad_hoc(&rows, 60_000)); + // The second bucket's total left `i64`: an exact Decimal. + assert!( + matches!(got[1].1[1], Value::Decimal(_)), + "{:?}", + got[1].1[1] + ); + } + + #[test] + fn rollup_tier_equals_ad_hoc_over_raw_rows() { + let mut mgr = ContinuousAggregateManager::new(); + mgr.register(exact_def("v_1m", "metrics", "1m")); + let mut tier2 = exact_def("v_1h", "v_1m", "1h"); + tier2.refresh_policy = RefreshPolicy::OnSeal; + mgr.register(tier2); + let rows = exact_rows(); + + let refreshed = mgr.on_flush(0, "metrics", &int_drain(&rows[4..]), T0); + assert_eq!(refreshed, vec!["v_1m", "v_1h"]); + mgr.on_flush(0, "metrics", &int_drain(&rows[..4]), T0); + + assert_eq!(materialized(&mgr, "v_1m"), ad_hoc(&rows, 60_000)); + assert_eq!(materialized(&mgr, "v_1h"), ad_hoc(&rows, 3_600_000)); + assert_eq!(mgr.get_watermark(0, "v_1h").unwrap().rows_aggregated, 7); + } + + #[test] + fn manual_refresh_rolls_up_into_downstream() { + let mut mgr = ContinuousAggregateManager::new(); + let mut tier1 = exact_def("v_1m", "metrics", "1m"); + tier1.refresh_policy = RefreshPolicy::Manual; + mgr.register(tier1); + mgr.register(exact_def("v_1h", "v_1m", "1h")); + let rows = exact_rows(); + + mgr.manual_refresh(0, "v_1m", &int_drain(&rows), T0); + assert_eq!(materialized(&mgr, "v_1h"), ad_hoc(&rows, 3_600_000)); + } + + #[test] + fn three_tier_chain_equals_ad_hoc() { + let mut mgr = ContinuousAggregateManager::new(); + mgr.register(exact_def("a", "metrics", "1m")); + mgr.register(exact_def("b", "a", "1m")); + mgr.register(exact_def("c", "b", "1h")); + let rows = exact_rows(); + + let refreshed = mgr.on_flush(0, "metrics", &int_drain(&rows), T0); + assert_eq!(refreshed, vec!["a", "b", "c"]); + assert_eq!(materialized(&mgr, "b"), ad_hoc(&rows, 60_000)); + assert_eq!(materialized(&mgr, "c"), ad_hoc(&rows, 3_600_000)); + } + + /// Two aggregates sourced from each other: a refresh of one reaches the + /// other and stops there instead of cycling. + #[test] + fn chain_that_loops_back_refreshes_each_aggregate_once() { + let mut mgr = ContinuousAggregateManager::new(); + let mut x = exact_def("x", "y", "1m"); + x.refresh_policy = RefreshPolicy::OnSeal; + mgr.register(x); + mgr.register(exact_def("y", "x", "1m")); + let rows = exact_rows(); + + mgr.manual_refresh(0, "x", &int_drain(&rows), T0); + assert_eq!(materialized(&mgr, "x"), ad_hoc(&rows, 60_000)); + assert_eq!(materialized(&mgr, "y"), ad_hoc(&rows, 60_000)); + } + + #[test] + fn reregister_with_another_shape_drops_old_buckets() { + let mut mgr = ContinuousAggregateManager::new(); + mgr.register(exact_def("v_1m", "metrics", "1m")); + mgr.on_flush(0, "metrics", &int_drain(&exact_rows()), T0); + assert!(!mgr.get_materialized(0, "v_1m").unwrap().is_empty()); + + // Same shape: buckets stay. + mgr.register(exact_def("v_1m", "metrics", "1m")); + assert!(!mgr.get_materialized(0, "v_1m").unwrap().is_empty()); + + // Another bucket interval: buckets and watermark reset. + mgr.register(exact_def("v_1m", "metrics", "5m")); + assert!(mgr.get_materialized(0, "v_1m").unwrap().is_empty()); + assert_eq!(mgr.get_watermark(0, "v_1m").unwrap().rows_aggregated, 0); + } } diff --git a/nodedb/src/engine/timeseries/continuous_agg/mod.rs b/nodedb/src/engine/timeseries/continuous_agg/mod.rs index a4e4e633a..9b85dc92d 100644 --- a/nodedb/src/engine/timeseries/continuous_agg/mod.rs +++ b/nodedb/src/engine/timeseries/continuous_agg/mod.rs @@ -4,9 +4,10 @@ pub mod definition; pub mod manager; pub mod partial; pub mod refresh; +pub mod rollup; pub mod watermark; pub use definition::{AggFunction, AggregateExpr, ContinuousAggregateDef, RefreshPolicy}; pub use manager::{AggregateInfo, ContinuousAggregateManager}; -pub use partial::PartialAggregate; +pub use partial::{ColumnLayout, PartialAggregate}; pub use watermark::WatermarkState; diff --git a/nodedb/src/engine/timeseries/continuous_agg/partial.rs b/nodedb/src/engine/timeseries/continuous_agg/partial.rs deleted file mode 100644 index eae85e90b..000000000 --- a/nodedb/src/engine/timeseries/continuous_agg/partial.rs +++ /dev/null @@ -1,233 +0,0 @@ -// SPDX-License-Identifier: BUSL-1.1 - -//! Partial aggregate state for incremental merging. -//! -//! Stores enough state per (bucket, group_key) to merge incrementally: -//! count, sum, min, max, first/last timestamps and values, plus optional -//! sketch state for approximate aggregations. - -use super::definition::AggFunction; -use nodedb_types::approx::{HyperLogLog, SpaceSaving, TDigest}; - -/// Partial aggregate state for a single (bucket, group_key) combination. -pub struct PartialAggregate { - pub bucket_ts: i64, - /// Symbol IDs for GROUP BY columns. - pub group_key: Vec, - pub count: u64, - pub sum: f64, - pub min: f64, - pub max: f64, - pub first_ts: i64, - pub first_val: f64, - pub last_ts: i64, - pub last_val: f64, - - // ── Sketch state (lazily initialized only when needed) ── - pub hll: Option, - pub tdigest: Option, - pub topk: Option, -} - -impl PartialAggregate { - /// Create from a single sample. - pub fn new(bucket_ts: i64, group_key: Vec, ts: i64, val: f64) -> Self { - Self { - bucket_ts, - group_key, - count: 1, - sum: val, - min: val, - max: val, - first_ts: ts, - first_val: val, - last_ts: ts, - last_val: val, - hll: None, - tdigest: None, - topk: None, - } - } - - /// Ensure sketch state is initialized for the given function. - pub fn ensure_sketch(&mut self, function: &AggFunction) { - match function { - AggFunction::CountDistinct if self.hll.is_none() => { - self.hll = Some(HyperLogLog::new()); - } - AggFunction::Percentile(_) if self.tdigest.is_none() => { - self.tdigest = Some(TDigest::new()); - } - AggFunction::TopK(k) if self.topk.is_none() => { - self.topk = Some(SpaceSaving::new(*k)); - } - _ => {} - } - } - - /// Merge another sample into this partial aggregate. - pub fn merge_sample(&mut self, ts: i64, val: f64) { - self.count += 1; - self.sum += val; - if val < self.min { - self.min = val; - } - if val > self.max { - self.max = val; - } - if ts < self.first_ts { - self.first_ts = ts; - self.first_val = val; - } - if ts > self.last_ts { - self.last_ts = ts; - self.last_val = val; - } - - // Feed into active sketches. - if let Some(hll) = &mut self.hll { - hll.add(val.to_bits()); - } - if let Some(td) = &mut self.tdigest { - td.add(val); - } - if let Some(ss) = &mut self.topk { - ss.add(val.to_bits()); - } - } - - /// Merge another partial aggregate (for cross-shard or incremental merge). - pub fn merge_partial(&mut self, other: &PartialAggregate) { - self.count += other.count; - self.sum += other.sum; - if other.min < self.min { - self.min = other.min; - } - if other.max > self.max { - self.max = other.max; - } - if other.first_ts < self.first_ts { - self.first_ts = other.first_ts; - self.first_val = other.first_val; - } - if other.last_ts > self.last_ts { - self.last_ts = other.last_ts; - self.last_val = other.last_val; - } - - // Merge sketch state. - if let Some(other_hll) = &other.hll { - self.hll - .get_or_insert_with(HyperLogLog::new) - .merge(other_hll); - } - if let Some(other_td) = &other.tdigest { - self.tdigest - .get_or_insert_with(TDigest::new) - .merge(other_td); - } - if let Some(other_ss) = &other.topk { - let k = other_ss.top_k().len().max(10); - self.topk - .get_or_insert_with(|| SpaceSaving::new(k)) - .merge(other_ss); - } - } - - /// Compute a final aggregate value from the partial state. - pub fn finalize(&self, function: &AggFunction) -> f64 { - match function { - AggFunction::Sum => self.sum, - AggFunction::Count => self.count as f64, - AggFunction::Min => self.min, - AggFunction::Max => self.max, - AggFunction::Avg => { - if self.count == 0 { - 0.0 - } else { - self.sum / self.count as f64 - } - } - AggFunction::First => self.first_val, - AggFunction::Last => self.last_val, - AggFunction::CountDistinct => self.hll.as_ref().map_or(0.0, |h| h.estimate()), - AggFunction::Percentile(q) => { - self.tdigest.as_ref().map_or(f64::NAN, |td| td.quantile(*q)) - } - AggFunction::TopK(_) => { - // TopK returns structured data; finalize as count of tracked items. - self.topk.as_ref().map_or(0.0, |ss| ss.top_k().len() as f64) - } - _ => f64::NAN, - } - } -} - -#[cfg(test)] -mod tests { - use super::*; - - #[test] - fn single_sample() { - let pa = PartialAggregate::new(0, vec![], 100, 42.0); - assert_eq!(pa.count, 1); - assert_eq!(pa.finalize(&AggFunction::Sum), 42.0); - assert_eq!(pa.finalize(&AggFunction::Avg), 42.0); - } - - #[test] - fn merge_samples() { - let mut pa = PartialAggregate::new(0, vec![], 100, 10.0); - pa.merge_sample(200, 20.0); - pa.merge_sample(300, 30.0); - - assert_eq!(pa.finalize(&AggFunction::Count), 3.0); - assert_eq!(pa.finalize(&AggFunction::Sum), 60.0); - assert_eq!(pa.finalize(&AggFunction::Min), 10.0); - assert_eq!(pa.finalize(&AggFunction::Max), 30.0); - assert!((pa.finalize(&AggFunction::Avg) - 20.0).abs() < f64::EPSILON); - assert_eq!(pa.finalize(&AggFunction::First), 10.0); - assert_eq!(pa.finalize(&AggFunction::Last), 30.0); - } - - #[test] - fn sketch_count_distinct() { - let mut pa = PartialAggregate::new(0, vec![], 100, 1.0); - pa.ensure_sketch(&AggFunction::CountDistinct); - for i in 1..100 { - pa.merge_sample(100 + i, i as f64); - } - let est = pa.finalize(&AggFunction::CountDistinct); - assert!(est > 80.0 && est < 120.0, "expected ~100, got {est}"); - } - - #[test] - fn sketch_percentile() { - let mut pa = PartialAggregate::new(0, vec![], 0, 0.0); - pa.ensure_sketch(&AggFunction::Percentile(0.5)); - for i in 1..1000 { - pa.merge_sample(i, i as f64); - } - let p50 = pa.finalize(&AggFunction::Percentile(0.5)); - assert!(p50 > 400.0 && p50 < 600.0, "expected ~500, got {p50}"); - } - - #[test] - fn merge_partials() { - let mut a = PartialAggregate::new(0, vec![], 100, 10.0); - a.merge_sample(200, 20.0); - - let mut b = PartialAggregate::new(0, vec![], 50, 5.0); - b.merge_sample(300, 30.0); - - a.merge_partial(&b); - assert_eq!(a.count, 4); - assert_eq!(a.sum, 65.0); - assert_eq!(a.min, 5.0); - assert_eq!(a.max, 30.0); - assert_eq!(a.first_ts, 50); - assert_eq!(a.first_val, 5.0); - assert_eq!(a.last_ts, 300); - assert_eq!(a.last_val, 30.0); - } -} diff --git a/nodedb/src/engine/timeseries/continuous_agg/partial/bucket.rs b/nodedb/src/engine/timeseries/continuous_agg/partial/bucket.rs new file mode 100644 index 000000000..37aa3e5f7 --- /dev/null +++ b/nodedb/src/engine/timeseries/continuous_agg/partial/bucket.rs @@ -0,0 +1,211 @@ +// SPDX-License-Identifier: BUSL-1.1 + +//! Partial aggregate state of one `(bucket, group_key)`. +//! +//! The bucket keeps its row count for `COUNT` and one [`ColumnPartial`] per +//! source column in the aggregate's [`ColumnLayout`]. Finalizing an +//! expression gives the value the ad-hoc aggregate gives over the same rows: +//! `COUNT` is an `Integer`, SUM is exact (`Integer`, `Decimal`, or `Float`), +//! AVG is a `Float` from the exact total, and MIN / MAX / FIRST / LAST are +//! the original cells. + +use nodedb_query::EvalError; +use nodedb_types::Value; + +use super::super::definition::{AggFunction, AggregateExpr}; +use super::column::ColumnPartial; +use super::layout::ColumnLayout; + +/// Partial aggregate state for a single `(bucket, group_key)` combination. +#[derive(Debug)] +pub struct PartialAggregate { + pub bucket_ts: i64, + /// Symbol IDs for GROUP BY columns. + pub group_key: Vec, + /// Rows aggregated into this bucket. + pub count: u64, + /// One state per column slot of the aggregate's layout. + pub columns: Vec, +} + +impl PartialAggregate { + /// An empty bucket shaped by `layout`. + pub fn new(bucket_ts: i64, group_key: Vec, layout: &ColumnLayout) -> Self { + Self { + bucket_ts, + group_key, + count: 0, + columns: (0..layout.len()) + .map(|slot| ColumnPartial::with_sketches(layout.sketches(slot))) + .collect(), + } + } + + /// Fold `other`, a bucket of the same layout, into this one. + pub fn merge(&mut self, other: &PartialAggregate) { + self.count += other.count; + for (mine, theirs) in self.columns.iter_mut().zip(&other.columns) { + mine.merge(theirs); + } + } + + /// Fold `other`, a bucket of another layout, into this one. `map[i]` is + /// the slot of `other` that feeds this bucket's slot `i`. A slot mapped to + /// `None` takes nothing. + pub fn merge_mapped(&mut self, other: &PartialAggregate, map: &[Option]) { + self.count += other.count; + for (mine, source) in self.columns.iter_mut().zip(map) { + if let Some(theirs) = source.and_then(|slot| other.columns.get(slot)) { + mine.merge(theirs); + } + } + } + + /// The final value of `expr` over this bucket. `layout` is the layout + /// this bucket was built with. A column the bucket has no numeric cell + /// for finalizes to NULL, except `COUNT`, which counts rows. + pub fn finalize( + &self, + expr: &AggregateExpr, + layout: &ColumnLayout, + ) -> Result { + if expr.function == AggFunction::Count { + return i64::try_from(self.count) + .map(Value::Integer) + .map_err(|_| EvalError::NumericOverflow { function: "count" }); + } + let Some(column) = layout + .slot(&expr.source_column) + .and_then(|slot| self.columns.get(slot)) + else { + return Ok(Value::Null); + }; + let cell = |v: Option<&Value>| v.cloned().unwrap_or(Value::Null); + Ok(match &expr.function { + AggFunction::Sum => column.sum.sum()?, + AggFunction::Avg => column.sum.avg()?, + AggFunction::Min => cell(column.min.as_ref()), + AggFunction::Max => cell(column.max.as_ref()), + AggFunction::First => cell(column.first.as_ref().map(|f| &f.value)), + AggFunction::Last => cell(column.last.as_ref().map(|l| &l.value)), + AggFunction::CountDistinct => column + .hll + .as_ref() + .filter(|_| column.sum.count() > 0) + .map_or(Value::Null, |h| Value::Float(h.estimate())), + AggFunction::Percentile(q) => column + .tdigest + .as_ref() + .filter(|_| column.sum.count() > 0) + .map_or(Value::Null, |td| Value::Float(td.quantile(*q))), + AggFunction::TopK(_) => match &column.topk { + Some(ss) => i64::try_from(ss.top_k().len()) + .map(Value::Integer) + .map_err(|_| EvalError::NumericOverflow { function: "topk" })?, + None => Value::Null, + }, + // `Count` returned above. `AggFunction` is non-exhaustive: a + // function this bucket keeps no state for has no value. + _ => Value::Null, + }) + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::engine::timeseries::continuous_agg::definition::{ + ContinuousAggregateDef, RefreshPolicy, + }; + + fn expr(function: AggFunction, column: &str) -> AggregateExpr { + AggregateExpr { + function, + source_column: column.into(), + output_column: String::new(), + } + } + + fn layout(aggregates: Vec) -> ColumnLayout { + ColumnLayout::of(&ContinuousAggregateDef { + database_id: 0, + name: "a".into(), + source: "s".into(), + bucket_interval: "1m".into(), + bucket_interval_ms: 60_000, + group_by: Vec::new(), + aggregates, + refresh_policy: RefreshPolicy::OnFlush, + retention_period_ms: 0, + stale: false, + }) + } + + #[test] + fn finalize_keeps_integer_types() { + let exprs = vec![ + expr(AggFunction::Count, "*"), + expr(AggFunction::Sum, "v"), + expr(AggFunction::Min, "v"), + expr(AggFunction::Max, "v"), + expr(AggFunction::Avg, "v"), + ]; + let layout = layout(exprs.clone()); + let mut p = PartialAggregate::new(0, Vec::new(), &layout); + for (ts, v) in [i64::MAX, i64::MAX].iter().enumerate() { + p.count += 1; + p.columns[0].add_int(ts as i64, *v); + } + let out: Vec = exprs + .iter() + .map(|e| p.finalize(e, &layout).unwrap()) + .collect(); + assert_eq!(out[0], Value::Integer(2)); + assert_eq!( + out[1], + Value::Decimal(rust_decimal::Decimal::from_i128_with_scale( + 2 * i128::from(i64::MAX), + 0 + )) + ); + assert_eq!(out[2], Value::Integer(i64::MAX)); + assert_eq!(out[3], Value::Integer(i64::MAX)); + assert_eq!(out[4], Value::Float(i64::MAX as f64)); + } + + #[test] + fn column_without_cells_is_null() { + let exprs = vec![expr(AggFunction::Sum, "v"), expr(AggFunction::Min, "v")]; + let layout = layout(exprs.clone()); + let mut p = PartialAggregate::new(0, Vec::new(), &layout); + p.count = 3; + for e in &exprs { + assert_eq!(p.finalize(e, &layout).unwrap(), Value::Null); + } + assert_eq!( + p.finalize(&expr(AggFunction::Count, "v"), &layout).unwrap(), + Value::Integer(3) + ); + } + + #[test] + fn merge_mapped_takes_matching_columns() { + let up_layout = layout(vec![ + expr(AggFunction::Sum, "a"), + expr(AggFunction::Sum, "b"), + ]); + let down_layout = layout(vec![expr(AggFunction::Sum, "b")]); + let mut up = PartialAggregate::new(0, Vec::new(), &up_layout); + up.count = 1; + up.columns[0].add_int(0, 5); + up.columns[1].add_int(0, 7); + let mut down = PartialAggregate::new(0, Vec::new(), &down_layout); + down.merge_mapped(&up, &down_layout.map_from(&up_layout)); + assert_eq!(down.count, 1); + assert_eq!( + down.finalize(&expr(AggFunction::Sum, "b"), &down_layout) + .unwrap(), + Value::Integer(7) + ); + } +} diff --git a/nodedb/src/engine/timeseries/continuous_agg/partial/column.rs b/nodedb/src/engine/timeseries/continuous_agg/partial/column.rs new file mode 100644 index 000000000..a2be0aea3 --- /dev/null +++ b/nodedb/src/engine/timeseries/continuous_agg/partial/column.rs @@ -0,0 +1,213 @@ +// SPDX-License-Identifier: BUSL-1.1 + +//! Partial state of one source column inside one aggregate bucket. +//! +//! The rule matches the ad-hoc aggregate path: +//! +//! - SUM / AVG total in an [`ExactSum`]: an integer cell adds exactly, a +//! float cell adds Kahan-compensated. +//! - MIN / MAX keep the original cell, compared exactly by +//! [`value_replaces`]. +//! - FIRST / LAST keep the original cell at the lowest / highest timestamp. +//! - The sketches (HLL, t-digest, top-K) are approximate by definition. +//! +//! Every part merges without loss, so a bucket built from several refreshes, +//! out-of-order flushes, or a lower-tier rollup equals one pass over the rows. + +use nodedb_query::ExactSum; +use nodedb_query::window::extremum::value_replaces; +use nodedb_types::Value; +use nodedb_types::approx::{HyperLogLog, SpaceSaving, TDigest}; + +use super::super::definition::AggFunction; +use crate::util::fnv1a_hash; + +/// A cell at its timestamp, for FIRST / LAST. +#[derive(Debug, Clone, PartialEq)] +pub struct TimedCell { + pub ts: i64, + pub value: Value, +} + +/// Partial aggregate state of one source column. +#[derive(Debug, Default)] +pub struct ColumnPartial { + /// Exact SUM / AVG state. Its count is the number of cells added. + pub sum: ExactSum, + /// Smallest cell, as added. + pub min: Option, + /// Largest cell, as added. + pub max: Option, + /// Cell at the lowest timestamp. An equal timestamp keeps the earlier. + pub first: Option, + /// Cell at the highest timestamp. An equal timestamp keeps the earlier. + pub last: Option, + pub hll: Option, + pub tdigest: Option, + pub topk: Option, +} + +impl ColumnPartial { + /// An empty column state feeding the sketches `functions` name. + pub fn with_sketches(functions: &[AggFunction]) -> Self { + let mut column = Self::default(); + for function in functions { + match function { + AggFunction::CountDistinct if column.hll.is_none() => { + column.hll = Some(HyperLogLog::new()); + } + AggFunction::Percentile(_) if column.tdigest.is_none() => { + column.tdigest = Some(TDigest::new()); + } + AggFunction::TopK(k) if column.topk.is_none() => { + column.topk = Some(SpaceSaving::new(*k)); + } + _ => {} + } + } + column + } + + /// Add one integer cell at timestamp `ts`, kept exact. + pub fn add_int(&mut self, ts: i64, v: i64) { + self.sum.add_i64(v); + // Distinct integers hash apart even above 2^53, where their `f64` + // readings collide. + self.add_cell( + ts, + Value::Integer(v), + fnv1a_hash(&v.to_le_bytes()), + v as f64, + ); + } + + /// Add one float cell at timestamp `ts`. + pub fn add_float(&mut self, ts: i64, v: f64) { + self.sum.add_f64(v); + self.add_cell( + ts, + Value::Float(v), + fnv1a_hash(&v.to_bits().to_le_bytes()), + v, + ); + } + + /// Feed the sketches and the extremes with a cell already added to the + /// sum. `hash` is its sketch identity, `reading` its `f64` reading. + fn add_cell(&mut self, ts: i64, cell: Value, hash: u64, reading: f64) { + if let Some(hll) = &mut self.hll { + hll.add(hash); + } + if let Some(td) = &mut self.tdigest { + td.add(reading); + } + if let Some(ss) = &mut self.topk { + ss.add(hash); + } + if value_replaces(&cell, self.min.as_ref(), false) { + self.min = Some(cell.clone()); + } + if value_replaces(&cell, self.max.as_ref(), true) { + self.max = Some(cell.clone()); + } + if self.first.as_ref().is_none_or(|f| ts < f.ts) { + self.first = Some(TimedCell { + ts, + value: cell.clone(), + }); + } + if self.last.as_ref().is_none_or(|l| ts > l.ts) { + self.last = Some(TimedCell { ts, value: cell }); + } + } + + /// Fold `other` into this state without loss. + pub fn merge(&mut self, other: &ColumnPartial) { + self.sum.merge(&other.sum); + if let Some(min) = &other.min + && value_replaces(min, self.min.as_ref(), false) + { + self.min = Some(min.clone()); + } + if let Some(max) = &other.max + && value_replaces(max, self.max.as_ref(), true) + { + self.max = Some(max.clone()); + } + if let Some(first) = &other.first + && self.first.as_ref().is_none_or(|f| first.ts < f.ts) + { + self.first = Some(first.clone()); + } + if let Some(last) = &other.last + && self.last.as_ref().is_none_or(|l| last.ts > l.ts) + { + self.last = Some(last.clone()); + } + if let (Some(mine), Some(theirs)) = (&mut self.hll, &other.hll) { + mine.merge(theirs); + } + if let (Some(mine), Some(theirs)) = (&mut self.tdigest, &other.tdigest) { + mine.merge(theirs); + } + if let (Some(mine), Some(theirs)) = (&mut self.topk, &other.topk) { + mine.merge(theirs); + } + } +} + +#[cfg(test)] +mod tests { + use super::*; + + const ABOVE: i64 = 9_007_199_254_740_993; + const AT: i64 = 9_007_199_254_740_992; + + #[test] + fn integers_past_two_pow_53_stay_exact() { + let mut c = ColumnPartial::default(); + c.add_int(2, ABOVE); + c.add_int(1, AT); + assert_eq!(c.sum.sum().unwrap(), Value::Integer(ABOVE + AT)); + assert_eq!(c.min, Some(Value::Integer(AT))); + assert_eq!(c.max, Some(Value::Integer(ABOVE))); + assert_eq!( + c.first.as_ref().map(|f| &f.value), + Some(&Value::Integer(AT)) + ); + assert_eq!( + c.last.as_ref().map(|l| &l.value), + Some(&Value::Integer(ABOVE)) + ); + } + + #[test] + fn merge_equals_single_pass() { + let cells = [ABOVE, AT, i64::MAX, 1_700_000_000_000_000_001]; + let mut single = ColumnPartial::default(); + for (ts, v) in cells.iter().enumerate() { + single.add_int(ts as i64, *v); + } + let mut a = ColumnPartial::default(); + let mut b = ColumnPartial::default(); + for (ts, v) in cells.iter().enumerate() { + let part = if ts % 2 == 0 { &mut a } else { &mut b }; + part.add_int(ts as i64, *v); + } + b.merge(&a); + assert_eq!(b.sum, single.sum); + assert_eq!(b.min, single.min); + assert_eq!(b.max, single.max); + assert_eq!(b.first, single.first); + assert_eq!(b.last, single.last); + } + + #[test] + fn distinct_integers_above_two_pow_53_count_apart() { + let mut c = ColumnPartial::with_sketches(&[AggFunction::CountDistinct]); + c.add_int(0, ABOVE); + c.add_int(1, AT); + let estimate = c.hll.as_ref().map_or(0.0, HyperLogLog::estimate); + assert!(estimate > 1.5, "two distinct values, estimate {estimate}"); + } +} diff --git a/nodedb/src/engine/timeseries/continuous_agg/partial/layout.rs b/nodedb/src/engine/timeseries/continuous_agg/partial/layout.rs new file mode 100644 index 000000000..ac6249110 --- /dev/null +++ b/nodedb/src/engine/timeseries/continuous_agg/partial/layout.rs @@ -0,0 +1,133 @@ +// SPDX-License-Identifier: BUSL-1.1 + +//! Column layout of an aggregate's partial state. +//! +//! A partial bucket keeps one [`super::ColumnPartial`] per distinct source +//! column the aggregate reads. `COUNT(*)` reads no column and uses the +//! bucket's row count. Every aggregate expression on one column shares that +//! column's state, so SUM, MIN, AVG and the rest of one column all derive +//! from the same inputs. + +use super::super::definition::{AggFunction, ContinuousAggregateDef}; + +/// The source columns of an aggregate, in first-use order, and the sketch +/// functions each column feeds. +#[derive(Debug, Clone, PartialEq)] +pub struct ColumnLayout { + columns: Vec, + sketches: Vec>, +} + +impl ColumnLayout { + /// The layout of `def`. + pub fn of(def: &ContinuousAggregateDef) -> Self { + let mut layout = Self { + columns: Vec::new(), + sketches: Vec::new(), + }; + for expr in &def.aggregates { + if expr.source_column == "*" { + continue; + } + let slot = match layout.slot(&expr.source_column) { + Some(slot) => slot, + None => { + layout.columns.push(expr.source_column.clone()); + layout.sketches.push(Vec::new()); + layout.columns.len() - 1 + } + }; + if expr.function.uses_sketch() { + layout.sketches[slot].push(expr.function.clone()); + } + } + layout + } + + /// The slot of `column`, or `None` when the aggregate does not read it. + pub fn slot(&self, column: &str) -> Option { + self.columns.iter().position(|c| c == column) + } + + /// The source column names, in slot order. + pub fn columns(&self) -> &[String] { + &self.columns + } + + /// The sketch functions fed by the column in `slot`. + pub fn sketches(&self, slot: usize) -> &[AggFunction] { + self.sketches.get(slot).map_or(&[], Vec::as_slice) + } + + /// Number of column slots. + pub fn len(&self) -> usize { + self.columns.len() + } + + /// Whether the aggregate reads no column. + pub fn is_empty(&self) -> bool { + self.columns.is_empty() + } + + /// For each slot of `self`, the slot of the same column in `upstream`. + /// A column `upstream` does not read maps to `None`. + pub fn map_from(&self, upstream: &ColumnLayout) -> Vec> { + self.columns.iter().map(|c| upstream.slot(c)).collect() + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::engine::timeseries::continuous_agg::definition::{AggregateExpr, RefreshPolicy}; + + fn expr(function: AggFunction, column: &str) -> AggregateExpr { + AggregateExpr { + function, + source_column: column.into(), + output_column: String::new(), + } + } + + fn def(aggregates: Vec) -> ContinuousAggregateDef { + ContinuousAggregateDef { + database_id: 0, + name: "a".into(), + source: "s".into(), + bucket_interval: "1m".into(), + bucket_interval_ms: 60_000, + group_by: Vec::new(), + aggregates, + refresh_policy: RefreshPolicy::OnFlush, + retention_period_ms: 0, + stale: false, + } + } + + #[test] + fn one_slot_per_distinct_column() { + let layout = ColumnLayout::of(&def(vec![ + expr(AggFunction::Count, "*"), + expr(AggFunction::Sum, "v"), + expr(AggFunction::Max, "w"), + expr(AggFunction::CountDistinct, "v"), + ])); + assert_eq!(layout.columns(), ["v".to_string(), "w".to_string()]); + assert_eq!(layout.sketches(0), [AggFunction::CountDistinct]); + assert!(layout.sketches(1).is_empty()); + assert_eq!(layout.slot("*"), None); + } + + #[test] + fn map_from_matches_columns_by_name() { + let up = ColumnLayout::of(&def(vec![ + expr(AggFunction::Sum, "a"), + expr(AggFunction::Sum, "b"), + ])); + let down = ColumnLayout::of(&def(vec![ + expr(AggFunction::Sum, "b"), + expr(AggFunction::Sum, "c"), + ])); + assert_eq!(down.map_from(&up), vec![Some(1), None]); + } +} diff --git a/nodedb/src/engine/timeseries/continuous_agg/partial/mod.rs b/nodedb/src/engine/timeseries/continuous_agg/partial/mod.rs new file mode 100644 index 000000000..2b8b68fbd --- /dev/null +++ b/nodedb/src/engine/timeseries/continuous_agg/partial/mod.rs @@ -0,0 +1,9 @@ +// SPDX-License-Identifier: BUSL-1.1 + +pub mod bucket; +pub mod column; +pub mod layout; + +pub use bucket::PartialAggregate; +pub use column::{ColumnPartial, TimedCell}; +pub use layout::ColumnLayout; diff --git a/nodedb/src/engine/timeseries/continuous_agg/refresh.rs b/nodedb/src/engine/timeseries/continuous_agg/refresh.rs index a325df578..e64c44b03 100644 --- a/nodedb/src/engine/timeseries/continuous_agg/refresh.rs +++ b/nodedb/src/engine/timeseries/continuous_agg/refresh.rs @@ -2,17 +2,26 @@ //! Incremental refresh engine for continuous aggregates. //! -//! Takes a `ColumnarDrainResult` (flushed data) and computes aggregates -//! by time_bucket + GROUP BY columns. O(flushed_rows), not O(total_rows). +//! Takes a `ColumnarDrainResult` (flushed data) and computes the partial +//! state of every touched `(time_bucket, group)` from those rows alone: +//! O(flushed_rows), not O(total_rows). The caller merges that delta into the +//! materialized state and rolls it up into downstream tiers. +//! +//! Cells feed the partial state exactly as the ad-hoc aggregate scan feeds +//! its accumulator: an `Int64` or `Timestamp` cell as an exact integer, a +//! `Float64` cell as a float, and a symbol cell not at all. use std::collections::HashMap; use super::definition::ContinuousAggregateDef; -use super::partial::PartialAggregate; +use super::partial::{ColumnLayout, PartialAggregate}; use super::watermark::WatermarkState; use crate::engine::timeseries::columnar_memtable::{ColumnData, ColumnarDrainResult}; use crate::engine::timeseries::time_bucket; +/// Partial buckets keyed by `(bucket_ts, group_key)`. +pub type Buckets = HashMap<(i64, Vec), PartialAggregate>; + /// Result of a single aggregate refresh. pub struct RefreshResult { /// Number of input rows processed. @@ -21,16 +30,16 @@ pub struct RefreshResult { pub max_ts: i64, /// Oldest O3 timestamp (below watermark), if any. pub o3_min_ts: Option, + /// Partial state of the flushed rows alone, to merge into the + /// materialized state. + pub delta: Buckets, } -/// Incrementally refresh an aggregate from flushed data. -/// -/// Merges new samples into the existing materialized state. +/// Compute the partial state of `drain`'s rows for `def`. pub fn refresh_from_drain( def: &ContinuousAggregateDef, drain: &ColumnarDrainResult, watermark: &WatermarkState, - materialized: &mut HashMap<(i64, Vec), PartialAggregate>, ) -> RefreshResult { let bucket_ms = def.bucket_interval_ms; if bucket_ms <= 0 || drain.row_count == 0 { @@ -38,95 +47,74 @@ pub fn refresh_from_drain( rows_processed: 0, max_ts: watermark.watermark_ts, o3_min_ts: None, + delta: Buckets::new(), }; } + let layout = ColumnLayout::of(def); let ts_idx = drain.schema.timestamp_idx; let timestamps = drain.columns[ts_idx].as_timestamps(); + let column_index = |name: &str| { + drain + .schema + .columns + .iter() + .position(|(column, _)| column == name) + }; - // Find value column index for the first aggregate expression. - // All expressions reference columns from the same drain. - let value_col_idx = def - .aggregates + // Drain column of each layout slot. + let slot_columns: Vec> = layout + .columns() .iter() - .find(|e| e.source_column != "*") - .and_then(|e| { - drain - .schema - .columns - .iter() - .position(|(name, _)| name == &e.source_column) - }); - - // Find GROUP BY column indices. - let group_col_indices: Vec> = def + .map(|name| column_index(name).map(|idx| &drain.columns[idx])) + .collect(); + + // Drain column of each GROUP BY column. + let group_columns: Vec> = def .group_by .iter() - .map(|col_name| { - drain - .schema - .columns - .iter() - .position(|(name, _)| name == col_name) - }) + .map(|name| column_index(name).map(|idx| &drain.columns[idx])) .collect(); let current_watermark = watermark.watermark_ts; let mut max_ts = current_watermark; let mut o3_min: Option = None; + let mut delta = Buckets::new(); for row in 0..drain.row_count as usize { let ts = timestamps[row]; let bucket = time_bucket::time_bucket(bucket_ms, ts); // O3 detection. - if ts <= current_watermark { - match o3_min { - Some(current) if ts < current => o3_min = Some(ts), - None => o3_min = Some(ts), - _ => {} - } + if ts <= current_watermark && o3_min.is_none_or(|current| ts < current) { + o3_min = Some(ts); } if ts > max_ts { max_ts = ts; } - // Build group key. - let group_key: Vec = group_col_indices + let group_key: Vec = group_columns .iter() - .map(|opt_idx| match opt_idx { - Some(idx) => match &drain.columns[*idx] { - ColumnData::Symbol(v) => v[row], - _ => 0, - }, - None => 0, + .map(|column| match column { + Some(ColumnData::Symbol(ids)) => ids[row], + _ => 0, }) .collect(); - // Get value. - let val = value_col_idx - .map(|idx| match &drain.columns[idx] { - ColumnData::Float64(v) => v[row], - ColumnData::Int64(v) => v[row] as f64, - ColumnData::Timestamp(v) => v[row] as f64, - ColumnData::Symbol(_) => 0.0, - ColumnData::DictEncoded { .. } => 0.0, - }) - .unwrap_or(1.0); // COUNT(*) - - // Merge into materialized state. - let key = (bucket, group_key.clone()); - match materialized.get_mut(&key) { - Some(partial) => partial.merge_sample(ts, val), - None => { - let mut partial = PartialAggregate::new(bucket, group_key, ts, val); - // Initialize sketches for any approximate aggregate expressions. - for expr in &def.aggregates { - if expr.function.uses_sketch() { - partial.ensure_sketch(&expr.function); - } + let partial = delta + .entry((bucket, group_key)) + .or_insert_with_key(|(bucket, key)| { + PartialAggregate::new(*bucket, key.clone(), &layout) + }); + partial.count += 1; + for (slot, column) in slot_columns.iter().enumerate() { + let state = &mut partial.columns[slot]; + match column { + Some(ColumnData::Float64(v)) => state.add_float(ts, v[row]), + Some(ColumnData::Int64(v)) | Some(ColumnData::Timestamp(v)) => { + state.add_int(ts, v[row]) } - materialized.insert(key, partial); + Some(ColumnData::Symbol(_)) | Some(ColumnData::DictEncoded { .. }) | None => {} } } } @@ -135,5 +123,18 @@ pub fn refresh_from_drain( rows_processed: drain.row_count, max_ts, o3_min_ts: o3_min, + delta, + } +} + +/// Merge `delta` into `materialized`. Both hold buckets of one layout. +pub fn merge_delta(materialized: &mut Buckets, delta: Buckets) { + for (key, partial) in delta { + match materialized.get_mut(&key) { + Some(existing) => existing.merge(&partial), + None => { + materialized.insert(key, partial); + } + } } } diff --git a/nodedb/src/engine/timeseries/continuous_agg/rollup.rs b/nodedb/src/engine/timeseries/continuous_agg/rollup.rs new file mode 100644 index 000000000..8711fcbf8 --- /dev/null +++ b/nodedb/src/engine/timeseries/continuous_agg/rollup.rs @@ -0,0 +1,59 @@ +// SPDX-License-Identifier: BUSL-1.1 + +//! Roll an upstream aggregate's refresh delta up into a downstream tier. +//! +//! A downstream aggregate sourced from another aggregate (a retention-policy +//! tier, or `CREATE CONTINUOUS AGGREGATE ... ON `) never re-reads +//! raw rows. It merges the upstream's partial buckets into its own coarser +//! buckets. Every partial part merges without loss, so the downstream bucket +//! equals one pass over the raw rows it covers. +//! +//! An upstream bucket lands whole in the downstream bucket that holds its +//! start. A downstream interval that is a multiple of the upstream interval +//! covers whole upstream buckets, so the rollup is exact. + +use super::definition::ContinuousAggregateDef; +use super::partial::{ColumnLayout, PartialAggregate}; +use super::refresh::Buckets; +use crate::engine::timeseries::time_bucket; + +/// The downstream buckets `delta`, a refresh delta of `upstream`, adds to +/// `downstream`. +/// +/// A downstream column or GROUP BY column the upstream does not carry takes +/// no input: its column state stays empty and its key part is `0`. +pub fn rollup_delta( + upstream: &ContinuousAggregateDef, + downstream: &ContinuousAggregateDef, + delta: &Buckets, +) -> Buckets { + let mut out = Buckets::new(); + if downstream.bucket_interval_ms <= 0 { + return out; + } + let up_layout = ColumnLayout::of(upstream); + let down_layout = ColumnLayout::of(downstream); + let column_map = down_layout.map_from(&up_layout); + let key_map: Vec> = downstream + .group_by + .iter() + .map(|name| upstream.group_by.iter().position(|g| g == name)) + .collect(); + + for partial in delta.values() { + let bucket = time_bucket::time_bucket(downstream.bucket_interval_ms, partial.bucket_ts); + let key: Vec = key_map + .iter() + .map(|pos| { + pos.and_then(|p| partial.group_key.get(p).copied()) + .unwrap_or(0) + }) + .collect(); + out.entry((bucket, key)) + .or_insert_with_key(|(bucket, key)| { + PartialAggregate::new(*bucket, key.clone(), &down_layout) + }) + .merge_mapped(partial, &column_map); + } + out +} diff --git a/nodedb/src/engine/timeseries/grouped_scan/types.rs b/nodedb/src/engine/timeseries/grouped_scan/types.rs index a033ae539..42e3962de 100644 --- a/nodedb/src/engine/timeseries/grouped_scan/types.rs +++ b/nodedb/src/engine/timeseries/grouped_scan/types.rs @@ -132,14 +132,13 @@ pub(super) fn accumulate_row( accums[agg_idx].feed_count_only(); } AggColInfo::Numeric(col_idx) => { - if let Some(data) = columns[*col_idx] { - let val = match data { - ColumnData::Float64(v) => v[row_idx], - ColumnData::Int64(v) => v[row_idx] as f64, - ColumnData::Timestamp(v) => v[row_idx] as f64, - _ => continue, - }; - accums[agg_idx].feed(val); + // An integer cell feeds exactly: no rounding through `f64`. + match columns[*col_idx] { + Some(ColumnData::Float64(v)) => accums[agg_idx].feed(v[row_idx]), + Some(ColumnData::Int64(v)) | Some(ColumnData::Timestamp(v)) => { + accums[agg_idx].feed_int(v[row_idx]) + } + _ => {} } } AggColInfo::Skip => {} diff --git a/nodedb/src/event/consumer/run.rs b/nodedb/src/event/consumer/run.rs index b756cfb2d..e47c10478 100644 --- a/nodedb/src/event/consumer/run.rs +++ b/nodedb/src/event/consumer/run.rs @@ -460,8 +460,8 @@ mod tests { #[derive(Debug, PartialEq)] struct Totals { audit_rows: usize, - mv_count: f64, - mv_sum: f64, + mv_count: i64, + mv_sum: i64, crdt_events: u64, } @@ -520,11 +520,16 @@ mod tests { .unwrap_or_else(|p| p.into_inner()) .query_by_event(&AuditEvent::DmlAudit) .len(); + let int = |v: &nodedb_types::Value| match v { + nodedb_types::Value::Integer(i) => *i, + nodedb_types::Value::Null => 0, + other => panic!("an integer COUNT / SUM, got {other:?}"), + }; let (mv_count, mv_sum) = shared .mv_registry .get_state(DatabaseId::DEFAULT, TENANT, "orders_totals") - .and_then(|state| state.read_results().into_iter().next()) - .map_or((0.0, 0.0), |(_, row)| (row[0].1, row[1].1)); + .and_then(|state| state.read_results().unwrap().into_iter().next()) + .map_or((0, 0), |(_, row)| (int(&row[0].1), int(&row[1].1))); let crdt_events = shared.delta_packager.deltas_skipped.load(Ordering::Relaxed) + shared .delta_packager @@ -563,6 +568,7 @@ mod tests { user_id: None, statement_digest: None, commit_hlc: Some(crate::event::test_utils::test_commit_hlc()), + image_fault: None, } } @@ -659,8 +665,8 @@ mod tests { async fn catch_up_mid_run_delivers_every_event_exactly_once() { let expected = Totals { audit_rows: WRITES as usize, - mv_count: WRITES as f64, - mv_sum: (1..=WRITES).sum::() as f64, + mv_count: WRITES as i64, + mv_sum: (1..=WRITES).sum::() as i64, crdt_events: WRITES, }; @@ -780,7 +786,7 @@ mod tests { async fn wait_for_mv_count(shared: &SharedState, count: u64) { tokio::time::timeout(Duration::from_secs(10), async { - while totals(shared).mv_count < count as f64 { + while totals(shared).mv_count < count as i64 { tokio::time::sleep(Duration::from_millis(10)).await; } }) @@ -896,10 +902,10 @@ mod tests { tokio::time::sleep(Duration::from_millis(300)).await; let seen = totals(&node.shared); assert_eq!( - seen.mv_count, WRITES as f64, + seen.mv_count, WRITES as i64, "a view counted an event twice" ); - assert_eq!(seen.mv_sum, (1..=WRITES).sum::() as f64); + assert_eq!(seen.mv_sum, (1..=WRITES).sum::() as i64); assert_eq!( durable_audit_rows(&node.wal), WRITES as usize, diff --git a/nodedb/src/event/streaming_mv/persist.rs b/nodedb/src/event/streaming_mv/persist.rs index a9aed3d49..3ebb836c8 100644 --- a/nodedb/src/event/streaming_mv/persist.rs +++ b/nodedb/src/event/streaming_mv/persist.rs @@ -370,6 +370,19 @@ pub fn spawn_persist_task( mod tests { use super::*; + use crate::event::streaming_mv::state::{AggInput, MvState}; + use crate::event::streaming_mv::types::{AggDef, AggFunction}; + use nodedb_types::Value; + + /// A group state that took `values` for `func`. + fn state_of(func: AggFunction, values: &[i64]) -> GroupState { + let mut state = GroupState::default(); + for v in values { + state.update(func, &AggInput::Value(Value::Integer(*v))); + } + state + } + #[test] fn save_and_load_roundtrip() { let dir = tempfile::tempdir().unwrap(); @@ -378,25 +391,11 @@ mod tests { let snapshot = vec![ ( "INSERT".to_string(), - vec![GroupState { - count: 5, - sum: 100.0, - min: Some(10.0), - max: Some(50.0), - finalized: false, - latest_event_time: 0, - }], + vec![state_of(AggFunction::Sum, &[10, 20, 30, 40])], ), ( "UPDATE".to_string(), - vec![GroupState { - count: 3, - sum: 30.0, - min: Some(5.0), - max: Some(15.0), - finalized: false, - latest_event_time: 0, - }], + vec![state_of(AggFunction::Sum, &[5, 15])], ), ]; @@ -408,9 +407,73 @@ mod tests { .load(DatabaseId::new(1), 1, "order_stats") .unwrap() .unwrap(); - assert_eq!(loaded.len(), 2); - assert_eq!(loaded[0].0, "INSERT"); - assert_eq!(loaded[0].1[0].count, 5); + assert_eq!(loaded, snapshot); + } + + /// Exact state survives a save, a fresh open of the store, and a restore: + /// the restored view reads the same exact values it held before. + #[test] + fn restart_round_trip_keeps_exact_values() { + const ABOVE: i64 = 9_007_199_254_740_993; + const AT: i64 = 9_007_199_254_740_992; + let aggregates = vec![ + AggDef { + output_name: "s".into(), + function: AggFunction::Sum, + input_expr: "v".into(), + }, + AggDef { + output_name: "lo".into(), + function: AggFunction::Min, + input_expr: "v".into(), + }, + AggDef { + output_name: "hi".into(), + function: AggFunction::Max, + input_expr: "v".into(), + }, + AggDef { + output_name: "mean".into(), + function: AggFunction::Avg, + input_expr: "v".into(), + }, + ]; + let before = MvState::new("m".into(), Vec::new(), aggregates.clone()); + for v in [ + Value::Integer(ABOVE), + Value::Integer(AT), + Value::Integer(1_700_000_000_000_000_001), + Value::from_u64(u64::MAX), + ] { + let input = AggInput::Value(v); + before.update_with_time("", &vec![input; 4], 1); + } + let expected = before.read_results().unwrap(); + + let dir = tempfile::tempdir().unwrap(); + { + let persist = MvPersistence::open(dir.path()).unwrap(); + persist + .save(DatabaseId::new(1), 1, "m", &before.snapshot()) + .unwrap(); + } + let persist = MvPersistence::open(dir.path()).unwrap(); + let after = MvState::new("m".into(), Vec::new(), aggregates); + after.restore(persist.load(DatabaseId::new(1), 1, "m").unwrap().unwrap()); + + assert_eq!(after.read_results().unwrap(), expected); + assert_eq!( + expected[0].1[0].1, + Value::Decimal(rust_decimal::Decimal::from_i128_with_scale( + i128::from(ABOVE) + + i128::from(AT) + + 1_700_000_000_000_000_001 + + i128::from(u64::MAX), + 0 + )) + ); + assert_eq!(expected[0].1[1].1, Value::Integer(AT)); + assert_eq!(expected[0].1[2].1, Value::from_u64(u64::MAX)); } #[test] diff --git a/nodedb/src/event/streaming_mv/processor.rs b/nodedb/src/event/streaming_mv/processor.rs index d078db559..70108d4f8 100644 --- a/nodedb/src/event/streaming_mv/processor.rs +++ b/nodedb/src/event/streaming_mv/processor.rs @@ -14,7 +14,7 @@ use crate::event::cdc::event::CdcEvent; use crate::event::types::WriteEvent; use super::registry::MvRegistry; -use super::state::MvState; +use super::state::{AggInput, MvState}; use super::types::{AggDef, AggFunction}; /// Process a WriteEvent against all streaming MVs sourced from a given stream. @@ -109,14 +109,14 @@ fn update_mv(event: &CdcEvent, mv_state: &MvState) { // Extract GROUP BY key from the event. let group_key = extract_group_key(event, &mv_state.group_by_columns); - // Extract aggregate values from the event. - let agg_values: Vec = mv_state + // Extract aggregate inputs from the event. + let inputs: Vec = mv_state .aggregates .iter() - .map(|agg| extract_agg_value(event, agg)) + .map(|agg| extract_agg_input(event, agg)) .collect(); - mv_state.update_with_time(&group_key, &agg_values, event.event_time); + mv_state.update_with_time(&group_key, &inputs, event.event_time); trace!( mv = %mv_state.name, @@ -157,36 +157,32 @@ fn extract_group_key(event: &CdcEvent, group_by_columns: &[String]) -> String { parts.join(":") } -/// Extract a numeric value for an aggregate function from the event. +/// Extract one aggregate's input from the event. /// -/// For COUNT, always returns 1.0 (each event counts as one). -/// For SUM/MIN/MAX/AVG, the aggregate's source column — the same one the -/// definition-time redaction refusal decides on — is read from new_value. -fn extract_agg_value(event: &CdcEvent, agg: &AggDef) -> f64 { +/// COUNT takes the event itself. SUM / MIN / MAX / AVG take the value of the +/// aggregate's source column — the same one the definition-time redaction +/// refusal decides on — read from new_value as the exact `Value` it holds: +/// an integer stays an integer, a `u64` above `i64::MAX` stays exact. A +/// missing or null field is [`AggInput::Absent`]. Which values a function +/// takes is the state's rule (`GroupState::update`). +fn extract_agg_input(event: &CdcEvent, agg: &AggDef) -> AggInput { if agg.function == AggFunction::Count { - return 1.0; // Each event counts as one. + return AggInput::Event; } let Some(field_name) = agg.source_field() else { // An aggregate naming no column has nothing to read. - return f64::NAN; // NaN → skipped by GroupState::update. + return AggInput::Absent; }; - - // Look up the field in new_value. - event - .new_value - .as_ref() - .and_then(|v| v.get(field_name)) - .and_then(|v| match v { - serde_json::Value::Number(n) => n.as_f64(), - serde_json::Value::String(s) => s.parse::().ok(), - _ => None, - }) - .unwrap_or(f64::NAN) // NaN → skipped by GroupState::update. + match event.new_value.as_ref().and_then(|v| v.get(field_name)) { + None | Some(serde_json::Value::Null) => AggInput::Absent, + Some(v) => AggInput::Value(nodedb_types::conversion::json_to_value_ref(v)), + } } #[cfg(test)] mod tests { use super::*; + use nodedb_types::Value; fn agg(function: AggFunction, input_expr: &str) -> AggDef { AggDef { @@ -241,40 +237,74 @@ mod tests { } #[test] - fn extract_agg_value_count() { + fn extract_agg_input_count() { let event = make_event("INSERT", 99.0); - assert_eq!(extract_agg_value(&event, &agg(AggFunction::Count, "")), 1.0); + assert_eq!( + extract_agg_input(&event, &agg(AggFunction::Count, "")), + AggInput::Event + ); } #[test] - fn extract_agg_value_sum() { + fn extract_agg_input_sum() { let event = make_event("INSERT", 42.5); assert_eq!( - extract_agg_value(&event, &agg(AggFunction::Sum, "total")), - 42.5 + extract_agg_input(&event, &agg(AggFunction::Sum, "total")), + AggInput::Value(Value::Float(42.5)) + ); + } + + /// An integer field stays an exact integer, and a `u64` above `i64::MAX` + /// keeps its number. + #[test] + fn extract_agg_input_keeps_integers_exact() { + let mut event = make_event("INSERT", 0.0); + event.new_value = Some(serde_json::json!({ + "big": 9_007_199_254_740_993_i64, + "huge": u64::MAX, + "none": null, + })); + assert_eq!( + extract_agg_input(&event, &agg(AggFunction::Sum, "big")), + AggInput::Value(Value::Integer(9_007_199_254_740_993)) + ); + assert_eq!( + extract_agg_input(&event, &agg(AggFunction::Max, "huge")), + AggInput::Value(Value::from_u64(u64::MAX)) + ); + assert_eq!( + extract_agg_input(&event, &agg(AggFunction::Min, "none")), + AggInput::Absent + ); + assert_eq!( + extract_agg_input(&event, &agg(AggFunction::Min, "missing")), + AggInput::Absent ); } /// A `doc_get` wrapper names the same stored column the plain form does, so /// both the extraction and the definition-time refusal see one field name. #[test] - fn extract_agg_value_reads_a_doc_get_wrapped_field() { + fn extract_agg_input_reads_a_doc_get_wrapped_field() { let event = make_event("INSERT", 7.5); assert_eq!( - extract_agg_value( + extract_agg_input( &event, &agg(AggFunction::Sum, "doc_get(new_value, '$.total')") ), - 7.5 + AggInput::Value(Value::Float(7.5)) ); } - /// A non-COUNT aggregate naming no column reads nothing, and must stay NaN - /// so `GroupState::update` skips it rather than counting it as a value. + /// A non-COUNT aggregate naming no column reads nothing, so + /// `GroupState::update` takes no input from it. #[test] - fn extract_agg_value_without_a_column_is_nan() { + fn extract_agg_input_without_a_column_is_absent() { let event = make_event("INSERT", 42.5); - assert!(extract_agg_value(&event, &agg(AggFunction::Sum, " ")).is_nan()); + assert_eq!( + extract_agg_input(&event, &agg(AggFunction::Sum, " ")), + AggInput::Absent + ); } #[test] @@ -315,13 +345,13 @@ mod tests { let state = registry .get_state(crate::types::DatabaseId::new(7), 1, "order_stats") .unwrap(); - let results = state.read_results(); + let results = state.read_results().unwrap(); let insert_row = results.iter().find(|(k, _)| k == "INSERT").unwrap(); - assert_eq!(insert_row.1[0].1, 2.0); // COUNT = 2 - assert_eq!(insert_row.1[1].1, 150.0); // SUM = 150 + assert_eq!(insert_row.1[0].1, Value::Integer(2)); // COUNT = 2 + assert_eq!(insert_row.1[1].1, Value::Float(150.0)); // SUM = 150 let update_row = results.iter().find(|(k, _)| k == "UPDATE").unwrap(); - assert_eq!(update_row.1[0].1, 1.0); // COUNT = 1 + assert_eq!(update_row.1[0].1, Value::Integer(1)); // COUNT = 1 } } diff --git a/nodedb/src/event/streaming_mv/query.rs b/nodedb/src/event/streaming_mv/query.rs index 492c4fe79..badce23a1 100644 --- a/nodedb/src/event/streaming_mv/query.rs +++ b/nodedb/src/event/streaming_mv/query.rs @@ -2,131 +2,258 @@ //! Query streaming MV results as Arrow RecordBatch. //! -//! Converts the in-memory aggregate state to a RecordBatch that -//! DataFusion can serve via MemTable. +//! Converts the in-memory aggregate state to a RecordBatch. +//! +//! Each aggregate column takes the narrowest Arrow type that holds every +//! value in it exactly: +//! +//! - `Int64` when every value is an `Integer`. +//! - `Decimal128(38, 0)` when every value is an integer, some past `i64`. +//! - `Float64` when the values are numbers and at least one is a float. +//! - `Utf8` (the value's display text) for any other mix, such as a MIN +//! over text. +//! +//! NULL stays NULL in every type. use std::sync::Arc; -use arrow::array::{Float64Array, StringArray}; -use arrow::datatypes::{DataType, Field, Schema, SchemaRef}; +use arrow::array::{ + ArrayRef, BooleanArray, Decimal128Array, Float64Array, Int64Array, StringArray, +}; +use arrow::datatypes::{DataType, Field, Schema}; use arrow::record_batch::RecordBatch; +use nodedb_types::Value; +use rust_decimal::prelude::ToPrimitive; use super::state::MvState; -/// Build the Arrow schema for an MV's results. -/// -/// Columns: one Utf8 column per GROUP BY key, then one Float64 column per aggregate. -pub fn mv_result_schema(mv_state: &MvState) -> SchemaRef { - let mut fields: Vec = mv_state - .group_by_columns - .iter() - .map(|col| Field::new(col, DataType::Utf8, false)) - .collect(); - - for agg in &mv_state.aggregates { - fields.push(Field::new(&agg.output_name, DataType::Float64, false)); +/// Precision of an integer column past `i64`: the widest `Decimal128`. +const DECIMAL_PRECISION: u8 = 38; + +/// The Arrow type of one aggregate column. +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +enum ColumnKind { + Int64, + Decimal, + Float64, + Utf8, +} + +impl ColumnKind { + fn data_type(self) -> DataType { + match self { + Self::Int64 => DataType::Int64, + Self::Decimal => DataType::Decimal128(DECIMAL_PRECISION, 0), + Self::Float64 => DataType::Float64, + Self::Utf8 => DataType::Utf8, + } } +} - // Finalization status column. - fields.push(Field::new("finalized", DataType::Boolean, false)); +/// An integral `Decimal` as its exact `i128`. +fn decimal_integer(d: &rust_decimal::Decimal) -> Option { + if d.fract().is_zero() { + d.trunc().to_i128() + } else { + None + } +} + +/// The narrowest kind that holds every value of a column exactly. +fn column_kind<'a>(values: impl Iterator) -> ColumnKind { + let mut kind = ColumnKind::Int64; + for value in values { + let this = match value { + Value::Null | Value::Integer(_) => ColumnKind::Int64, + Value::Decimal(d) if decimal_integer(d).is_some() => ColumnKind::Decimal, + Value::Float(_) | Value::Decimal(_) => ColumnKind::Float64, + _ => return ColumnKind::Utf8, + }; + kind = match (kind, this) { + (ColumnKind::Float64, _) | (_, ColumnKind::Float64) => ColumnKind::Float64, + (ColumnKind::Decimal, _) | (_, ColumnKind::Decimal) => ColumnKind::Decimal, + _ => ColumnKind::Int64, + }; + } + kind +} - Arc::new(Schema::new(fields)) +/// One aggregate column as an Arrow array of `kind`. +fn build_column(values: &[&Value], kind: ColumnKind) -> crate::Result { + Ok(match kind { + ColumnKind::Int64 => Arc::new(Int64Array::from( + values + .iter() + .map(|v| match v { + Value::Integer(i) => Some(*i), + _ => None, + }) + .collect::>(), + )), + ColumnKind::Decimal => Arc::new( + Decimal128Array::from( + values + .iter() + .map(|v| match v { + Value::Integer(i) => Some(i128::from(*i)), + Value::Decimal(d) => decimal_integer(d), + _ => None, + }) + .collect::>(), + ) + .with_precision_and_scale(DECIMAL_PRECISION, 0) + .map_err(|e| crate::Error::Serialization { + format: "arrow".into(), + detail: format!("streaming MV decimal column: {e}"), + })?, + ), + ColumnKind::Float64 => Arc::new(Float64Array::from( + values + .iter() + .map(|v| match v { + Value::Decimal(d) => d.to_f64(), + other => other.as_f64(), + }) + .collect::>(), + )), + ColumnKind::Utf8 => Arc::new(StringArray::from( + values + .iter() + .map(|v| (!v.is_null()).then(|| v.to_string())) + .collect::>(), + )), + }) } /// Convert MV state to a RecordBatch. /// -/// Returns None if the MV has no data yet. -pub fn mv_state_to_record_batch(mv_state: &MvState) -> Option { - let results = mv_state.read_results_with_status(); +/// Columns: one Utf8 column per GROUP BY key, then one column per aggregate +/// typed by its values, then the `finalized` flag. `None` when the MV has no +/// data yet. Fails when an aggregate total left the exact range, or Arrow +/// refuses a column. +pub fn mv_state_to_record_batch(mv_state: &MvState) -> crate::Result> { + let results = mv_state.read_results_with_status()?; if results.is_empty() { - return None; + return Ok(None); } - let schema = mv_result_schema(mv_state); let num_group_cols = mv_state.group_by_columns.len(); - let num_aggs = mv_state.aggregates.len(); + let mut fields: Vec = Vec::new(); + let mut columns: Vec = Vec::new(); - let mut group_arrays: Vec> = vec![Vec::new(); num_group_cols]; - let mut agg_arrays: Vec> = vec![Vec::new(); num_aggs]; - let mut finalized_array: Vec = Vec::new(); - - for (key, agg_values, finalized) in &results { - let parts: Vec<&str> = key.splitn(num_group_cols, ':').collect(); - for (i, col_values) in group_arrays.iter_mut().enumerate() { - col_values.push(parts.get(i).unwrap_or(&"").to_string()); - } - - for (i, agg_col) in agg_arrays.iter_mut().enumerate() { - let val = agg_values.get(i).map(|(_, v)| *v).unwrap_or(0.0); - agg_col.push(val); - } - - finalized_array.push(*finalized); + for (i, col) in mv_state.group_by_columns.iter().enumerate() { + let values: Vec = results + .iter() + .map(|(key, _, _)| { + key.splitn(num_group_cols, ':') + .nth(i) + .unwrap_or("") + .to_string() + }) + .collect(); + fields.push(Field::new(col, DataType::Utf8, false)); + columns.push(Arc::new(StringArray::from(values))); } - let mut columns: Vec> = Vec::new(); - for group_col in &group_arrays { - columns.push(Arc::new(StringArray::from( - group_col.iter().map(|s| s.as_str()).collect::>(), - ))); - } - for agg_col in &agg_arrays { - columns.push(Arc::new(Float64Array::from(agg_col.clone()))); + for (i, agg) in mv_state.aggregates.iter().enumerate() { + let values: Vec<&Value> = results + .iter() + .map(|(_, row, _)| row.get(i).map_or(&Value::Null, |(_, v)| v)) + .collect(); + let kind = column_kind(values.iter().copied()); + fields.push(Field::new(&agg.output_name, kind.data_type(), true)); + columns.push(build_column(&values, kind)?); } - columns.push(Arc::new(arrow::array::BooleanArray::from(finalized_array))); - RecordBatch::try_new(schema, columns).ok() + fields.push(Field::new("finalized", DataType::Boolean, false)); + columns.push(Arc::new(BooleanArray::from( + results.iter().map(|(_, _, f)| *f).collect::>(), + ))); + + RecordBatch::try_new(Arc::new(Schema::new(fields)), columns) + .map(Some) + .map_err(|e| crate::Error::Serialization { + format: "arrow".into(), + detail: format!("streaming MV record batch: {e}"), + }) } #[cfg(test)] mod tests { use super::*; + use crate::event::streaming_mv::state::AggInput; use crate::event::streaming_mv::types::{AggDef, AggFunction}; + use arrow::array::Array; + + fn agg(output_name: &str, function: AggFunction) -> AggDef { + AggDef { + output_name: output_name.into(), + function, + input_expr: "v".into(), + } + } #[test] - fn schema_reflects_definition() { + fn state_to_batch() { let state = MvState::new( "test".into(), vec!["event_type".into()], - vec![ - AggDef { - output_name: "cnt".into(), - function: AggFunction::Count, - input_expr: String::new(), - }, - AggDef { - output_name: "total".into(), - function: AggFunction::Sum, - input_expr: "amount".into(), - }, - ], + vec![agg("cnt", AggFunction::Count)], ); - let schema = mv_result_schema(&state); - assert_eq!(schema.fields().len(), 4); // event_type + cnt + total + finalized - assert_eq!(schema.field(0).name(), "event_type"); - assert_eq!(schema.field(1).name(), "cnt"); - assert_eq!(schema.field(2).name(), "total"); + state.update_with_time("INSERT", &[AggInput::Event], 0); + state.update_with_time("INSERT", &[AggInput::Event], 0); + state.update_with_time("DELETE", &[AggInput::Event], 0); + + let batch = mv_state_to_record_batch(&state).unwrap().unwrap(); + assert_eq!(batch.num_rows(), 2); + assert_eq!(batch.num_columns(), 3); // event_type + cnt + finalized + assert_eq!(batch.schema().field(0).name(), "event_type"); + assert_eq!(batch.schema().field(1).data_type(), &DataType::Int64); } #[test] - fn state_to_batch() { + fn integer_columns_stay_exact() { let state = MvState::new( "test".into(), - vec!["event_type".into()], - vec![AggDef { - output_name: "cnt".into(), - function: AggFunction::Count, - input_expr: String::new(), - }], + vec!["g".into()], + vec![agg("s", AggFunction::Sum), agg("hi", AggFunction::Max)], ); + let big = AggInput::Value(Value::Integer(i64::MAX)); + state.update_with_time("a", &[big.clone(), big.clone()], 0); + state.update_with_time("a", &[big.clone(), big], 0); + let small = AggInput::Value(Value::Integer(9_007_199_254_740_993)); + state.update_with_time("b", &[small.clone(), small], 0); - state.update_with_time("INSERT", &[1.0], 0); - state.update_with_time("INSERT", &[1.0], 0); - state.update_with_time("DELETE", &[1.0], 0); + let batch = mv_state_to_record_batch(&state).unwrap().unwrap(); + let sums = batch + .column(1) + .as_any() + .downcast_ref::() + .unwrap(); + assert_eq!(sums.value(0), 2 * i128::from(i64::MAX)); + assert_eq!(sums.value(1), 9_007_199_254_740_993); + let maxes = batch + .column(2) + .as_any() + .downcast_ref::() + .unwrap(); + assert_eq!(maxes.value(0), i64::MAX); + assert_eq!(maxes.value(1), 9_007_199_254_740_993); + } - let batch = mv_state_to_record_batch(&state).unwrap(); - assert_eq!(batch.num_rows(), 2); - assert_eq!(batch.num_columns(), 3); // event_type + cnt + finalized + #[test] + fn column_kind_widens_only_as_needed() { + let i = Value::Integer(1); + let d = Value::from_u64(u64::MAX); + let f = Value::Float(0.5); + let s = Value::String("x".into()); + assert_eq!( + column_kind([&i, &Value::Null].into_iter()), + ColumnKind::Int64 + ); + assert_eq!(column_kind([&i, &d].into_iter()), ColumnKind::Decimal); + assert_eq!(column_kind([&d, &f].into_iter()), ColumnKind::Float64); + assert_eq!(column_kind([&f, &s].into_iter()), ColumnKind::Utf8); } } diff --git a/nodedb/src/event/streaming_mv/state.rs b/nodedb/src/event/streaming_mv/state.rs index 1b84854b1..c3605ae76 100644 --- a/nodedb/src/event/streaming_mv/state.rs +++ b/nodedb/src/event/streaming_mv/state.rs @@ -5,27 +5,58 @@ //! Supports incremental updates: each incoming event updates only the //! affected group key's state. O(1) per event, not O(N) rescan. //! +//! The state follows the ad-hoc aggregate rule, so a view row equals the +//! same `SELECT ... GROUP BY` over the source rows: +//! +//! - COUNT counts events. +//! - SUM / AVG take numeric inputs into an [`ExactSum`]: an integer adds +//! exactly, a float adds Kahan-compensated. SUM is an `Integer`, a +//! `Decimal` past `i64`, or a `Float` when any input was a float. AVG is a +//! `Float` from the exact total. +//! - MIN / MAX keep the original non-null input, compared exactly by +//! [`value_replaces`]. +//! - An aggregate with no input is NULL, except COUNT, which is `0`. +//! //! State is stored in-memory (HashMap) and persisted to redb periodically. use std::collections::HashMap; use std::sync::RwLock; +use nodedb_query::window::extremum::value_replaces; +use nodedb_query::{EvalError, ExactSum}; +use nodedb_types::Value; + use super::types::{AggDef, AggFunction}; /// A row of aggregate results: (aggregate_name, value). -pub type AggRow = Vec<(String, f64)>; +pub type AggRow = Vec<(String, Value)>; /// MV result row: (group_key, aggregate_values, finalized). pub type MvResultRow = (String, AggRow, bool); +/// One event's input to one aggregate. +#[derive(Debug, Clone, PartialEq)] +pub enum AggInput { + /// The event itself, for COUNT. + Event, + /// The aggregate's non-null source value. + Value(Value), + /// The event carries no value for the aggregate. + Absent, +} + /// Partial aggregate state for one group key. -#[derive(Debug, Clone, Default, zerompk::ToMessagePack, zerompk::FromMessagePack)] +#[derive(Debug, Clone, Default, PartialEq, zerompk::ToMessagePack, zerompk::FromMessagePack)] #[msgpack(map)] pub struct GroupState { + /// Inputs taken: events for COUNT, values for every other function. pub count: u64, - pub sum: f64, - pub min: Option, - pub max: Option, + /// Exact SUM / AVG state. + pub sum: ExactSum, + /// Smallest input, as received. + pub min: Option, + /// Largest input, as received. + pub max: Option, /// Whether this bucket is finalized (all partitions have advanced past it). /// Once finalized, no more events will arrive for this group key. #[msgpack(default)] @@ -37,12 +68,30 @@ pub struct GroupState { } impl GroupState { - /// Update this state with a new value and event timestamp. - pub fn update(&mut self, value: f64) { - self.count += 1; - self.sum += value; - self.min = Some(self.min.map_or(value, |m| m.min(value))); - self.max = Some(self.max.map_or(value, |m| m.max(value))); + /// Take one input for `func`. Returns whether the input counted: SUM / + /// AVG take only a numeric value, MIN / MAX any value, COUNT the event. + pub fn update(&mut self, func: AggFunction, input: &AggInput) -> bool { + let taken = match (func, input) { + (AggFunction::Count, AggInput::Event) => true, + (AggFunction::Sum | AggFunction::Avg, AggInput::Value(v)) => self.sum.add_value(v), + (AggFunction::Min, AggInput::Value(v)) => { + if value_replaces(v, self.min.as_ref(), false) { + self.min = Some(v.clone()); + } + true + } + (AggFunction::Max, AggInput::Value(v)) => { + if value_replaces(v, self.max.as_ref(), true) { + self.max = Some(v.clone()); + } + true + } + _ => false, + }; + if taken { + self.count += 1; + } + taken } /// Update the latest event time for this group key. @@ -53,19 +102,15 @@ impl GroupState { } /// Compute a specific aggregate from this state. - pub fn compute(&self, func: AggFunction) -> f64 { + pub fn compute(&self, func: AggFunction) -> Result { match func { - AggFunction::Count => self.count as f64, - AggFunction::Sum => self.sum, - AggFunction::Min => self.min.unwrap_or(0.0), - AggFunction::Max => self.max.unwrap_or(0.0), - AggFunction::Avg => { - if self.count > 0 { - self.sum / self.count as f64 - } else { - 0.0 - } - } + AggFunction::Count => i64::try_from(self.count) + .map(Value::Integer) + .map_err(|_| EvalError::NumericOverflow { function: "count" }), + AggFunction::Sum => self.sum.sum(), + AggFunction::Avg => self.sum.avg(), + AggFunction::Min => Ok(self.min.clone().unwrap_or(Value::Null)), + AggFunction::Max => Ok(self.max.clone().unwrap_or(Value::Null)), } } } @@ -97,48 +142,48 @@ impl MvState { /// Update the MV state with a new event. /// /// `group_key` is the concatenated GROUP BY values (e.g., "INSERT" or "orders:INSERT"). - /// `agg_values` is one value per aggregate definition (NaN if not applicable). + /// `inputs` is one input per aggregate definition. /// `event_time_ms` is the wall-clock timestamp of the event. - pub fn update_with_time(&self, group_key: &str, agg_values: &[f64], event_time_ms: u64) { + pub fn update_with_time(&self, group_key: &str, inputs: &[AggInput], event_time_ms: u64) { let mut groups = self.groups.write().unwrap_or_else(|p| p.into_inner()); let states = groups .entry(group_key.to_string()) .or_insert_with(|| vec![GroupState::default(); self.aggregates.len()]); - for (i, &value) in agg_values.iter().enumerate() { - if i < states.len() && !value.is_nan() { - states[i].update(value); - states[i].update_event_time(event_time_ms); + for ((state, agg), input) in states.iter_mut().zip(&self.aggregates).zip(inputs) { + if state.update(agg.function, input) { + state.update_event_time(event_time_ms); } } } - /// Read the current aggregate results. + /// The aggregate values of one group's states, in definition order. + fn row(&self, states: &[GroupState]) -> Result { + self.aggregates + .iter() + .enumerate() + .map(|(i, agg)| { + let val = match states.get(i) { + Some(state) => state.compute(agg.function)?, + None => GroupState::default().compute(agg.function)?, + }; + Ok((agg.output_name.clone(), val)) + }) + .collect() + } + + /// Read the current aggregate results, sorted by group key. /// - /// Returns: Vec<(group_key, Vec<(agg_name, value)>)> - pub fn read_results(&self) -> Vec<(String, AggRow)> { + /// Fails with [`EvalError::NumericOverflow`] when a SUM or AVG total left + /// the exact range. + pub fn read_results(&self) -> Result, EvalError> { let groups = self.groups.read().unwrap_or_else(|p| p.into_inner()); - let mut results: Vec<(String, AggRow)> = groups + let mut results = groups .iter() - .map(|(key, states)| { - let values: Vec<(String, f64)> = self - .aggregates - .iter() - .enumerate() - .map(|(i, agg)| { - let val = if i < states.len() { - states[i].compute(agg.function) - } else { - 0.0 - }; - (agg.output_name.clone(), val) - }) - .collect(); - (key.clone(), values) - }) - .collect(); + .map(|(key, states)| Ok((key.clone(), self.row(states)?))) + .collect::, EvalError>>()?; results.sort_by(|a, b| a.0.cmp(&b.0)); - results + Ok(results) } /// Finalize time buckets whose latest event_time is below the watermark time. @@ -167,33 +212,20 @@ impl MvState { finalized_count } - /// Read results with finalization status. + /// Read results with finalization status, sorted by group key. /// - /// Returns: Vec<(group_key, Vec<(agg_name, value)>, finalized)> - pub fn read_results_with_status(&self) -> Vec { + /// Fails like [`Self::read_results`]. + pub fn read_results_with_status(&self) -> Result, EvalError> { let groups = self.groups.read().unwrap_or_else(|p| p.into_inner()); - let mut results: Vec = groups + let mut results = groups .iter() .map(|(key, states)| { let finalized = states.iter().all(|s| s.finalized); - let values: Vec<(String, f64)> = self - .aggregates - .iter() - .enumerate() - .map(|(i, agg)| { - let val = if i < states.len() { - states[i].compute(agg.function) - } else { - 0.0 - }; - (agg.output_name.clone(), val) - }) - .collect(); - (key.clone(), values, finalized) + Ok((key.clone(), self.row(states)?, finalized)) }) - .collect(); + .collect::, EvalError>>()?; results.sort_by(|a, b| a.0.cmp(&b.0)); - results + Ok(results) } /// Number of distinct group keys. @@ -232,18 +264,103 @@ impl MvState { mod tests { use super::*; + const ABOVE: i64 = 9_007_199_254_740_993; + const AT: i64 = 9_007_199_254_740_992; + + fn agg(output_name: &str, function: AggFunction) -> AggDef { + AggDef { + output_name: output_name.into(), + function, + input_expr: "v".into(), + } + } + + fn value(i: i64) -> AggInput { + AggInput::Value(Value::Integer(i)) + } + #[test] fn group_state_incremental() { let mut gs = GroupState::default(); - gs.update(10.0); - gs.update(20.0); - gs.update(5.0); + for v in [10.0, 20.0, 5.0] { + assert!(gs.update(AggFunction::Avg, &AggInput::Value(Value::Float(v)))); + } assert_eq!(gs.count, 3); - assert_eq!(gs.sum, 35.0); - assert_eq!(gs.min, Some(5.0)); - assert_eq!(gs.max, Some(20.0)); - assert!((gs.compute(AggFunction::Avg) - 11.666666).abs() < 0.01); + assert_eq!(gs.compute(AggFunction::Sum).unwrap(), Value::Float(35.0)); + let Value::Float(avg) = gs.compute(AggFunction::Avg).unwrap() else { + panic!("AVG is a float"); + }; + assert!((avg - 11.666666).abs() < 0.01); + } + + #[test] + fn integers_above_two_pow_53_stay_exact() { + let state = MvState::new( + "m".into(), + Vec::new(), + vec![ + agg("cnt", AggFunction::Count), + agg("s", AggFunction::Sum), + agg("lo", AggFunction::Min), + agg("hi", AggFunction::Max), + agg("mean", AggFunction::Avg), + ], + ); + for v in [ABOVE, AT, i64::MAX] { + state.update_with_time( + "", + &[AggInput::Event, value(v), value(v), value(v), value(v)], + 0, + ); + } + let rows = state.read_results().unwrap(); + let row: Vec = rows[0].1.iter().map(|(_, v)| v.clone()).collect(); + assert_eq!(row[0], Value::Integer(3)); + assert_eq!( + row[1], + Value::Decimal(rust_decimal::Decimal::from_i128_with_scale( + i128::from(ABOVE) + i128::from(AT) + i128::from(i64::MAX), + 0 + )) + ); + assert_eq!(row[2], Value::Integer(AT)); + assert_eq!(row[3], Value::Integer(i64::MAX)); + let exact_mean = (i128::from(ABOVE) + i128::from(AT) + i128::from(i64::MAX)) / 3; + assert_eq!(row[4], Value::Float(exact_mean as f64)); + } + + #[test] + fn absent_and_non_numeric_inputs_do_not_count() { + let mut gs = GroupState::default(); + assert!(!gs.update(AggFunction::Sum, &AggInput::Absent)); + assert!(!gs.update( + AggFunction::Sum, + &AggInput::Value(Value::String("x".into())) + )); + assert_eq!(gs.count, 0); + assert_eq!(gs.compute(AggFunction::Sum).unwrap(), Value::Null); + assert_eq!(gs.compute(AggFunction::Avg).unwrap(), Value::Null); + assert_eq!(gs.compute(AggFunction::Min).unwrap(), Value::Null); + assert_eq!(gs.compute(AggFunction::Count).unwrap(), Value::Integer(0)); + } + + #[test] + fn snapshot_round_trips_through_msgpack() { + let mut gs = GroupState::default(); + for v in [ABOVE, i64::MAX, i64::MAX] { + gs.update(AggFunction::Sum, &value(v)); + } + gs.min = Some(Value::Integer(AT)); + gs.max = Some(Value::from_u64(u64::MAX)); + gs.latest_event_time = 7; + let bytes = zerompk::to_msgpack_vec(&gs).unwrap(); + let back: GroupState = zerompk::from_msgpack(&bytes).unwrap(); + assert_eq!(back, gs); + assert_eq!( + back.compute(AggFunction::Sum).unwrap(), + gs.compute(AggFunction::Sum).unwrap() + ); } #[test] @@ -258,17 +375,17 @@ mod tests { }], ); - state.update_with_time("INSERT", &[1.0], 0); - state.update_with_time("INSERT", &[1.0], 0); - state.update_with_time("UPDATE", &[1.0], 0); + state.update_with_time("INSERT", &[AggInput::Event], 0); + state.update_with_time("INSERT", &[AggInput::Event], 0); + state.update_with_time("UPDATE", &[AggInput::Event], 0); - let results = state.read_results(); + let results = state.read_results().unwrap(); assert_eq!(results.len(), 2); let insert_row = results.iter().find(|(k, _)| k == "INSERT").unwrap(); - assert_eq!(insert_row.1[0].1, 2.0); // COUNT = 2 + assert_eq!(insert_row.1[0].1, Value::Integer(2)); let update_row = results.iter().find(|(k, _)| k == "UPDATE").unwrap(); - assert_eq!(update_row.1[0].1, 1.0); // COUNT = 1 + assert_eq!(update_row.1[0].1, Value::Integer(1)); } } diff --git a/nodedb/tests/inproc/cases/event_streaming_mv.rs b/nodedb/tests/inproc/cases/event_streaming_mv.rs index 051783422..56e01d7e2 100644 --- a/nodedb/tests/inproc/cases/event_streaming_mv.rs +++ b/nodedb/tests/inproc/cases/event_streaming_mv.rs @@ -6,37 +6,45 @@ //! finalization, backfill from buffer, state persistence + restore. use nodedb::event::streaming_mv::persist::MvPersistence; -use nodedb::event::streaming_mv::state::{GroupState, MvState}; +use nodedb::event::streaming_mv::state::{AggInput, GroupState, MvState}; use nodedb::event::streaming_mv::types::{AggDef, AggFunction}; use nodedb::types::DatabaseId; +use nodedb_types::Value; + +fn count_def() -> Vec { + vec![AggDef { + output_name: "cnt".to_string(), + function: AggFunction::Count, + input_expr: String::new(), + }] +} + +fn int(v: i64) -> AggInput { + AggInput::Value(Value::Integer(v)) +} #[test] fn incremental_count() { let state = MvState::new( "order_counts".to_string(), vec!["op".to_string()], - vec![AggDef { - output_name: "cnt".to_string(), - function: AggFunction::Count, - input_expr: String::new(), - }], + count_def(), ); - // Each update passes &[1.0]; GroupState::update increments count by 1 per call. - state.update_with_time("group_a", &[1.0], 0); - state.update_with_time("group_a", &[1.0], 0); - state.update_with_time("group_b", &[1.0], 0); + state.update_with_time("group_a", &[AggInput::Event], 0); + state.update_with_time("group_a", &[AggInput::Event], 0); + state.update_with_time("group_b", &[AggInput::Event], 0); - let results = state.read_results_with_status(); + let results = state.read_results_with_status().unwrap(); assert_eq!(results.len(), 2); // MvResultRow = (group_key, AggRow, finalized) // AggRow = Vec<(output_name, value)> let group_a = results.iter().find(|r| r.0 == "group_a").unwrap(); - assert_eq!(group_a.1[0].1, 2.0); // COUNT = 2 + assert_eq!(group_a.1[0].1, Value::Integer(2)); let group_b = results.iter().find(|r| r.0 == "group_b").unwrap(); - assert_eq!(group_b.1[0].1, 1.0); // COUNT = 1 + assert_eq!(group_b.1[0].1, Value::Integer(1)); } #[test] @@ -64,19 +72,16 @@ fn incremental_sum_min_max() { ); // Pass the same value to all three aggregate slots. - state.update_with_time("bucket", &[10.0, 10.0, 10.0], 0); - state.update_with_time("bucket", &[30.0, 30.0, 30.0], 0); - state.update_with_time("bucket", &[20.0, 20.0, 20.0], 0); + for v in [10, 30, 20] { + state.update_with_time("bucket", &[int(v), int(v), int(v)], 0); + } - let results = state.read_results_with_status(); + let results = state.read_results_with_status().unwrap(); let bucket = results.iter().find(|r| r.0 == "bucket").unwrap(); - // Sum aggregate (index 0): SUM of 10+30+20 = 60. - assert!((bucket.1[0].1 - 60.0).abs() < f64::EPSILON); - // Min aggregate (index 1): MIN of 10, 30, 20 = 10. - assert!((bucket.1[1].1 - 10.0).abs() < f64::EPSILON); - // Max aggregate (index 2): MAX of 10, 30, 20 = 30. - assert!((bucket.1[2].1 - 30.0).abs() < f64::EPSILON); + assert_eq!(bucket.1[0].1, Value::Integer(60)); + assert_eq!(bucket.1[1].1, Value::Integer(10)); + assert_eq!(bucket.1[2].1, Value::Integer(30)); } #[test] @@ -91,14 +96,14 @@ fn incremental_avg() { }], ); - state.update_with_time("g", &[10.0], 0); - state.update_with_time("g", &[20.0], 0); - state.update_with_time("g", &[30.0], 0); + for v in [10, 20, 30] { + state.update_with_time("g", &[int(v)], 0); + } - let results = state.read_results_with_status(); + let results = state.read_results_with_status().unwrap(); let g = results.iter().find(|r| r.0 == "g").unwrap(); // AVG = SUM / COUNT = 60 / 3 = 20. - assert!((g.1[0].1 - 20.0).abs() < f64::EPSILON); + assert_eq!(g.1[0].1, Value::Float(20.0)); } #[test] @@ -106,22 +111,18 @@ fn watermark_finalization() { let state = MvState::new( "event_counts".to_string(), vec!["group".to_string()], - vec![AggDef { - output_name: "cnt".to_string(), - function: AggFunction::Count, - input_expr: String::new(), - }], + count_def(), ); // Use update_with_time so latest_event_time is populated for finalization. - state.update_with_time("early", &[1.0], 1000); - state.update_with_time("late", &[1.0], 5000); + state.update_with_time("early", &[AggInput::Event], 1000); + state.update_with_time("late", &[AggInput::Event], 5000); // Finalize groups with latest_event_time < 3000. let finalized = state.finalize_buckets(3000); assert_eq!(finalized, 1); // Only "early" finalized. - let results = state.read_results_with_status(); + let results = state.read_results_with_status().unwrap(); // MvResultRow = (group_key, AggRow, finalized_bool) let early = results.iter().find(|r| r.0 == "early").unwrap(); assert!(early.2); // finalized = true @@ -135,15 +136,11 @@ fn snapshot_and_restore() { let state = MvState::new( "snap_mv".to_string(), vec!["group".to_string()], - vec![AggDef { - output_name: "cnt".to_string(), - function: AggFunction::Count, - input_expr: String::new(), - }], + count_def(), ); - state.update_with_time("g1", &[1.0], 0); - state.update_with_time("g1", &[1.0], 0); - state.update_with_time("g2", &[1.0], 0); + state.update_with_time("g1", &[AggInput::Event], 0); + state.update_with_time("g1", &[AggInput::Event], 0); + state.update_with_time("g2", &[AggInput::Event], 0); let snapshot = state.snapshot(); assert_eq!(snapshot.len(), 2); @@ -152,17 +149,23 @@ fn snapshot_and_restore() { let restored = MvState::new( "snap_mv".to_string(), vec!["group".to_string()], - vec![AggDef { - output_name: "cnt".to_string(), - function: AggFunction::Count, - input_expr: String::new(), - }], + count_def(), ); restored.restore(snapshot); - let results = restored.read_results_with_status(); + let results = restored.read_results_with_status().unwrap(); let g1 = results.iter().find(|r| r.0 == "g1").unwrap(); - assert_eq!(g1.1[0].1, 2.0); // COUNT = 2 + assert_eq!(g1.1[0].1, Value::Integer(2)); +} + +/// A group state that took `values` for `func`, at `event_time`. +fn state_of(func: AggFunction, values: &[i64], event_time: u64) -> GroupState { + let mut state = GroupState::default(); + for v in values { + state.update(func, &int(*v)); + } + state.update_event_time(event_time); + state } #[test] @@ -170,29 +173,14 @@ fn persistence_save_and_load() { let dir = tempfile::tempdir().unwrap(); let persist = MvPersistence::open(dir.path()).unwrap(); + let mut update = state_of(AggFunction::Sum, &[5, 10, 15], 3000); + update.finalized = true; let snapshot = vec![ ( "INSERT".to_string(), - vec![GroupState { - count: 5, - sum: 100.0, - min: Some(10.0), - max: Some(50.0), - finalized: false, - latest_event_time: 5000, - }], - ), - ( - "UPDATE".to_string(), - vec![GroupState { - count: 3, - sum: 30.0, - min: Some(5.0), - max: Some(15.0), - finalized: true, - latest_event_time: 3000, - }], + vec![state_of(AggFunction::Sum, &[10, 20, 30, 40], 5000)], ), + ("UPDATE".to_string(), vec![update]), ]; persist @@ -202,8 +190,8 @@ fn persistence_save_and_load() { .load(DatabaseId::DEFAULT, 1, "order_stats") .unwrap() .unwrap(); - assert_eq!(loaded.len(), 2); - assert_eq!(loaded[0].1[0].count, 5); + assert_eq!(loaded, snapshot); + assert_eq!(loaded[0].1[0].count, 4); assert!(loaded[1].1[0].finalized); } diff --git a/nodedb/tests/inproc/cases/sql_streaming_mv.rs b/nodedb/tests/inproc/cases/sql_streaming_mv.rs index f63e13291..d63dfb1a1 100644 --- a/nodedb/tests/inproc/cases/sql_streaming_mv.rs +++ b/nodedb/tests/inproc/cases/sql_streaming_mv.rs @@ -15,7 +15,9 @@ use std::time::Duration; +use nodedb::event::streaming_mv::state::AggRow; use nodedb_test_support::pgwire_harness::TestServer; +use nodedb_types::Value; /// The default harness superuser (`nodedb`) is provisioned under tenant id 1. const TENANT_ID: u64 = 1; @@ -70,25 +72,10 @@ async fn streaming_mv_incrementally_aggregates_source_writes() { // The Event Plane consumes WriteEvents asynchronously. Poll the registry // until the MV state materializes both group keys (or time out). This is a // deterministic convergence poll, not a blind fixed sleep. - let mut results: Vec<(String, Vec<(String, f64)>)> = Vec::new(); - for _ in 0..80 { - if let Some(state) = server.shared.mv_registry.get_state( - nodedb::types::DatabaseId::DEFAULT, - TENANT_ID, - "smv_order_stats", - ) { - results = state.read_results(); - if results.len() >= 2 { - break; - } - } - tokio::time::sleep(Duration::from_millis(50)).await; - } + let results = await_groups(&server, "smv_order_stats", 2).await; // A correctly-wired streaming MV registers its definition on CREATE and // then incrementally aggregates the three source writes into two groups. - // Today the neutral DDL handler never registers a `StreamingMvDef`, so the - // registry has no state for the view and this assertion fails. assert_eq!( results.len(), 2, @@ -100,12 +87,14 @@ async fn streaming_mv_incrementally_aggregates_source_writes() { .find(|(k, _)| k == "active") .expect("`active` group must be present in streaming MV state"); // Aggregate order matches the SELECT list: index 0 = COUNT(*), 1 = SUM(amount). - assert!( - (active.1[0].1 - 2.0).abs() < f64::EPSILON, + assert_eq!( + active.1[0].1, + Value::Integer(2), "active COUNT(*) must be 2; got {active:?}" ); - assert!( - (active.1[1].1 - 30.0).abs() < f64::EPSILON, + assert_eq!( + active.1[1].1, + Value::Integer(30), "active SUM(amount) must be 10 + 20 = 30; got {active:?}" ); @@ -113,16 +102,148 @@ async fn streaming_mv_incrementally_aggregates_source_writes() { .iter() .find(|(k, _)| k == "pending") .expect("`pending` group must be present in streaming MV state"); - assert!( - (pending.1[0].1 - 1.0).abs() < f64::EPSILON, + assert_eq!( + pending.1[0].1, + Value::Integer(1), "pending COUNT(*) must be 1; got {pending:?}" ); - assert!( - (pending.1[1].1 - 5.0).abs() < f64::EPSILON, + assert_eq!( + pending.1[1].1, + Value::Integer(5), "pending SUM(amount) must be 5; got {pending:?}" ); } +/// Poll `view`'s state until it holds at least `groups` group keys, or time +/// out. The Event Plane consumes WriteEvents asynchronously, so this is a +/// convergence poll, not a blind fixed sleep. +async fn await_groups(server: &TestServer, view: &str, groups: usize) -> Vec<(String, AggRow)> { + let mut results = Vec::new(); + for _ in 0..80 { + if let Some(state) = + server + .shared + .mv_registry + .get_state(nodedb::types::DatabaseId::DEFAULT, TENANT_ID, view) + { + results = state.read_results().expect("view totals stay in range"); + if results.len() >= groups { + break; + } + } + tokio::time::sleep(Duration::from_millis(50)).await; + } + results +} + +/// Poll `view` until its `group` row's COUNT (aggregate 0) reaches `count`. +async fn await_count(server: &TestServer, view: &str, group: &str, count: i64) -> AggRow { + let mut row = Vec::new(); + for _ in 0..80 { + if let Some((_, found)) = await_groups(server, view, 1) + .await + .into_iter() + .find(|(k, _)| k == group) + { + row = found; + if row.first().map(|(_, v)| v) == Some(&Value::Integer(count)) { + break; + } + } + tokio::time::sleep(Duration::from_millis(50)).await; + } + row +} + +/// The cells of the ad-hoc `sql` result's one row. +async fn ad_hoc_row(server: &TestServer, sql: &str) -> Vec { + let rows = server.query_rows(sql).await.unwrap(); + assert_eq!(rows.len(), 1, "{sql}: one row, got {rows:?}"); + rows[0].clone() +} + +/// A view value as the text the ad-hoc query returns for it. +fn text(v: &Value) -> String { + v.to_string() +} + +/// A streaming view over integers above 2^53 and nanosecond timestamps holds +/// the same COUNT / SUM / MIN / MAX / AVG the ad-hoc query computes over the +/// source rows, including a SUM past `i64::MAX`. +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn streaming_mv_integer_aggregates_equal_ad_hoc() { + let server = TestServer::start().await; + server.exec("CREATE COLLECTION smv_exact").await.unwrap(); + server + .exec("CREATE CHANGE STREAM smv_exact_changes ON smv_exact") + .await + .unwrap(); + server + .exec( + "CREATE MATERIALIZED VIEW smv_exact_stats ON smv_exact STREAMING AS \ + SELECT g, COUNT(*) AS cnt, SUM(v) AS total, MIN(v) AS lo, MAX(v) AS hi, \ + AVG(v) AS mean FROM smv_exact_changes GROUP BY g", + ) + .await + .unwrap(); + + let groups: [(&str, &[i64]); 2] = [ + ( + "big", + &[9_007_199_254_740_993, 9_007_199_254_740_992, i64::MAX], + ), + ( + "nanos", + &[ + 1_700_000_000_000_000_002, + 1_700_000_000_000_000_001, + 1_700_000_000_000_000_003, + ], + ), + ]; + for (g, values) in groups { + for (i, v) in values.iter().enumerate() { + server + .exec(&format!( + "INSERT INTO smv_exact {{ id: '{g}{i}', g: '{g}', v: {v} }}" + )) + .await + .unwrap(); + } + } + + for (g, values) in groups { + let view = await_count(&server, "smv_exact_stats", g, values.len() as i64).await; + let ad_hoc = ad_hoc_row( + &server, + &format!( + "SELECT COUNT(*), SUM(v), MIN(v), MAX(v), AVG(v) FROM smv_exact WHERE g = '{g}'" + ), + ) + .await; + let cells: Vec<&Value> = view.iter().map(|(_, v)| v).collect(); + assert_eq!(cells.len(), 5, "{g}: view row {view:?}"); + for (i, name) in ["COUNT", "SUM", "MIN", "MAX"].iter().enumerate() { + assert_eq!( + text(cells[i]), + ad_hoc[i], + "{g}: view {name} must equal the ad-hoc {name}" + ); + } + let view_avg = cells[4].as_f64().expect("AVG is a float"); + let ad_hoc_avg: f64 = ad_hoc[4].parse().expect("ad-hoc AVG is numeric"); + assert_eq!( + view_avg, ad_hoc_avg, + "{g}: view AVG must equal the ad-hoc AVG" + ); + } + + // The SUM past `i64::MAX` is exact, not rounded or wrapped. + let big = await_count(&server, "smv_exact_stats", "big", 3).await; + let exact: i128 = 9_007_199_254_740_993 + 9_007_199_254_740_992 + i128::from(i64::MAX); + assert_eq!(text(&big[1].1), exact.to_string()); +} + /// Install a mask on `collection`.`field` for a single role, exactly as the /// metadata applier does when a policy replicates. fn install_mask(server: &TestServer, collection: &str, field: &str) { From 31bec311f830fc0730faedd9660f9398355fc411 Mon Sep 17 00:00:00 2001 From: Farhan Syah Date: Thu, 8 Oct 2026 16:51:29 +0800 Subject: [PATCH 05/24] feat(types): enforce DECIMAL(p,s) typmods ColumnType::Decimal carries an optional typmod. A plain DECIMAL keeps every digit it is given. DECIMAL(p,s) rounds a written value to s digits and refuses one with more than p-s integer digits with 22003. A typmod outside the supported range is refused at CREATE and ALTER with 22023. --- fuzz/src/targets/strict_tuple.rs | 10 +- nodedb-sql/src/parser/type_expr.rs | 26 +- .../src/planner/declared_type_coerce.rs | 148 ++++++++- nodedb-sql/src/planner/predicate_coerce.rs | 3 +- nodedb-sql/src/types_expr.rs | 8 +- nodedb-types/src/columnar/column_parse.rs | 137 +++++--- nodedb-types/src/columnar/column_type.rs | 173 +++++----- nodedb-types/src/columnar/decimal_typmod.rs | 289 +++++++++++++++++ nodedb-types/src/columnar/mod.rs | 5 + nodedb-types/src/columnar/schema.rs | 18 +- nodedb-types/src/kv.rs | 5 +- nodedb-types/tests/wire_enum_lock.rs | 16 +- .../neutral/collection/alter/add_column.rs | 35 +- .../ddl/neutral/collection/create/build.rs | 5 + .../shared/ddl/neutral/convert/column_defs.rs | 10 +- .../shared/ddl/neutral/convert/type_map.rs | 13 +- .../shared/ddl/neutral/declared_typmod.rs | 101 ++++++ .../control/server/shared/ddl/neutral/mod.rs | 1 + nodedb/tests/wire/cases/sql_decimal_typmod.rs | 306 ++++++++++++++++++ 19 files changed, 1090 insertions(+), 219 deletions(-) create mode 100644 nodedb-types/src/columnar/decimal_typmod.rs create mode 100644 nodedb/src/control/server/shared/ddl/neutral/declared_typmod.rs create mode 100644 nodedb/tests/wire/cases/sql_decimal_typmod.rs 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-sql/src/parser/type_expr.rs b/nodedb-sql/src/parser/type_expr.rs index 7bcba19dc..558c39755 100644 --- a/nodedb-sql/src/parser/type_expr.rs +++ b/nodedb-sql/src/parser/type_expr.rs @@ -7,7 +7,7 @@ //! at write time. use nodedb_types::Value; -use nodedb_types::columnar::ColumnType; +use nodedb_types::columnar::{ColumnType, DecimalTypmod}; use crate::error::SqlError; @@ -43,12 +43,9 @@ pub enum SimpleType { /// Engine-assigned instant. Matching mirrors /// `ColumnType::SystemTimestamp`: a client writes an instant, never text. SystemTimestamp, - /// Declared precision and scale carry the spelling the author wrote. + /// The declared typmod, `None` for a plain `DECIMAL`. /// Matching is variant-level, as it is for [`SimpleType::Vector`]. - Decimal { - precision: u8, - scale: u8, - }, + Decimal(Option), Uuid, Ulid, Geometry, @@ -285,7 +282,7 @@ fn simple_from_column_type( ColumnType::Timestamp => SimpleType::Timestamp, ColumnType::Timestamptz => SimpleType::Timestamptz, ColumnType::SystemTimestamp => SimpleType::SystemTimestamp, - ColumnType::Decimal { precision, scale } => SimpleType::Decimal { precision, scale }, + ColumnType::Decimal(typmod) => SimpleType::Decimal(typmod), ColumnType::Geometry => SimpleType::Geometry, ColumnType::Vector(dim) => SimpleType::Vector(dim), ColumnType::SparseVector => SimpleType::SparseVector, @@ -383,7 +380,7 @@ fn value_matches_simple(value: &Value, simple: &SimpleType) -> bool { // Mirrors `ColumnType::SystemTimestamp.accepts`: an engine-assigned // instant takes no text form. SimpleType::SystemTimestamp => matches!(value, Value::DateTime(_) | Value::Integer(_)), - SimpleType::Decimal { .. } => matches!( + SimpleType::Decimal(_) => matches!( value, Value::Decimal(_) | Value::Float(_) | Value::Integer(_) | Value::String(_) ), @@ -600,18 +597,15 @@ mod tests { fn parse_decimal_carries_declared_params() { assert_eq!( parse_type_expr("DECIMAL(10, 2)").unwrap(), - TypeExpr::Simple(SimpleType::Decimal { - precision: 10, - scale: 2 - }) + TypeExpr::Simple(SimpleType::Decimal(Some( + DecimalTypmod::new(10, 2).expect("valid typmod") + ))) ); assert_eq!( parse_type_expr("NUMERIC").unwrap(), - TypeExpr::Simple(SimpleType::Decimal { - precision: 38, - scale: 10 - }) + TypeExpr::Simple(SimpleType::Decimal(None)) ); + assert!(parse_type_expr("DECIMAL(1001,0)").is_err()); assert!(parse_type_expr("DECIMAL(0)").is_err()); } diff --git a/nodedb-sql/src/planner/declared_type_coerce.rs b/nodedb-sql/src/planner/declared_type_coerce.rs index d098ab358..3dbd953b6 100644 --- a/nodedb-sql/src/planner/declared_type_coerce.rs +++ b/nodedb-sql/src/planner/declared_type_coerce.rs @@ -13,9 +13,9 @@ //! storing it — the strict document encoder parses that string back into a //! float — so the stored cell always matches what the column declared. //! -//! Two engines have no typed write path — key-value and document-schemaless. -//! Both store the value bytes they are handed, and the declared schema lives -//! only in the catalog, in the Control Plane. Left uncoerced, a column declared +//! Two engines have no typed schema — key-value and document-schemaless. The +//! Data Plane re-types only their declared integer, float, and `DECIMAL(p,s)` +//! columns, and stores every other value as handed. Left uncoerced, a column declared //! `REAL` / `DOUBLE` holds the string `"1.5"` while `RowDescription` advertises //! float4 / float8, and the pgwire encoder — which correctly refuses to encode //! a non-number under a numeric OID — transmits SQL NULL. The client asked for @@ -51,15 +51,32 @@ //! parsed here for the same reason, so an unparseable spelling is refused at //! the statement instead of being stored as text that no read can render. //! -//! Every other declared type either has one unambiguous literal form already -//! or (like `DECIMAL`) is deliberately carried as text for exactness, and +//! A [`SqlDataType::Decimal`] column with a declared `DECIMAL(p,s)` typmod +//! is fitted to it here, for every engine. The value rounds to `s` digits and +//! one with more than `p - s` integer digits is refused. A plain `DECIMAL` +//! keeps every digit it is given. +//! +//! Every other declared type has one unambiguous literal form already and //! passes through untouched. //! //! The numeric conversions mirror the strict document encoder's //! `coerce_value`, so the two engines accept and reject exactly the same //! literals for a given declared type. - +//! +//! # Relation to the Data Plane rule +//! +//! The key-value and document-schemaless write paths re-type every value +//! stored under a declared integer, float, or `DECIMAL(p,s)` column on the +//! Data Plane, through the strict encoder's `coerce_value` and the declared +//! width. That rule covers computed values, which this pass never sees. This +//! pass still runs for literals: it refuses a bad literal at plan time with +//! the column named, and the DDL `DEFAULT` gate and the predicate coercion +//! share it. Fitting is idempotent, so a literal fitted here is unchanged by +//! the Data Plane pass and is never rounded twice. + +use nodedb_types::columnar::DecimalTypmod; use nodedb_types::datetime::NdbDateTime; +use rust_decimal::Decimal; use rust_decimal::prelude::ToPrimitive; use super::dml_helpers::{check_declared_float_ranges, check_declared_int_ranges}; @@ -144,8 +161,8 @@ pub(super) fn coerce_rows_to_declared_types( /// [`coerce_row_to_declared_types`] for `SET col = ` assignments. /// /// Only literal assignments carry a value at plan time; a computed assignment -/// (`SET n = n + 1`) is evaluated by the engine and is not this pass's to -/// re-type. `exempt_column` carries the same primary-key exemption, for the +/// (`SET n = n + 1`) is evaluated by the engine, which re-types the result by +/// the same numeric rule before it stores the row. `exempt_column` carries the same primary-key exemption, for the /// same reason: `SET pk = ` rewrites the row's identity, which the /// engines derive from the literal's own rendering. pub(super) fn coerce_assignments_to_declared_types( @@ -196,13 +213,15 @@ pub(crate) fn coerce_value( SqlDataType::Timestamp | SqlDataType::Timestamptz => { coerce_to_instant(column, value, declared) } + SqlDataType::Decimal(Some(typmod)) => coerce_to_decimal(column, value, *typmod), SqlDataType::String | SqlDataType::Bool | SqlDataType::Bytes - | SqlDataType::Decimal + | SqlDataType::Decimal(None) | SqlDataType::Uuid | SqlDataType::Vector(_) | SqlDataType::Geometry + | SqlDataType::Json | SqlDataType::Unknown => Ok(value), } } @@ -253,6 +272,33 @@ fn coerce_to_float(column: &str, value: SqlValue) -> Result { } } +/// `DECIMAL(p,s)` column: a numeric literal or numeric text becomes the exact +/// decimal the column stores, fitted by [`DecimalTypmod::fit`]. +/// +/// `fit` is the rule the strict and columnar encoders apply, so every engine +/// rounds and refuses the same values. A value that does not fit is refused +/// as [`SqlError::DecimalOutOfRange`]. `NULL` and non-numeric kinds are left +/// alone, as for the other numeric columns. +fn coerce_to_decimal(column: &str, value: SqlValue, typmod: DecimalTypmod) -> Result { + let exact = match value { + SqlValue::Decimal(d) => d, + SqlValue::Int(i) => Decimal::from(i), + SqlValue::Float(f) => Decimal::try_from(f) + .map_err(|_| not_representable(column, &f.to_string(), "DECIMAL"))?, + SqlValue::String(s) => s + .parse::() + .map_err(|_| not_representable(column, &s, "DECIMAL"))?, + other => return Ok(other), + }; + typmod + .fit(exact) + .map(SqlValue::Decimal) + .map_err(|source| SqlError::DecimalOutOfRange { + column: column.to_string(), + source, + }) +} + /// Timestamp column: every accepted literal becomes the one typed instant /// the column stores, tagged as the declared kind. /// @@ -275,10 +321,11 @@ fn coerce_to_instant(column: &str, value: SqlValue, declared: &SqlDataType) -> R | SqlDataType::Bool | SqlDataType::Bytes | SqlDataType::Timestamp - | SqlDataType::Decimal + | SqlDataType::Decimal(_) | SqlDataType::Uuid | SqlDataType::Vector(_) | SqlDataType::Geometry + | SqlDataType::Json | SqlDataType::Unknown => SqlValue::Timestamp(at), }; match value { @@ -321,10 +368,11 @@ fn instant_declared_name(declared: &SqlDataType) -> &'static str { | SqlDataType::Bool | SqlDataType::Bytes | SqlDataType::Timestamp - | SqlDataType::Decimal + | SqlDataType::Decimal(_) | SqlDataType::Uuid | SqlDataType::Vector(_) | SqlDataType::Geometry + | SqlDataType::Json | SqlDataType::Unknown => "TIMESTAMP", } } @@ -482,25 +530,95 @@ mod tests { ); } - /// Non-numeric declared types impose no representation of their own: a - /// `DECIMAL` column keeps its exact decimal, and a `TEXT` column keeps - /// whatever literal it was given. + /// A plain `DECIMAL` column keeps its exact decimal with every digit, and + /// a `TEXT` column keeps whatever literal it was given. #[test] fn non_numeric_declared_types_pass_through_untouched() { let columns = [ - column("d", SqlDataType::Decimal), + column("d", SqlDataType::Decimal(None)), column("t", SqlDataType::String), ]; assert_eq!( coerced(&columns, "d", decimal("1.5")).expect("decimal columns keep exact decimals"), decimal("1.5") ); + assert_eq!( + coerced(&columns, "d", decimal("12345678901234567890.123456789")) + .expect("a plain DECIMAL has no digit limit"), + decimal("12345678901234567890.123456789") + ); assert_eq!( coerced(&columns, "t", SqlValue::Int(9)).expect("text columns are untouched"), SqlValue::Int(9) ); } + fn decimal_5_2() -> SqlDataType { + SqlDataType::Decimal(Some(DecimalTypmod::new(5, 2).expect("valid typmod"))) + } + + /// A `DECIMAL(5,2)` column rounds every numeric literal to two digits, + /// half away from zero, as the strict and columnar encoders do. + #[test] + fn constrained_decimal_rounds_to_its_scale() { + let columns = [column("d", decimal_5_2())]; + for (input, expected) in [ + (decimal("1.005"), "1.01"), + (decimal("-1.005"), "-1.01"), + (SqlValue::Int(7), "7.00"), + (SqlValue::Float(2.5), "2.50"), + (SqlValue::String("999.994".into()), "999.99"), + ] { + let label = format!("{input:?}"); + let got = coerced(&columns, "d", input).expect(&label); + assert_eq!(got, decimal(expected), "{label}"); + let SqlValue::Decimal(d) = got else { + panic!("{label}: expected a decimal"); + }; + assert_eq!(d.to_string(), expected, "{label} keeps scale 2"); + } + assert_eq!( + coerced(&columns, "d", SqlValue::Null).expect("null passes"), + SqlValue::Null + ); + } + + /// A value whose rounded integer part exceeds `p - s` digits is refused + /// as `DecimalOutOfRange`, on `VALUES`, `SET` and a `DEFAULT` literal. + #[test] + fn constrained_decimal_past_precision_is_out_of_range() { + let columns = [column("d", decimal_5_2())]; + for input in [ + decimal("123456.789"), + decimal("999.995"), + SqlValue::Int(1000), + ] { + let label = format!("{input:?}"); + let err = coerced(&columns, "d", input).expect_err(&label); + assert!( + matches!(err, SqlError::DecimalOutOfRange { ref column, .. } if column == "d"), + "{label}: {err:?}" + ); + } + + let mut assignments = vec![("d".to_string(), SqlExpr::Literal(decimal("123456.789")))]; + let err = coerce_assignments_to_declared_types(&columns, &mut assignments, None) + .expect_err("SET past precision is refused"); + assert!(matches!(err, SqlError::DecimalOutOfRange { .. }), "{err:?}"); + + let err = coerce_write_literal(&columns[0], decimal("123456.789")) + .expect_err("a DEFAULT past precision is refused"); + assert!(matches!(err, SqlError::DecimalOutOfRange { .. }), "{err:?}"); + } + + /// Text that names no decimal is a type mismatch naming the column. + #[test] + fn constrained_decimal_refuses_non_numeric_text() { + let columns = [column("d", decimal_5_2())]; + let err = coerced(&columns, "d", SqlValue::String("abc".into())).expect_err("abc"); + assert!(matches!(err, SqlError::TypeMismatch { .. }), "{err:?}"); + } + /// `2020-03-05T10:00:00Z` as microseconds since the Unix epoch. const EARLY_MICROS: i64 = 1_583_402_400_000_000; diff --git a/nodedb-sql/src/planner/predicate_coerce.rs b/nodedb-sql/src/planner/predicate_coerce.rs index 085b58b6e..b4ebc94cc 100644 --- a/nodedb-sql/src/planner/predicate_coerce.rs +++ b/nodedb-sql/src/planner/predicate_coerce.rs @@ -167,10 +167,11 @@ fn is_instant(declared: &SqlDataType) -> bool { | SqlDataType::String | SqlDataType::Bool | SqlDataType::Bytes - | SqlDataType::Decimal + | SqlDataType::Decimal(_) | SqlDataType::Uuid | SqlDataType::Vector(_) | SqlDataType::Geometry + | SqlDataType::Json | SqlDataType::Unknown => false, } } diff --git a/nodedb-sql/src/types_expr.rs b/nodedb-sql/src/types_expr.rs index f4d7ecca7..9e5c65004 100644 --- a/nodedb-sql/src/types_expr.rs +++ b/nodedb-sql/src/types_expr.rs @@ -5,6 +5,7 @@ //! Re-exported from `types` so downstream `use crate::types::*` continues //! to resolve these symbols without change. +use nodedb_types::columnar::DecimalTypmod; use nodedb_types::datetime::NdbDateTime; use crate::types::SqlPlan; @@ -146,9 +147,14 @@ pub enum SqlDataType { Timestamp, /// Timezone-aware timestamp. Timestamptz, - Decimal, + /// Exact decimal. `Some` carries the declared `DECIMAL(p,s)` typmod that + /// every written value is fitted to. `None` is a plain `DECIMAL`. + Decimal(Option), Uuid, Vector(usize), Geometry, + /// A structured value: a JSON document, or a typed `ARRAY`, `SET`, + /// `RANGE` or `RECORD` cell. It reads back as its JSON text. + Json, Unknown, } diff --git a/nodedb-types/src/columnar/column_parse.rs b/nodedb-types/src/columnar/column_parse.rs index b184352da..54a03a7eb 100644 --- a/nodedb-types/src/columnar/column_parse.rs +++ b/nodedb-types/src/columnar/column_parse.rs @@ -7,6 +7,8 @@ use std::fmt; use std::str::FromStr; use super::column_type::ColumnType; +use super::decimal_typmod::{DecimalTypmod, DecimalTypmodError}; +use crate::error::sqlstate; /// Every declared spelling that resolves to [`ColumnType::Int64`]. /// @@ -51,9 +53,29 @@ pub enum ColumnTypeParseError { #[error("invalid VARCHAR length: '{0}' (must be a positive integer)")] InvalidCharLength(String), #[error( - "invalid DECIMAL/NUMERIC params: '{0}' (expected DECIMAL(precision, scale) with precision 1-38 and scale <= precision)" + "invalid DECIMAL/NUMERIC params: '{0}' (expected DECIMAL(precision) or DECIMAL(precision, scale))" )] InvalidDecimalParams(String), + /// A well-formed `DECIMAL(p,s)` whose precision or scale is out of range. + #[error("{0}")] + InvalidDecimalTypmod(#[from] DecimalTypmodError), +} + +impl ColumnTypeParseError { + /// The SQLSTATE a DDL statement refuses this declared type with. + /// + /// An out-of-range typmod is `22023` (invalid_parameter_value), as + /// PostgreSQL reports it. Every other refusal is `42601` (syntax_error). + pub fn sqlstate(&self) -> &'static str { + match self { + Self::InvalidDecimalTypmod(_) => sqlstate::INVALID_PARAMETER_VALUE, + Self::Unknown(_) + | Self::UseTimestamp + | Self::InvalidVectorDim(_) + | Self::InvalidCharLength(_) + | Self::InvalidDecimalParams(_) => sqlstate::SYNTAX_ERROR, + } + } } impl fmt::Display for ColumnType { @@ -67,7 +89,10 @@ impl fmt::Display for ColumnType { Self::Timestamp => f.write_str("TIMESTAMP"), Self::Timestamptz => f.write_str("TIMESTAMPTZ"), Self::SystemTimestamp => f.write_str("SYSTEM_TIMESTAMP"), - Self::Decimal { precision, scale } => write!(f, "DECIMAL({precision},{scale})"), + Self::Decimal(Some(typmod)) => { + write!(f, "DECIMAL({},{})", typmod.precision(), typmod.scale()) + } + Self::Decimal(None) => f.write_str("DECIMAL"), Self::Geometry => f.write_str("GEOMETRY"), Self::Vector(dim) => write!(f, "VECTOR({dim})"), Self::SparseVector => f.write_str("SPARSEVECTOR"), @@ -90,46 +115,20 @@ impl FromStr for ColumnType { fn from_str(s: &str) -> Result { let upper = s.trim().to_uppercase(); - // NUMERIC(p,s) / DECIMAL(p,s) special case. - if upper.starts_with("NUMERIC") || upper.starts_with("DECIMAL") { - let base = if upper.starts_with("NUMERIC") { - "NUMERIC" - } else { - "DECIMAL" - }; - let rest = upper[base.len()..].trim(); + // NUMERIC / DECIMAL, bare or with a `(p)` / `(p,s)` typmod. A longer + // word such as `NUMERIC_MONEY` is not this keyword and falls through + // to the unknown-type arm. + if let Some(rest) = upper + .strip_prefix("NUMERIC") + .or_else(|| upper.strip_prefix("DECIMAL")) + { + let rest = rest.trim(); if rest.is_empty() { - return Ok(Self::Decimal { - precision: 38, - scale: 10, - }); + return Ok(Self::Decimal(None)); } - if rest.starts_with('(') && rest.ends_with(')') { - let inner = &rest[1..rest.len() - 1]; - let parts: Vec<&str> = inner.splitn(2, ',').collect(); - let precision: u8 = parts[0] - .trim() - .parse() - .map_err(|_| ColumnTypeParseError::InvalidDecimalParams(rest.to_string()))?; - let scale: u8 = parts - .get(1) - .map(|p| p.trim()) - .unwrap_or("0") - .parse() - .map_err(|_| ColumnTypeParseError::InvalidDecimalParams(rest.to_string()))?; - if precision == 0 || precision > 38 { - return Err(ColumnTypeParseError::InvalidDecimalParams(format!( - "precision {precision} out of range 1-38" - ))); - } - if scale > precision { - return Err(ColumnTypeParseError::InvalidDecimalParams(format!( - "scale {scale} must be <= precision {precision}" - ))); - } - return Ok(Self::Decimal { precision, scale }); + if rest.starts_with('(') { + return parse_decimal_typmod(rest).map(|typmod| Self::Decimal(Some(typmod))); } - return Err(ColumnTypeParseError::InvalidDecimalParams(rest.to_string())); } // VECTOR(N) special case. @@ -221,10 +220,38 @@ impl ColumnType { /// here as catalog text resolves to [`ColumnType::Timestamp`]. Pass such a /// spelling to [`str::parse`] instead to resolve it whole. pub fn from_declared_type(declared: &str) -> Option { - bare_declared_token(declared).parse().ok() + Self::parse_declared_type(declared).ok() + } + + /// [`ColumnType::from_declared_type`] with the parse error kept. + /// + /// A DDL gate uses this to refuse a declared type it cannot hold, such as + /// `DECIMAL(1001,0)`, instead of reading it as an unknown type. + pub fn parse_declared_type(declared: &str) -> Result { + bare_declared_token(declared).parse() } } +/// The typmod of `DECIMAL(p)` or `DECIMAL(p,s)`, from the `(p)` or `(p,s)` +/// text after the keyword. A bare `(p)` has scale 0. +/// +/// Text that is not one or two integers in parentheses is malformed. Integers +/// out of range are an [`ColumnTypeParseError::InvalidDecimalTypmod`]. +fn parse_decimal_typmod(params: &str) -> Result { + let malformed = || ColumnTypeParseError::InvalidDecimalParams(params.to_string()); + let inner = params + .strip_prefix('(') + .and_then(|p| p.strip_suffix(')')) + .ok_or_else(malformed)?; + let (precision_text, scale_text) = match inner.split_once(',') { + Some((precision, scale)) => (precision, scale), + None => (inner, "0"), + }; + let precision: i64 = precision_text.trim().parse().map_err(|_| malformed())?; + let scale: i64 = scale_text.trim().parse().map_err(|_| malformed())?; + Ok(DecimalTypmod::new(precision, scale)?) +} + /// The leading type token of a declared DDL type string. /// /// The cut is the first whitespace outside parentheses, so a parameter list @@ -317,10 +344,9 @@ mod tests { fn declared_type_keeps_a_spaced_parameter_list() { assert_eq!( ColumnType::from_declared_type("DECIMAL(10, 2) NOT NULL"), - Some(ColumnType::Decimal { - precision: 10, - scale: 2 - }) + Some(ColumnType::Decimal(Some( + DecimalTypmod::new(10, 2).expect("valid typmod") + ))) ); assert_eq!( ColumnType::from_declared_type("VECTOR(768)"), @@ -328,6 +354,29 @@ mod tests { ); } + /// An out-of-range typmod keeps its typed error through the declared-type + /// parse, and refuses as `22023`. A malformed one refuses as `42601`. + #[test] + fn declared_decimal_typmod_errors_are_typed() { + let err = ColumnType::parse_declared_type("DECIMAL(1001,0) NOT NULL") + .expect_err("precision 1001 is refused"); + assert_eq!( + err, + ColumnTypeParseError::InvalidDecimalTypmod(DecimalTypmodError::PrecisionOutOfRange { + precision: 1001 + }) + ); + assert_eq!(err.sqlstate(), "22023"); + assert_eq!(ColumnType::from_declared_type("DECIMAL(1001,0)"), None); + + let err = ColumnType::parse_declared_type("NUMERIC(x)").expect_err("malformed"); + assert_eq!(err.sqlstate(), "42601"); + assert_eq!( + ColumnType::parse_declared_type("NUMERIC DEFAULT 1"), + Ok(ColumnType::Decimal(None)) + ); + } + /// The two zone spellings resolve whole, so a parser that reads the full /// string reads the zone the author wrote. #[test] diff --git a/nodedb-types/src/columnar/column_type.rs b/nodedb-types/src/columnar/column_type.rs index 7eb45829c..170c6d915 100644 --- a/nodedb-types/src/columnar/column_type.rs +++ b/nodedb-types/src/columnar/column_type.rs @@ -4,15 +4,15 @@ use serde::{Deserialize, Serialize}; +use super::decimal_typmod::DecimalTypmod; use crate::InstantKind; use crate::value::Value; /// Typed column definition for strict document and columnar collections. /// -/// `#[non_exhaustive]` — this enum grows with each type system expansion -/// (e.g. future variants may add `Decimal { precision, scale }` or split -/// `Timestamp`/`TimestampTz`). External exhaustive `match` arms must handle -/// future variants via a typed error arm rather than `_ => unreachable!()`. +/// `#[non_exhaustive]` — this enum grows with each type system expansion. +/// External exhaustive `match` arms must handle future variants via a typed +/// error arm rather than `_ => unreachable!()`. #[non_exhaustive] #[derive( Debug, @@ -42,14 +42,10 @@ pub enum ColumnType { /// layer can reject user-supplied values — the column is populated by the /// engine from HLC at commit. SystemTimestamp, - /// Arbitrary-precision decimal with explicit precision and scale. - /// - /// `precision`: total significant digits, 1–38. `scale`: digits after the - /// decimal point, 0–precision. Default when unspecified: `{38, 10}`. - Decimal { - precision: u8, - scale: u8, - }, + /// Exact decimal. `Some` carries the declared `DECIMAL(p,s)` typmod, and + /// every value is fitted to it. `None` is a plain `DECIMAL`, and its + /// values are stored as given. + Decimal(Option), Geometry, /// Fixed-dimension float32 vector. Vector(u32), @@ -273,6 +269,12 @@ mod tests { crate::datetime::NdbDateTime::from_micros(0) } + fn decimal(precision: i64, scale: i64) -> ColumnType { + ColumnType::Decimal(Some( + DecimalTypmod::new(precision, scale).expect("test typmod is valid"), + )) + } + #[test] fn instant_kind_names_the_variant_a_time_column_reads_back_as() { assert_eq!( @@ -301,14 +303,8 @@ mod tests { assert_eq!(ColumnType::Timestamptz.to_pg_oid(), 1184); assert_eq!(ColumnType::SystemTimestamp.to_pg_oid(), 1114); assert_eq!(ColumnType::Duration.to_pg_oid(), 1186); - assert_eq!( - ColumnType::Decimal { - precision: 38, - scale: 10 - } - .to_pg_oid(), - 1700 - ); + assert_eq!(decimal(28, 10).to_pg_oid(), 1700); + assert_eq!(ColumnType::Decimal(None).to_pg_oid(), 1700); assert_eq!(ColumnType::Uuid.to_pg_oid(), 2950); assert_eq!(ColumnType::Ulid.to_pg_oid(), 2950); assert_eq!(ColumnType::Json.to_pg_oid(), 3802); @@ -458,14 +454,9 @@ mod tests { ColumnType::Timestamp, ColumnType::Timestamptz, ColumnType::Vector(768), - ColumnType::Decimal { - precision: 10, - scale: 2, - }, - ColumnType::Decimal { - precision: 38, - scale: 10, - }, + decimal(10, 2), + decimal(28, 10), + ColumnType::Decimal(None), ] { let s = ct.to_string(); let parsed: ColumnType = s.parse().unwrap(); @@ -477,63 +468,93 @@ mod tests { fn decimal_parse_with_params() { assert_eq!( "NUMERIC(10,2)".parse::().unwrap(), - ColumnType::Decimal { - precision: 10, - scale: 2 - } + decimal(10, 2) ); assert_eq!( - "DECIMAL(38,10)".parse::().unwrap(), - ColumnType::Decimal { - precision: 38, - scale: 10 - } + "DECIMAL(28,10)".parse::().unwrap(), + decimal(28, 10) ); + assert_eq!("DECIMAL(7)".parse::().unwrap(), decimal(7, 0)); assert_eq!( "NUMERIC".parse::().unwrap(), - ColumnType::Decimal { - precision: 38, - scale: 10 - } + ColumnType::Decimal(None) ); assert_eq!( "DECIMAL".parse::().unwrap(), - ColumnType::Decimal { - precision: 38, - scale: 10 - } + ColumnType::Decimal(None) ); + assert_eq!(ColumnType::Decimal(None).to_string(), "DECIMAL"); + assert_eq!(decimal(5, 2).to_string(), "DECIMAL(5,2)"); } #[test] fn decimal_parse_invalid() { - assert!("DECIMAL(5,6)".parse::().is_err()); - assert!("DECIMAL(0,0)".parse::().is_err()); - assert!("DECIMAL(39,0)".parse::().is_err()); + use super::super::column_parse::ColumnTypeParseError; + use super::super::decimal_typmod::DecimalTypmodError; + + assert_eq!( + "DECIMAL(5,6)".parse::(), + Err(ColumnTypeParseError::InvalidDecimalTypmod( + DecimalTypmodError::ScaleOutOfRange { + precision: 5, + scale: 6 + } + )) + ); + assert_eq!( + "DECIMAL(0,0)".parse::(), + Err(ColumnTypeParseError::InvalidDecimalTypmod( + DecimalTypmodError::PrecisionOutOfRange { precision: 0 } + )) + ); + assert_eq!( + "DECIMAL(1001,0)".parse::(), + Err(ColumnTypeParseError::InvalidDecimalTypmod( + DecimalTypmodError::PrecisionOutOfRange { precision: 1001 } + )) + ); + assert_eq!( + "DECIMAL(39,0)".parse::(), + Err(ColumnTypeParseError::InvalidDecimalTypmod( + DecimalTypmodError::PrecisionUnsupported { precision: 39 } + )) + ); + for malformed in ["DECIMAL()", "DECIMAL(a,2)", "DECIMAL(5,2,1)", "NUMERIC(5"] { + assert!( + matches!( + malformed.parse::(), + Err(ColumnTypeParseError::InvalidDecimalParams(_)) + ), + "{malformed}" + ); + } + assert_eq!( + "NUMERIC_MONEY".parse::(), + Err(ColumnTypeParseError::Unknown("NUMERIC_MONEY".into())) + ); } #[test] fn decimal_fixed_size() { - assert_eq!( - ColumnType::Decimal { - precision: 10, - scale: 2 - } - .fixed_size(), - Some(16) - ); + assert_eq!(decimal(10, 2).fixed_size(), Some(16)); + assert_eq!(ColumnType::Decimal(None).fixed_size(), Some(16)); } #[test] fn decimal_to_pg_oid_is_1700() { - assert_eq!( - ColumnType::Decimal { - precision: 10, - scale: 2 - } - .to_pg_oid(), - 1700 - ); + assert_eq!(decimal(10, 2).to_pg_oid(), 1700); + } + + #[test] + fn decimal_wire_forms_roundtrip() { + for ct in [decimal(10, 2), ColumnType::Decimal(None)] { + let bytes = zerompk::to_msgpack_vec(&ct).expect("encode"); + let back: ColumnType = zerompk::from_msgpack(&bytes).expect("decode"); + assert_eq!(back, ct); + let json = serde_json::to_string(&ct).expect("encode json"); + let back: ColumnType = serde_json::from_str(&json).expect("decode json"); + assert_eq!(back, ct); + } } #[test] @@ -547,13 +568,7 @@ mod tests { assert!( ColumnType::Uuid.accepts(&Value::Uuid("550e8400-e29b-41d4-a716-446655440000".into())) ); - assert!( - ColumnType::Decimal { - precision: 38, - scale: 10 - } - .accepts(&Value::Decimal(rust_decimal::Decimal::ZERO)) - ); + assert!(decimal(28, 10).accepts(&Value::Decimal(rust_decimal::Decimal::ZERO))); let naive = Value::NaiveDateTime(nodedb_types_datetime_epoch()); let tz = Value::DateTime(nodedb_types_datetime_epoch()); @@ -576,20 +591,8 @@ mod tests { assert!(ColumnType::Uuid.accepts(&Value::String( "550e8400-e29b-41d4-a716-446655440000".into() ))); - assert!( - ColumnType::Decimal { - precision: 10, - scale: 2 - } - .accepts(&Value::String("99.95".into())) - ); - assert!( - ColumnType::Decimal { - precision: 10, - scale: 2 - } - .accepts(&Value::Float(99.95)) - ); + assert!(decimal(10, 2).accepts(&Value::String("99.95".into()))); + assert!(decimal(10, 2).accepts(&Value::Float(99.95))); assert!(ColumnType::Geometry.accepts(&Value::String("POINT(0 0)".into()))); } diff --git a/nodedb-types/src/columnar/decimal_typmod.rs b/nodedb-types/src/columnar/decimal_typmod.rs new file mode 100644 index 000000000..8103b5cdb --- /dev/null +++ b/nodedb-types/src/columnar/decimal_typmod.rs @@ -0,0 +1,289 @@ +// SPDX-License-Identifier: Apache-2.0 + +//! [`DecimalTypmod`]: the declared precision and scale of a `DECIMAL(p,s)` +//! column, and the rule that fits a value to it. + +use rust_decimal::{Decimal, RoundingStrategy}; +use serde::{Deserialize, Serialize}; + +/// The largest precision PostgreSQL accepts in a `NUMERIC` typmod. +pub const PG_MAX_DECIMAL_PRECISION: i64 = 1000; + +/// The largest precision the engine stores exactly. +/// +/// A stored decimal is a `rust_decimal::Decimal`: a 96-bit mantissa and a +/// scale of at most 28. Every value of 28 digits fits that mantissa. Some +/// values of 29 digits do not. +pub const MAX_DECIMAL_PRECISION: u8 = 28; + +/// The precision and scale a `DECIMAL(p,s)` column declares. +/// +/// A value of this type always holds `1 <= precision <= 28` and +/// `scale <= precision`. [`DecimalTypmod::new`] is the only constructor, and +/// both decoders run it. +#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize)] +#[serde(try_from = "RawDecimalTypmod", into = "RawDecimalTypmod")] +pub struct DecimalTypmod { + precision: u8, + scale: u8, +} + +/// The serde form of [`DecimalTypmod`], before validation. +#[derive(Serialize, Deserialize)] +struct RawDecimalTypmod { + precision: i64, + scale: i64, +} + +/// A declared `DECIMAL(p,s)` that is not a valid typmod. +/// +/// The messages follow PostgreSQL's `numeric` typmod errors. +#[derive(Debug, Clone, PartialEq, Eq, thiserror::Error)] +#[non_exhaustive] +pub enum DecimalTypmodError { + #[error( + "NUMERIC precision {precision} must be between 1 and {max}", + max = PG_MAX_DECIMAL_PRECISION + )] + PrecisionOutOfRange { precision: i64 }, + #[error("NUMERIC scale {scale} must be between 0 and precision {precision}")] + ScaleOutOfRange { precision: i64, scale: i64 }, + #[error( + "NUMERIC precision {precision} exceeds {max}, \ + the largest precision this server stores exactly", + max = MAX_DECIMAL_PRECISION + )] + PrecisionUnsupported { precision: i64 }, +} + +/// A value whose rounded integer part has more digits than its +/// `DECIMAL(p,s)` column holds. +#[derive(Debug, Clone, PartialEq, Eq, thiserror::Error)] +#[error( + "{value} does not fit DECIMAL({precision},{scale}): \ + the value must round to an absolute value less than 10^{max_integer_digits}" +)] +pub struct DecimalOutOfRange { + pub value: Decimal, + pub precision: u8, + pub scale: u8, + pub max_integer_digits: u32, +} + +impl DecimalTypmod { + /// Validate a declared precision and scale. + /// + /// PostgreSQL's range applies first: precision `1..=1000`, scale + /// `0..=precision`. A precision past [`MAX_DECIMAL_PRECISION`] is then + /// refused, because the engine cannot store every value it admits. + pub fn new(precision: i64, scale: i64) -> Result { + if !(1..=PG_MAX_DECIMAL_PRECISION).contains(&precision) { + return Err(DecimalTypmodError::PrecisionOutOfRange { precision }); + } + if !(0..=precision).contains(&scale) { + return Err(DecimalTypmodError::ScaleOutOfRange { precision, scale }); + } + let (Ok(precision_digits), Ok(scale_digits)) = + (u8::try_from(precision), u8::try_from(scale)) + else { + return Err(DecimalTypmodError::PrecisionUnsupported { precision }); + }; + if precision_digits > MAX_DECIMAL_PRECISION { + return Err(DecimalTypmodError::PrecisionUnsupported { precision }); + } + Ok(Self { + precision: precision_digits, + scale: scale_digits, + }) + } + + /// Total significant digits. + pub fn precision(self) -> u8 { + self.precision + } + + /// Digits after the decimal point. + pub fn scale(self) -> u8 { + self.scale + } + + /// `value` fitted to this typmod, as PostgreSQL stores it. + /// + /// The value rounds to `scale` fractional digits, half away from zero, + /// and carries exactly `scale` digits. A result with more than + /// `precision - scale` integer digits does not fit and is refused. + pub fn fit(self, value: Decimal) -> Result { + let scale = u32::from(self.scale); + let mut fitted = + value.round_dp_with_strategy(scale, RoundingStrategy::MidpointAwayFromZero); + let max_integer_digits = u32::from(self.precision - self.scale); + if integer_digits(fitted) > max_integer_digits { + return Err(DecimalOutOfRange { + value, + precision: self.precision, + scale: self.scale, + max_integer_digits, + }); + } + // The fitted value has at most `precision <= 28` digits, so scaling + // up only appends zeros and loses no digit. + fitted.rescale(scale); + Ok(fitted) + } +} + +impl TryFrom for DecimalTypmod { + type Error = DecimalTypmodError; + + fn try_from(raw: RawDecimalTypmod) -> Result { + Self::new(raw.precision, raw.scale) + } +} + +impl From for RawDecimalTypmod { + fn from(typmod: DecimalTypmod) -> Self { + Self { + precision: i64::from(typmod.precision), + scale: i64::from(typmod.scale), + } + } +} + +impl zerompk::ToMessagePack for DecimalTypmod { + fn write(&self, writer: &mut W) -> zerompk::Result<()> { + writer.write_array_len(2)?; + writer.write_u8(self.precision)?; + writer.write_u8(self.scale) + } +} + +impl<'de> zerompk::FromMessagePack<'de> for DecimalTypmod { + fn read>(reader: &mut R) -> zerompk::Result { + reader.check_array_len(2)?; + let precision = reader.read_u8()?; + let scale = reader.read_u8()?; + // An encoded typmod that fails validation is refused, the way an + // unknown marker is. + Self::new(i64::from(precision), i64::from(scale)) + .map_err(|_| zerompk::Error::InvalidMarker(precision)) + } +} + +/// The count of digits left of the decimal point in `d`; zero for `|d| < 1`. +fn integer_digits(d: Decimal) -> u32 { + // A `Decimal` is `mantissa / 10^scale` with `scale <= 28`, so both + // operands fit `u128`. + let integer_part = d.mantissa().unsigned_abs() / 10u128.pow(d.scale()); + integer_part.checked_ilog10().map_or(0, |log| log + 1) +} + +#[cfg(test)] +mod tests { + use super::*; + + fn dec(text: &str) -> Decimal { + text.parse().expect("test decimal parses") + } + + fn typmod(precision: i64, scale: i64) -> DecimalTypmod { + DecimalTypmod::new(precision, scale).expect("test typmod is valid") + } + + #[test] + fn new_accepts_the_engine_range() { + for (precision, scale) in [(1, 0), (1, 1), (5, 2), (28, 0), (28, 28)] { + let t = typmod(precision, scale); + assert_eq!(i64::from(t.precision()), precision); + assert_eq!(i64::from(t.scale()), scale); + } + } + + #[test] + fn new_refuses_what_postgresql_refuses() { + for precision in [0, -1, 1001, i64::MAX] { + assert_eq!( + DecimalTypmod::new(precision, 0), + Err(DecimalTypmodError::PrecisionOutOfRange { precision }) + ); + } + assert_eq!( + DecimalTypmod::new(5, 6), + Err(DecimalTypmodError::ScaleOutOfRange { + precision: 5, + scale: 6 + }) + ); + assert_eq!( + DecimalTypmod::new(5, -1), + Err(DecimalTypmodError::ScaleOutOfRange { + precision: 5, + scale: -1 + }) + ); + } + + #[test] + fn new_refuses_a_precision_the_engine_cannot_store() { + for precision in [29, 38, 300, 1000] { + assert_eq!( + DecimalTypmod::new(precision, 0), + Err(DecimalTypmodError::PrecisionUnsupported { precision }) + ); + } + } + + #[test] + fn fit_rounds_half_away_from_zero_to_the_scale() { + for (input, expected) in [ + ("1.005", "1.01"), + ("-1.005", "-1.01"), + ("1.004", "1.00"), + ("12.5", "12.50"), + ("7", "7.00"), + ("999.994", "999.99"), + ] { + let fitted = typmod(5, 2).fit(dec(input)).expect(input); + assert_eq!(fitted.to_string(), expected, "{input}"); + } + } + + #[test] + fn fit_refuses_too_many_integer_digits() { + for input in ["123456.789", "1000", "999.995", "-1000.00"] { + let err = typmod(5, 2).fit(dec(input)).expect_err(input); + assert_eq!(err.max_integer_digits, 3, "{input}"); + } + assert!(typmod(2, 2).fit(dec("0.995")).is_err()); + assert_eq!(typmod(2, 2).fit(dec("0.994")), Ok(dec("0.99"))); + } + + #[test] + fn fit_holds_every_28_digit_value() { + let widest = dec("9999999999999999999999999999"); + assert_eq!(typmod(28, 0).fit(widest), Ok(widest)); + let fraction = dec("0.9999999999999999999999999999"); + assert_eq!(typmod(28, 28).fit(fraction), Ok(fraction)); + } + + #[test] + fn msgpack_roundtrip_and_validation() { + let t = typmod(10, 2); + let bytes = zerompk::to_msgpack_vec(&t).expect("encode"); + let back: DecimalTypmod = zerompk::from_msgpack(&bytes).expect("decode"); + assert_eq!(back, t); + + let invalid = zerompk::to_msgpack_vec(&(5u8, 6u8)).expect("encode raw pair"); + assert!(zerompk::from_msgpack::(&invalid).is_err()); + } + + #[test] + fn serde_roundtrip_and_validation() { + let t = typmod(10, 2); + let json = serde_json::to_string(&t).expect("encode"); + assert_eq!(json, r#"{"precision":10,"scale":2}"#); + let back: DecimalTypmod = serde_json::from_str(&json).expect("decode"); + assert_eq!(back, t); + assert!(serde_json::from_str::(r#"{"precision":0,"scale":0}"#).is_err()); + assert!(serde_json::from_str::(r#"{"precision":39,"scale":0}"#).is_err()); + } +} diff --git a/nodedb-types/src/columnar/mod.rs b/nodedb-types/src/columnar/mod.rs index 5b2bca556..5e57070c3 100644 --- a/nodedb-types/src/columnar/mod.rs +++ b/nodedb-types/src/columnar/mod.rs @@ -3,6 +3,7 @@ pub mod column_def; pub mod column_parse; pub mod column_type; +pub mod decimal_typmod; pub mod declared_type_keyword; pub mod dml_wal_record; pub mod float_width; @@ -17,6 +18,10 @@ pub mod wal_record; pub use column_def::{ColumnDef, ColumnModifier}; pub use column_parse::{ColumnTypeParseError, DECLARED_FLOAT_KEYWORDS, DECLARED_INT_KEYWORDS}; pub use column_type::ColumnType; +pub use decimal_typmod::{ + DecimalOutOfRange, DecimalTypmod, DecimalTypmodError, MAX_DECIMAL_PRECISION, + PG_MAX_DECIMAL_PRECISION, +}; pub use declared_type_keyword::declared_type_matches; pub use dml_wal_record::ColumnarDmlWalRecord; pub use float_width::FloatWidth; diff --git a/nodedb-types/src/columnar/schema.rs b/nodedb-types/src/columnar/schema.rs index 77ac980eb..975bbb922 100644 --- a/nodedb-types/src/columnar/schema.rs +++ b/nodedb-types/src/columnar/schema.rs @@ -370,13 +370,7 @@ mod tests { let schema = StrictSchema::new(vec![ ColumnDef::required("id", ColumnType::Int64).with_primary_key(), ColumnDef::nullable("name", ColumnType::String), - ColumnDef::nullable( - "balance", - ColumnType::Decimal { - precision: 18, - scale: 4, - }, - ), + ColumnDef::nullable("balance", ColumnType::Decimal(None)), ]) .unwrap(); assert_eq!(schema.len(), 3); @@ -390,13 +384,7 @@ mod tests { let schema = StrictSchema::new(vec![ ColumnDef::required("id", ColumnType::Int64).with_primary_key(), ColumnDef::nullable("name", ColumnType::String), - ColumnDef::nullable( - "balance", - ColumnType::Decimal { - precision: 18, - scale: 4, - }, - ), + ColumnDef::nullable("balance", ColumnType::Decimal(None)), ColumnDef::nullable("bio", ColumnType::String), ]) .unwrap(); @@ -427,6 +415,8 @@ mod tests { generated_expr: None, generated_deps: Vec::new(), added_at_version: 1, + int_width: None, + float_width: None, }]; assert!(matches!( StrictSchema::new(cols), diff --git a/nodedb-types/src/kv.rs b/nodedb-types/src/kv.rs index b7d6963b3..37fbea4a1 100644 --- a/nodedb-types/src/kv.rs +++ b/nodedb-types/src/kv.rs @@ -161,10 +161,7 @@ mod tests { assert!(!is_valid_kv_key_type(&ColumnType::Bool)); assert!(!is_valid_kv_key_type(&ColumnType::Geometry)); assert!(!is_valid_kv_key_type(&ColumnType::Vector(128))); - assert!(!is_valid_kv_key_type(&ColumnType::Decimal { - precision: 18, - scale: 4 - })); + assert!(!is_valid_kv_key_type(&ColumnType::Decimal(None))); } #[test] diff --git a/nodedb-types/tests/wire_enum_lock.rs b/nodedb-types/tests/wire_enum_lock.rs index 794b4cf72..4a1d4d502 100644 --- a/nodedb-types/tests/wire_enum_lock.rs +++ b/nodedb-types/tests/wire_enum_lock.rs @@ -162,12 +162,19 @@ fn column_type_wire_forms() { } // Parametric variants. - let dec = ColumnType::Decimal { - precision: 10, - scale: 2, - }; + let dec = ColumnType::Decimal(Some( + nodedb_types::columnar::DecimalTypmod::new(10, 2).expect("valid typmod"), + )); let v = serde_json::to_value(dec).expect("serialize"); assert_eq!(v["type"], "Decimal", "ColumnType::Decimal wire tag"); + assert_eq!( + v["params"]["precision"], 10, + "DecimalTypmod precision field" + ); + assert_eq!(v["params"]["scale"], 2, "DecimalTypmod scale field"); + let plain = serde_json::to_value(ColumnType::Decimal(None)).expect("serialize"); + assert_eq!(plain["type"], "Decimal", "plain DECIMAL wire tag"); + assert!(plain["params"].is_null(), "plain DECIMAL carries no typmod"); let vec = ColumnType::Vector(128); let v = serde_json::to_value(vec).expect("serialize"); @@ -454,7 +461,6 @@ fn query_mode_wire_forms() { match qm { QueryMode::Or => {} QueryMode::And => {} - _ => panic!("unrecognized QueryMode — update wire_enum_lock.rs"), } let v = serde_json::to_value(qm).expect("serialize"); assert_eq!(v, json!(expected), "QueryMode::{expected} wire form"); diff --git a/nodedb/src/control/server/shared/ddl/neutral/collection/alter/add_column.rs b/nodedb/src/control/server/shared/ddl/neutral/collection/alter/add_column.rs index 2f3acef1e..47f69cdcd 100644 --- a/nodedb/src/control/server/shared/ddl/neutral/collection/alter/add_column.rs +++ b/nodedb/src/control/server/shared/ddl/neutral/collection/alter/add_column.rs @@ -16,6 +16,7 @@ use crate::control::server::shared::ddl::neutral::collection::helpers::parse_ori use crate::control::server::shared::ddl::neutral::column_default::{ DeclaredColumn, validate_column_default, }; +use crate::control::server::shared::ddl::neutral::declared_typmod::validate_declared_typmod; use crate::control::server::shared::ddl::result::{DdlError, DdlResult}; use crate::control::state::SharedState; @@ -31,16 +32,26 @@ pub(super) async fn alter_table_add_column( ) -> Result, DdlError> { let tenant_id = identity.tenant_id; + // The declared type as written, with its modifiers: `SMALLINT NOT NULL` + // from `age SMALLINT NOT NULL`, the same text `CREATE` records for a + // column. `ColumnDef::column_type` cannot supply this: it has one `Int64` + // variant for every integer width. A spaced parameter list such as + // `DECIMAL(10, 2)` stays whole. + let written_type = col_def_str + .trim_start() + .split_once(char::is_whitespace) + .map(|(name, declared)| (name, declared.trim())); + // An invalid `DECIMAL(p,s)` typmod is refused with the SQLSTATE `CREATE` + // gives it. + if let Some((name, declared)) = written_type { + validate_declared_typmod(name, declared)?; + } let column = parse_origin_column_def(col_def_str).map_err(|e| err("42601", e.to_string()))?; let column_name = column.name.clone(); - // The declared type as written, e.g. `SMALLINT` from `age SMALLINT NOT - // NULL`. `ColumnDef::column_type` cannot supply this: it has one `Int64` - // variant for every integer width. Falls back to the resolved type's own - // name when the definition has no separate type token to quote. - let declared_type = col_def_str - .split_whitespace() - .nth(1) - .map(str::to_string) + // Falls back to the resolved type's own name when the definition has no + // separate type text to quote. + let declared_type = written_type + .map(|(_, declared)| declared.to_string()) .unwrap_or_else(|| column.column_type.to_string()); // Validate: new column must be nullable or have a default. @@ -90,11 +101,9 @@ pub(super) async fn alter_table_add_column( let mut updated = coll; updated.collection_type = nodedb_types::CollectionType::strict(schema.clone()); updated.timeseries_config = sonic_rs::to_string(&schema).ok(); - // Record the column's *declared* type alongside the - // resolved one — see `strict_schema::retype_field`. Without - // this the added column has no declared width and falls - // back to `BIGINT` on the wire, unlike an identical column - // declared at CREATE time. + // Record the column's *declared* type alongside the schema + // column — see `strict_schema::retype_field`. Catalog + // introspection reads the width from this spelling. super::strict_schema::add_field( &mut updated, &column_name, diff --git a/nodedb/src/control/server/shared/ddl/neutral/collection/create/build.rs b/nodedb/src/control/server/shared/ddl/neutral/collection/create/build.rs index 0f6278dec..f79b9b5d3 100644 --- a/nodedb/src/control/server/shared/ddl/neutral/collection/create/build.rs +++ b/nodedb/src/control/server/shared/ddl/neutral/collection/create/build.rs @@ -41,6 +41,7 @@ use super::build_flags::{ use super::build_post_create::{create_serial_sequences, log_vector_fields}; use super::build_primary_engine::resolve_primary_engine; use crate::control::server::shared::ddl::neutral::column_default::validate_column_defaults; +use crate::control::server::shared::ddl::neutral::declared_typmod::validate_declared_typmods; /// Per-surface configuration. The fields are the entire surface-level /// difference between `CREATE COLLECTION` and `CREATE TABLE`. @@ -102,6 +103,10 @@ pub async fn build_and_persist( )); } + // Refuse an invalid `DECIMAL(p,s)` typmod on every engine. It runs + // before the DEFAULT gate, which reads the declared type to check a + // literal against it. + validate_declared_typmods(columns)?; // Refuse a DEFAULT the server cannot evaluate, or that the declared // column type cannot hold, here, not at the first INSERT. It runs before // any lifecycle guard or predecessor purge, so a rejected declaration diff --git a/nodedb/src/control/server/shared/ddl/neutral/convert/column_defs.rs b/nodedb/src/control/server/shared/ddl/neutral/convert/column_defs.rs index 4e571907f..e53108057 100644 --- a/nodedb/src/control/server/shared/ddl/neutral/convert/column_defs.rs +++ b/nodedb/src/control/server/shared/ddl/neutral/convert/column_defs.rs @@ -160,7 +160,8 @@ fn parse_column_defs(s: &str) -> Result, ColumnDef::nullable(col_name, ct) } else { ColumnDef::required(col_name, ct) - }; + } + .with_declared_width(&col_type); if primary_key { col = col.with_primary_key(); } @@ -258,10 +259,9 @@ mod tests { ); assert_eq!( cols[1].column_type, - ColumnType::Decimal { - precision: 10, - scale: 2 - } + ColumnType::Decimal(Some( + nodedb_types::columnar::DecimalTypmod::new(10, 2).expect("valid typmod") + )) ); } } diff --git a/nodedb/src/control/server/shared/ddl/neutral/convert/type_map.rs b/nodedb/src/control/server/shared/ddl/neutral/convert/type_map.rs index f9c3eaf5d..f7d620c30 100644 --- a/nodedb/src/control/server/shared/ddl/neutral/convert/type_map.rs +++ b/nodedb/src/control/server/shared/ddl/neutral/convert/type_map.rs @@ -32,9 +32,9 @@ fn resolve_declared_type( /// `BINARY` resolves ahead of that call. It is a CONVERT-only spelling for /// `ColumnType::Bytes` that the shared parser rejects. /// -/// A spelling the shared parser rejects raises `42601`. The author names the -/// stored column type, so an unresolvable spelling must never become -/// `ColumnType::String`. +/// A spelling the shared parser rejects raises `42601`, and an out-of-range +/// `DECIMAL(p,s)` typmod raises `22023`. The author names the stored column +/// type, so an unresolvable spelling must never become `ColumnType::String`. pub(super) fn sql_type_to_column_type( sql_type: &str, ) -> Result { @@ -45,7 +45,7 @@ pub(super) fn sql_type_to_column_type( return Ok(ColumnType::Bytes); } resolve_declared_type(declared) - .map_err(|e| err("42601", format!("column type '{declared}': {e}"))) + .map_err(|e| err(e.sqlstate(), format!("column type '{declared}': {e}"))) } /// Map a typeguard type expression string to a `ColumnType`. @@ -285,10 +285,7 @@ mod tests { SimpleType::Timestamp => ColumnType::Timestamp, SimpleType::Timestamptz => ColumnType::Timestamptz, SimpleType::SystemTimestamp => ColumnType::SystemTimestamp, - SimpleType::Decimal { precision, scale } => ColumnType::Decimal { - precision: *precision, - scale: *scale, - }, + SimpleType::Decimal(typmod) => ColumnType::Decimal(*typmod), SimpleType::Uuid => ColumnType::Uuid, SimpleType::Ulid => ColumnType::Ulid, SimpleType::Geometry => ColumnType::Geometry, diff --git a/nodedb/src/control/server/shared/ddl/neutral/declared_typmod.rs b/nodedb/src/control/server/shared/ddl/neutral/declared_typmod.rs new file mode 100644 index 000000000..f7af05bbc --- /dev/null +++ b/nodedb/src/control/server/shared/ddl/neutral/declared_typmod.rs @@ -0,0 +1,101 @@ +// SPDX-License-Identifier: BUSL-1.1 + +//! DDL-time gate for a declared `DECIMAL(p,s)` typmod. +//! +//! A schemaless or key-value collection keeps each declared column type as +//! text. Without this gate, `DECIMAL(1001,0)` is accepted there and the +//! column then reads as an unknown type. The gate runs on every engine's +//! `CREATE` path and on `ALTER ADD COLUMN`, so all engines refuse the same +//! declarations. +//! +//! An out-of-range precision or scale is refused with SQLSTATE `22023`, as +//! PostgreSQL refuses it. A malformed parameter list is refused with `42601`. + +use nodedb_types::columnar::{ColumnType, ColumnTypeParseError}; + +use super::super::result::DdlError; + +/// Refuse every column whose declared type is a `DECIMAL` or `NUMERIC` with +/// an invalid typmod. +/// +/// Each pair carries a column name and its declared type text, modifiers +/// included. +pub(super) fn validate_declared_typmods(columns: &[(String, String)]) -> Result<(), DdlError> { + for (column, declared_type) in columns { + validate_declared_typmod(column, declared_type)?; + } + Ok(()) +} + +/// Refuse one column whose declared type is a `DECIMAL` or `NUMERIC` with an +/// invalid typmod. +/// +/// Every other parse error is left to the engine's own schema builder. A +/// schemaless column can name a custom type this parser does not know. +pub(super) fn validate_declared_typmod(column: &str, declared_type: &str) -> Result<(), DdlError> { + let Err(error) = ColumnType::parse_declared_type(declared_type) else { + return Ok(()); + }; + if matches!( + error, + ColumnTypeParseError::InvalidDecimalTypmod(_) + | ColumnTypeParseError::InvalidDecimalParams(_) + ) { + return Err(DdlError::new( + error.sqlstate(), + format!("column '{column}': {error}"), + )); + } + Ok(()) +} + +#[cfg(test)] +mod tests { + use super::*; + + fn declared(name: &str, declared_type: &str) -> (String, String) { + (name.to_string(), declared_type.to_string()) + } + + #[test] + fn valid_and_plain_decimals_pass() { + let columns = [ + declared("a", "DECIMAL(5,2)"), + declared("b", "NUMERIC(28, 28) NOT NULL"), + declared("c", "DECIMAL"), + declared("d", "NUMERIC DEFAULT 1.5"), + declared("e", "TEXT"), + declared("f", "my_custom_type"), + ]; + assert!(validate_declared_typmods(&columns).is_ok()); + } + + #[test] + fn out_of_range_typmod_is_22023() { + for declared_type in [ + "DECIMAL(1001,0)", + "DECIMAL(0)", + "NUMERIC(5,6)", + "DECIMAL(29,0) NOT NULL", + "DECIMAL(38, 10)", + ] { + let error = validate_declared_typmods(&[declared("v", declared_type)]) + .expect_err(declared_type); + assert_eq!(error.sqlstate, "22023", "{declared_type}"); + assert!( + error.message.contains("'v'"), + "{declared_type}: {}", + error.message + ); + } + } + + #[test] + fn malformed_typmod_is_42601() { + for declared_type in ["DECIMAL(a,2)", "NUMERIC(5,2,1)", "DECIMAL()"] { + let error = validate_declared_typmods(&[declared("v", declared_type)]) + .expect_err(declared_type); + assert_eq!(error.sqlstate, "42601", "{declared_type}"); + } + } +} diff --git a/nodedb/src/control/server/shared/ddl/neutral/mod.rs b/nodedb/src/control/server/shared/ddl/neutral/mod.rs index eae9717e6..a7094a850 100644 --- a/nodedb/src/control/server/shared/ddl/neutral/mod.rs +++ b/nodedb/src/control/server/shared/ddl/neutral/mod.rs @@ -27,6 +27,7 @@ pub mod convert; pub mod crdt_ops; pub mod custom_type; pub mod database; +mod declared_typmod; pub mod deferred_effects; pub mod dsl; pub mod emergency_ddl; diff --git a/nodedb/tests/wire/cases/sql_decimal_typmod.rs b/nodedb/tests/wire/cases/sql_decimal_typmod.rs new file mode 100644 index 000000000..83c87778c --- /dev/null +++ b/nodedb/tests/wire/cases/sql_decimal_typmod.rs @@ -0,0 +1,306 @@ +// SPDX-License-Identifier: BUSL-1.1 + +//! A `DECIMAL(p,s)` column enforces its declared precision and scale on every +//! engine, as PostgreSQL does. A value rounds to `s` fractional digits, half +//! away from zero. A value whose rounded integer part has more than `p - s` +//! digits is refused with SQLSTATE 22003. A plain `DECIMAL` keeps every digit +//! it is given. A typmod outside precision `1..=1000` and scale +//! `0..=precision` is refused at DDL with SQLSTATE 22023. + +use crate::harness::TestServer; + +/// One engine under test: the collection name, its key column, and the +/// `CREATE` statement with a `v` column of the given declared type. +struct Engine { + name: String, + key: &'static str, + create: String, +} + +/// Every engine that accepts a declared column, each with a `v` column of +/// `declared` type. +fn engines(suffix: &str, declared: &str) -> Vec { + let engine = |kind: &str, key: &'static str, columns: &str, engine_name: &str| { + let name = format!("dec_{suffix}_{kind}"); + Engine { + create: format!("CREATE COLLECTION {name} {columns} WITH (engine='{engine_name}')"), + name, + key, + } + }; + vec![ + engine( + "strict", + "id", + &format!("(id TEXT PRIMARY KEY, v {declared})"), + "document_strict", + ), + engine( + "columnar", + "id", + &format!("COLUMNS (id TEXT, v {declared})"), + "columnar", + ), + engine( + "schemaless", + "id", + &format!("(id STRING PRIMARY KEY, v {declared})"), + "document_schemaless", + ), + engine( + "kv", + "key", + &format!("(key STRING PRIMARY KEY, v {declared})"), + "kv", + ), + ] +} + +/// The `v` value stored under key `'a'`. +async fn read_v(srv: &TestServer, engine: &Engine) -> Vec> { + let Engine { name, key, .. } = engine; + srv.query_rows(&format!("SELECT v FROM {name} WHERE {key} = 'a'")) + .await + .unwrap() +} + +/// A value with more integer digits than `DECIMAL(5,2)` holds is refused as +/// numeric value out of range, on INSERT, UPDATE and UPSERT. +#[tokio::test] +async fn decimal_past_declared_precision_is_22003() { + let srv = TestServer::start().await; + for engine in engines("range", "DECIMAL(5,2)") { + let Engine { name, key, create } = &engine; + srv.exec(create).await.unwrap(); + srv.expect_error( + &format!("INSERT INTO {name} ({key}, v) VALUES ('a', 123456.789)"), + "SQLSTATE 22003", + ) + .await; + srv.exec(&format!("INSERT INTO {name} ({key}, v) VALUES ('a', 1.5)")) + .await + .unwrap(); + srv.expect_error( + &format!("UPDATE {name} SET v = 123456.789 WHERE {key} = 'a'"), + "SQLSTATE 22003", + ) + .await; + srv.expect_error( + &format!("UPSERT INTO {name} ({key}, v) VALUES ('a', 123456.789)"), + "SQLSTATE 22003", + ) + .await; + assert_eq!( + read_v(&srv, &engine).await, + vec![vec!["1.50".to_string()]], + "{name}" + ); + } +} + +/// A value with more fractional digits than the declared scale rounds half +/// away from zero, and reads back at the declared scale. +#[tokio::test] +async fn decimal_rounds_to_declared_scale() { + let srv = TestServer::start().await; + for engine in engines("round", "DECIMAL(5,2)") { + let Engine { name, key, create } = &engine; + srv.exec(create).await.unwrap(); + srv.exec(&format!( + "INSERT INTO {name} ({key}, v) VALUES ('a', 1.005)" + )) + .await + .unwrap(); + assert_eq!( + read_v(&srv, &engine).await, + vec![vec!["1.01".to_string()]], + "{name}" + ); + srv.exec(&format!("UPDATE {name} SET v = 2.675 WHERE {key} = 'a'")) + .await + .unwrap(); + assert_eq!( + read_v(&srv, &engine).await, + vec![vec!["2.68".to_string()]], + "{name}" + ); + } +} + +/// A plain `DECIMAL` has no digit limit and keeps every digit it is given. +#[tokio::test] +async fn plain_decimal_keeps_every_digit() { + let srv = TestServer::start().await; + let wide = "1234567890123456789.123456789"; + for engine in engines("plain", "DECIMAL") { + let Engine { name, key, create } = &engine; + srv.exec(create).await.unwrap(); + srv.exec(&format!( + "INSERT INTO {name} ({key}, v) VALUES ('a', {wide})" + )) + .await + .unwrap(); + assert_eq!( + read_v(&srv, &engine).await, + vec![vec![wide.to_string()]], + "{name}" + ); + } +} + +/// The engines that store a row as a field map, so the Data Plane re-types a +/// computed value to the declared column: schemaless document and KV. +fn map_engines(suffix: &str, declared: &str) -> Vec { + engines(suffix, declared) + .into_iter() + .filter(|engine| engine.name.ends_with("_schemaless") || engine.name.ends_with("_kv")) + .collect() +} + +/// A computed assignment to `v` for `engine`'s row `'a'`. Schemaless runs it +/// as an `UPDATE`. KV refuses a computed `UPDATE` at plan time, so it runs as +/// the conflict branch of an `INSERT ... ON CONFLICT DO UPDATE`. +fn computed_set(engine: &Engine, expr: &str) -> String { + let Engine { name, key, .. } = engine; + if *key == "key" { + format!( + "INSERT INTO {name} ({key}, v) VALUES ('a', 0) \ + ON CONFLICT ({key}) DO UPDATE SET v = {expr}" + ) + } else { + format!("UPDATE {name} SET v = {expr} WHERE {key} = 'a'") + } +} + +/// A computed value past `DECIMAL(5,2)` is refused with 22003, on `UPDATE` +/// and on the conflict branch of an upsert, and the row keeps its value. +#[tokio::test] +async fn computed_decimal_past_declared_precision_is_22003() { + let srv = TestServer::start().await; + for engine in map_engines("computed_range", "DECIMAL(5,2)") { + let Engine { name, key, create } = &engine; + srv.exec(create).await.unwrap(); + srv.exec(&format!("INSERT INTO {name} ({key}, v) VALUES ('a', 1.5)")) + .await + .unwrap(); + srv.expect_error(&computed_set(&engine, "v * 1000"), "SQLSTATE 22003") + .await; + srv.expect_error( + &format!( + "INSERT INTO {name} ({key}, v) VALUES ('a', 0) \ + ON CONFLICT ({key}) DO UPDATE SET v = v * 1000" + ), + "SQLSTATE 22003", + ) + .await; + assert_eq!( + read_v(&srv, &engine).await, + vec![vec!["1.50".to_string()]], + "{name}" + ); + } +} + +/// A computed value with more fractional digits than the declared scale +/// rounds to it. +#[tokio::test] +async fn computed_decimal_rounds_to_declared_scale() { + let srv = TestServer::start().await; + for engine in map_engines("computed_round", "DECIMAL(5,2)") { + let Engine { name, key, create } = &engine; + srv.exec(create).await.unwrap(); + srv.exec(&format!("INSERT INTO {name} ({key}, v) VALUES ('a', 1)")) + .await + .unwrap(); + srv.exec(&computed_set(&engine, "v / 3")).await.unwrap(); + assert_eq!( + read_v(&srv, &engine).await, + vec![vec!["0.33".to_string()]], + "{name}" + ); + } +} + +/// `INSERT ... SELECT` copies a value with more fractional digits than the +/// target's declared scale, and the target stores it rounded. KV refuses an +/// `INSERT ... SELECT` target, so this runs on schemaless. +#[tokio::test] +async fn insert_select_rounds_to_declared_scale() { + let srv = TestServer::start().await; + srv.exec( + "CREATE COLLECTION dec_copy_src (id STRING PRIMARY KEY, v DECIMAL) \ + WITH (engine='document_schemaless')", + ) + .await + .unwrap(); + srv.exec("INSERT INTO dec_copy_src (id, v) VALUES ('a', 1.005)") + .await + .unwrap(); + for engine in map_engines("copy", "DECIMAL(5,2)") + .into_iter() + .filter(|engine| engine.key == "id") + { + let Engine { name, create, .. } = &engine; + srv.exec(create).await.unwrap(); + srv.exec(&format!("INSERT INTO {name} SELECT * FROM dec_copy_src")) + .await + .unwrap(); + assert_eq!( + read_v(&srv, &engine).await, + vec![vec!["1.01".to_string()]], + "{name}" + ); + } + srv.exec("INSERT INTO dec_copy_src (id, v) VALUES ('b', 123456.789)") + .await + .unwrap(); + srv.exec( + "CREATE COLLECTION dec_copy_narrow (id STRING PRIMARY KEY, v DECIMAL(5,2)) \ + WITH (engine='document_schemaless')", + ) + .await + .unwrap(); + srv.expect_error( + "INSERT INTO dec_copy_narrow SELECT * FROM dec_copy_src", + "SQLSTATE 22003", + ) + .await; +} + +/// A computed value past a declared `INT2` is refused with 22003, and the row +/// keeps its value. +#[tokio::test] +async fn computed_value_past_int2_is_22003() { + let srv = TestServer::start().await; + for engine in map_engines("computed_int2", "INT2") { + let Engine { name, key, create } = &engine; + srv.exec(create).await.unwrap(); + srv.exec(&format!("INSERT INTO {name} ({key}, v) VALUES ('a', 1)")) + .await + .unwrap(); + srv.expect_error(&computed_set(&engine, "v + 39999"), "SQLSTATE 22003") + .await; + assert_eq!( + read_v(&srv, &engine).await, + vec![vec!["1".to_string()]], + "{name}" + ); + } +} + +/// A typmod past PostgreSQL's range, or past the 28 digits the engine stores +/// exactly, is refused at DDL on every engine. +#[tokio::test] +async fn invalid_typmod_is_refused_at_ddl() { + let srv = TestServer::start().await; + for (suffix, declared) in [ + ("p1001", "DECIMAL(1001,0)"), + ("p0", "NUMERIC(0)"), + ("s6", "DECIMAL(5,6)"), + ("p29", "DECIMAL(29,2)"), + ] { + for engine in engines(suffix, declared) { + srv.expect_error(&engine.create, "SQLSTATE 22023").await; + } + } +} From 4e2f537c48edf40809f6a155240931e45e791274 Mon Sep 17 00:00:00 2001 From: Farhan Syah Date: Thu, 8 Oct 2026 16:51:29 +0800 Subject: [PATCH 06/24] feat(schema): enforce declared column types on every write path Column definitions record the declared integer width (SMALLINT, INT) and float width (REAL). Strict, columnar, schemaless and KV writes re-type values under declared numeric columns and refuse one past the declared width, including computed assignments, counters, upserts, staged transaction writes and WAL replay. A narrowing ALTER COLUMN TYPE is refused. Declared VECTOR(n) columns of a schemaless collection are indexed like strict vector columns. The strict tuple encoder reports a value of the wrong kind and text that does not convert as distinct errors. --- .../physical_plan/document/declared_column.rs | 130 +++ .../src/physical_plan/document/mod.rs | 2 + .../src/physical_plan/document/op.rs | 15 + nodedb-physical/src/physical_plan/mod.rs | 15 +- nodedb-sql/src/ddl_ast/collection_type/kv.rs | 3 +- .../src/ddl_ast/collection_type/strict.rs | 41 +- .../src/planner/dml_helpers/range_check.rs | 7 +- nodedb-strict/src/decode.rs | 14 +- nodedb-strict/src/encode.rs | 341 +++----- nodedb-strict/src/encode/value.rs | 241 ++++++ nodedb-strict/src/error.rs | 9 + nodedb-types/src/columnar/column_def.rs | 82 +- .../planner/catalog_adapter/type_convert.rs | 220 +++--- .../sql_plan_convert/dml/insert/schema.rs | 60 +- .../pgwire/handler/shape_encode/array_text.rs | 310 ++++++++ .../neutral/collection/alter/alter_type.rs | 64 +- .../neutral/collection/alter/strict_schema.rs | 11 +- .../collection/dml/indexed_vector_fields.rs | 24 + .../ddl/neutral/collection/dml/insert.rs | 5 +- .../shared/ddl/neutral/collection/helpers.rs | 3 +- .../shared/ddl/neutral/collection/register.rs | 90 +++ .../ddl/neutral/convert/typeguard_columns.rs | 10 +- .../executor/core_loop/declared_columns.rs | 40 + nodedb/src/data/executor/core_loop/mod.rs | 2 + nodedb/src/data/executor/dispatch/document.rs | 26 +- .../executor/dispatch/document_register.rs | 62 ++ nodedb/src/data/executor/dispatch/kv.rs | 11 +- nodedb/src/data/executor/dispatch/mod.rs | 1 + .../handlers/columnar_resolved_mutation.rs | 175 +++- .../executor/handlers/columnar_write/mod.rs | 10 +- .../handlers/columnar_write/schema.rs | 314 +++++--- .../handlers/document/write/register.rs | 65 +- .../executor/handlers/kv/conflict_merge.rs | 103 ++- .../executor/handlers/kv/crud/write_upsert.rs | 20 +- .../executor/handlers/kv/declared_body.rs | 277 +++++++ .../executor/handlers/kv/field_compute.rs | 34 +- nodedb/src/data/executor/handlers/kv/mod.rs | 1 + .../handlers/kv/resolve/atomic_ops.rs | 49 +- .../handlers/kv/resolve/predicate_ops.rs | 4 +- .../handlers/kv/resolve/transfer_ops.rs | 22 +- .../executor/handlers/kv/resolve/write_ops.rs | 11 +- .../handlers/point/apply_put/stored_body.rs | 20 +- .../handlers/point/apply_put/vector/fields.rs | 50 +- .../handlers/point/apply_put/vector/mod.rs | 5 +- .../handlers/point/update/post_image.rs | 69 +- .../handlers/transaction/stage_write/body.rs | 13 + .../stage_columnar_resolved_dml.rs | 16 +- .../transaction/stage_write/stage_kv.rs | 7 + .../stage_write/stage_kv_atomic.rs | 35 +- .../stage_write/stage_kv_conflict.rs | 19 +- .../stage_write/stage_kv_predicate.rs | 8 +- .../stage_write/stage_kv_transfer.rs | 46 +- .../transaction/stage_write/stage_upsert.rs | 13 +- .../src/data/executor/strict_format/coerce.rs | 744 ++++++++++++++---- .../data/executor/strict_format/declared.rs | 259 ++++++ nodedb/src/data/executor/strict_format/mod.rs | 7 + nodedb/src/data/executor/wal_replay/kv_put.rs | 38 + .../src/data/executor/wal_replay_kv_atomic.rs | 61 +- .../src/data/executor/wal_replay_kv_field.rs | 10 +- .../src/data/executor/wal_replay_kv_incr.rs | 29 +- .../executor/wal_replay_kv_insert_conflict.rs | 50 +- .../data/executor/wal_replay_kv_predicate.rs | 7 +- nodedb/src/engine/document/store/config.rs | 17 + nodedb/src/engine/kv/counter_refit.rs | 91 +++ nodedb/src/engine/kv/engine_atomic.rs | 157 +++- nodedb/src/engine/kv/mod.rs | 4 +- .../test_conflict_policy_register.rs | 3 + .../executor_tests/test_generated_columns.rs | 3 + .../test_range_scan_bitemporal.rs | 6 + .../cases/declared_width_computed_writes.rs | 312 ++++++++ .../sql_transactions_kv_atomic_overlay.rs | 11 +- 71 files changed, 4113 insertions(+), 921 deletions(-) create mode 100644 nodedb-physical/src/physical_plan/document/declared_column.rs create mode 100644 nodedb-strict/src/encode/value.rs create mode 100644 nodedb/src/control/server/pgwire/handler/shape_encode/array_text.rs create mode 100644 nodedb/src/data/executor/core_loop/declared_columns.rs create mode 100644 nodedb/src/data/executor/dispatch/document_register.rs create mode 100644 nodedb/src/data/executor/handlers/kv/declared_body.rs create mode 100644 nodedb/src/data/executor/strict_format/declared.rs create mode 100644 nodedb/src/engine/kv/counter_refit.rs create mode 100644 nodedb/tests/wire/cases/declared_width_computed_writes.rs 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/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-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/planner/dml_helpers/range_check.rs b/nodedb-sql/src/planner/dml_helpers/range_check.rs index 1a8967983..3ec6bfd6d 100644 --- a/nodedb-sql/src/planner/dml_helpers/range_check.rs +++ b/nodedb-sql/src/planner/dml_helpers/range_check.rs @@ -157,10 +157,9 @@ pub(crate) fn check_declared_float_ranges( /// [`check_declared_int_ranges`] for `UPDATE ... SET col = `. /// /// Only literal assignments are checkable at plan time; a computed assignment -/// (`SET n = n + 1`) has no value until the Data Plane evaluates it. Those are -/// caught on the read path instead, where the encoder refuses to transmit a -/// value that does not fit the column's advertised width — so an out-of-range -/// value can never reach a client silently by either route. +/// (`SET n = n + 1`) has no value until the Data Plane evaluates it. The Data +/// Plane checks the evaluated value against the same declared width before it +/// stores the row, so a computed value out of range is refused there. pub(crate) fn check_declared_int_ranges_in_assignments( columns: &[ColumnInfo], assignments: &[(String, SqlExpr)], diff --git a/nodedb-strict/src/decode.rs b/nodedb-strict/src/decode.rs index d8c99a7d6..8d26cd20a 100644 --- a/nodedb-strict/src/decode.rs +++ b/nodedb-strict/src/decode.rs @@ -480,10 +480,9 @@ mod tests { ColumnDef::nullable("email", ColumnType::String), ColumnDef::required( "balance", - ColumnType::Decimal { - precision: 18, - scale: 4, - }, + ColumnType::Decimal(Some( + nodedb_types::columnar::DecimalTypmod::new(18, 4).expect("valid typmod"), + )), ), ColumnDef::nullable("active", ColumnType::Bool), ]) @@ -766,10 +765,9 @@ mod tests { ColumnDef::required("tstz", ColumnType::Timestamptz), ColumnDef::required( "dec", - ColumnType::Decimal { - precision: 18, - scale: 4, - }, + ColumnType::Decimal(Some( + nodedb_types::columnar::DecimalTypmod::new(18, 4).expect("valid typmod"), + )), ), ColumnDef::required("uid", ColumnType::Uuid), ColumnDef::required("vec", ColumnType::Vector(2)), diff --git a/nodedb-strict/src/encode.rs b/nodedb-strict/src/encode.rs index b9a797dd7..4cb249a85 100644 --- a/nodedb-strict/src/encode.rs +++ b/nodedb-strict/src/encode.rs @@ -19,11 +19,15 @@ pub const MAGIC: u32 = 0x5453_444E; /// Current Binary Tuple format version. pub const FORMAT_VERSION: u8 = 1; -use nodedb_types::columnar::{ColumnType, StrictSchema}; +use nodedb_types::columnar::StrictSchema; use nodedb_types::value::Value; use crate::error::StrictError; +#[path = "encode/value.rs"] +mod value_encode; +use value_encode::{encode_fixed, encode_variable}; + /// Encodes rows into Binary Tuples according to a fixed schema. /// /// Reusable: create once per schema, encode many rows. Internal buffers @@ -121,12 +125,7 @@ impl TupleEncoder { // Write fixed-size value. if let Some(offset) = self.fixed_offsets[i] { let dst = fixed_start + offset; - encode_fixed(&mut buf[dst..], &col.column_type, val).map_err(|()| { - StrictError::TypeMismatch { - column: col.name.clone(), - expected: col.column_type, - } - })?; + encode_fixed(&mut buf[dst..], col, val)?; } // Variable-length values are handled in the offset table pass below. } @@ -143,15 +142,7 @@ impl TupleEncoder { let val = &values[col_idx]; if !matches!(val, Value::Null) { - encode_variable( - &mut var_data, - &self.schema.columns[col_idx].column_type, - val, - ) - .map_err(|()| StrictError::TypeMismatch { - column: self.schema.columns[col_idx].name.clone(), - expected: self.schema.columns[col_idx].column_type, - })?; + encode_variable(&mut var_data, &self.schema.columns[col_idx], val)?; } // If null: offset stays the same as next entry → zero length. } @@ -207,211 +198,130 @@ impl TupleEncoder { } } -/// Encode a fixed-size value into the buffer at the given position. -/// -/// Handles both native Value types and SQL coercion sources. -fn encode_fixed(dst: &mut [u8], col_type: &ColumnType, value: &Value) -> Result<(), ()> { - match (col_type, value) { - // Int64: native. - (ColumnType::Int64, Value::Integer(v)) => { - dst[..8].copy_from_slice(&v.to_le_bytes()); - } - // Float64: native + Int64→Float64 coercion. - (ColumnType::Float64, Value::Float(v)) => { - dst[..8].copy_from_slice(&v.to_le_bytes()); - } - (ColumnType::Float64, Value::Integer(v)) => { - dst[..8].copy_from_slice(&(*v as f64).to_le_bytes()); - } - // Bool: native. - (ColumnType::Bool, Value::Bool(v)) => { - dst[0] = *v as u8; - } - // Timestamp (naive): NaiveDateTime + Integer (micros) + String (ISO 8601 parse). - (ColumnType::Timestamp, Value::NaiveDateTime(dt)) => { - dst[..8].copy_from_slice(&dt.micros.to_le_bytes()); - } - (ColumnType::Timestamp, Value::Integer(micros)) => { - dst[..8].copy_from_slice(µs.to_le_bytes()); - } - (ColumnType::Timestamp, Value::String(s)) => { - let micros = nodedb_types::NdbDateTime::parse(s) - .map(|dt| dt.micros) - .unwrap_or(0); - dst[..8].copy_from_slice(µs.to_le_bytes()); - } - // Timestamptz (TZ-aware): DateTime + Integer (micros) + String (ISO 8601 parse). - (ColumnType::Timestamptz, Value::DateTime(dt)) => { - dst[..8].copy_from_slice(&dt.micros.to_le_bytes()); - } - (ColumnType::Timestamptz, Value::Integer(micros)) => { - dst[..8].copy_from_slice(µs.to_le_bytes()); - } - (ColumnType::Timestamptz, Value::String(s)) => { - let micros = nodedb_types::NdbDateTime::parse(s) - .map(|dt| dt.micros) - .unwrap_or(0); - dst[..8].copy_from_slice(µs.to_le_bytes()); - } - // System timestamp: UTC micros, decoded canonically as DateTime. - (ColumnType::SystemTimestamp, Value::DateTime(dt)) => { - dst[..8].copy_from_slice(&dt.micros.to_le_bytes()); - } - (ColumnType::SystemTimestamp, Value::Integer(micros)) => { - dst[..8].copy_from_slice(µs.to_le_bytes()); - } - // ULID: parse its textual representation and retain the canonical 16-byte ID. - (ColumnType::Ulid, Value::Ulid(s) | Value::String(s)) => { - let id = ulid::Ulid::from_string(s).map_err(|_| ())?; - dst[..16].copy_from_slice(&id.to_bytes()); - } - // Duration: native duration, microseconds, or a human-readable literal. - (ColumnType::Duration, Value::Duration(duration)) => { - dst[..8].copy_from_slice(&duration.micros.to_le_bytes()); - } - (ColumnType::Duration, Value::Integer(micros)) => { - dst[..8].copy_from_slice(µs.to_le_bytes()); - } - (ColumnType::Duration, Value::String(s)) => { - let duration = nodedb_types::NdbDuration::parse(s).ok_or(())?; - dst[..8].copy_from_slice(&duration.micros.to_le_bytes()); - } - // Decimal: native Decimal + String/Float/Integer coercion. - (ColumnType::Decimal { .. }, Value::Decimal(d)) => { - dst[..16].copy_from_slice(&d.serialize()); - } - (ColumnType::Decimal { .. }, Value::String(s)) => { - let d: rust_decimal::Decimal = s.parse().unwrap_or_default(); - dst[..16].copy_from_slice(&d.serialize()); - } - (ColumnType::Decimal { .. }, Value::Float(f)) => { - let d = rust_decimal::Decimal::try_from(*f).unwrap_or_default(); - dst[..16].copy_from_slice(&d.serialize()); - } - (ColumnType::Decimal { .. }, Value::Integer(i)) => { - let d = rust_decimal::Decimal::from(*i); - dst[..16].copy_from_slice(&d.serialize()); - } - // Uuid: native Uuid string + String coercion. - (ColumnType::Uuid, Value::Uuid(s) | Value::String(s)) => { - if let Ok(parsed) = uuid::Uuid::parse_str(s) { - dst[..16].copy_from_slice(parsed.as_bytes()); - } - } - // Vector: Array of floats + Bytes (packed f32). - (ColumnType::Vector(dim), Value::Array(arr)) => { - let d = *dim as usize; - for (i, v) in arr.iter().take(d).enumerate() { - let f = match v { - Value::Float(f) => *f as f32, - Value::Integer(n) => *n as f32, - _ => 0.0, - }; - dst[i * 4..(i + 1) * 4].copy_from_slice(&f.to_le_bytes()); - } - } - (ColumnType::Vector(dim), Value::Bytes(b)) => { - let byte_len = (*dim as usize) * 4; - let copy_len = b.len().min(byte_len); - dst[..copy_len].copy_from_slice(&b[..copy_len]); - } - _ => {} // Type mismatch caught earlier by accepts(). +#[cfg(test)] +mod tests { + use nodedb_types::columnar::{ColumnDef, ColumnType}; + use nodedb_types::datetime::NdbDateTime; + + use super::*; + + /// Encode one value into a single-column schema of `column_type`. + fn encode_one(column_type: ColumnType, value: Value) -> Result, StrictError> { + let schema = StrictSchema::new(vec![ColumnDef::required("c", column_type)]).unwrap(); + TupleEncoder::new(&schema).encode(&[value]) } - Ok(()) -} -/// Encode a variable-length value, appending to the data buffer. -/// -/// Handles both native Value types and SQL coercion sources. -fn encode_variable(var_data: &mut Vec, col_type: &ColumnType, value: &Value) -> Result<(), ()> { - match (col_type, value) { - (ColumnType::String, Value::String(s)) => { - var_data.extend_from_slice(s.as_bytes()); - } - (ColumnType::Bytes, Value::Bytes(b)) => { - var_data.extend_from_slice(b); - } - // Geometry: native Geometry (JSON-serialized) + String (WKT/GeoJSON passthrough). - (ColumnType::Geometry, Value::Geometry(g)) => { - if let Ok(json) = sonic_rs::to_vec(g) { - var_data.extend_from_slice(&json); + fn assert_invalid_value(result: Result, StrictError>, expected: ColumnType) { + match result { + Err(StrictError::InvalidValue { + column, + expected: got, + .. + }) => { + assert_eq!(column, "c"); + assert_eq!(got, expected); } + other => panic!("expected InvalidValue for {expected}, got {other:?}"), } - (ColumnType::Geometry, Value::String(s)) => { - var_data.extend_from_slice(s.as_bytes()); - } - (ColumnType::Json, Value::String(s)) => { - // String input for JSON column: parse as JSON, then serialize as MessagePack. - // This handles VALUES ('{"key":"val"}') where the SQL planner passes a string literal. - let parsed = sonic_rs::from_str::(s) - .ok() - .map(nodedb_types::Value::from); - let to_encode = parsed.as_ref().unwrap_or(value); - if let Ok(bytes) = nodedb_types::value_to_msgpack(to_encode) { - var_data.extend_from_slice(&bytes); - } - } - (ColumnType::Json, value) => { - // Non-string input (Object, Array, etc.): serialize directly as MessagePack. - if let Ok(bytes) = nodedb_types::value_to_msgpack(value) { - var_data.extend_from_slice(&bytes); - } - } - // SparseVector: a `'{id: weight}'` literal stored as raw UTF-8 bytes, - // mirroring the String path (parsed at index-build time). Raw bytes - // pass through unchanged. - (ColumnType::SparseVector, Value::String(s)) => { - var_data.extend_from_slice(s.as_bytes()); - } - (ColumnType::SparseVector, Value::Bytes(b)) => { - var_data.extend_from_slice(b); - } - // Typed variable columns use tagged NodeDB MessagePack so their Value - // variant survives storage. Coercion inputs are converted first. - (ColumnType::Array, Value::Array(_)) - | (ColumnType::Set, Value::Set(_)) - | (ColumnType::Regex, Value::Regex(_)) - | (ColumnType::Range, Value::Range { .. }) - | (ColumnType::Record, Value::Record { .. }) => { - append_msgpack(var_data, value)?; - } - (ColumnType::Set, Value::Array(items)) => { - append_msgpack(var_data, &Value::Set(items.clone()))?; - } - (ColumnType::Regex, Value::String(pattern)) => { - append_msgpack(var_data, &Value::Regex(pattern.clone()))?; - } - (ColumnType::Record, Value::String(reference)) => { - let (table, id) = reference.split_once(':').ok_or(())?; - if table.is_empty() || id.is_empty() { - return Err(()); - } - append_msgpack( - var_data, - &Value::Record { - table: table.to_owned(), - id: id.to_owned(), - }, - )?; + } + + #[test] + fn unparsable_uuid_text_errors() { + assert_invalid_value( + encode_one(ColumnType::Uuid, Value::String("not-a-uuid".into())), + ColumnType::Uuid, + ); + assert_invalid_value( + encode_one(ColumnType::Uuid, Value::Uuid("zzzz".into())), + ColumnType::Uuid, + ); + } + + #[test] + fn valid_uuid_text_roundtrips() { + let schema = StrictSchema::new(vec![ColumnDef::required("c", ColumnType::Uuid)]).unwrap(); + let text = "67e55044-10b1-426f-9247-bb680e5fe0c8"; + let tuple = TupleEncoder::new(&schema) + .encode(&[Value::String(text.into())]) + .unwrap(); + let decoded = crate::decode::TupleDecoder::new(&schema) + .extract_value(&tuple, 0) + .unwrap(); + assert_eq!(decoded, Value::Uuid(text.into())); + } + + #[test] + fn unparsable_ulid_and_duration_text_error() { + assert_invalid_value( + encode_one(ColumnType::Ulid, Value::String("not-a-ulid".into())), + ColumnType::Ulid, + ); + assert_invalid_value( + encode_one(ColumnType::Ulid, Value::Ulid("01ARZ3NDEK".into())), + ColumnType::Ulid, + ); + assert_invalid_value( + encode_one(ColumnType::Duration, Value::String("ten minutes".into())), + ColumnType::Duration, + ); + } + + #[test] + fn unparsable_timestamp_text_errors() { + for column_type in [ColumnType::Timestamp, ColumnType::Timestamptz] { + assert_invalid_value( + encode_one(column_type, Value::String("yesterday".into())), + column_type, + ); } - _ => {} // Type mismatch caught earlier by accepts(). } - Ok(()) -} -/// Append a lossless NodeDB MessagePack representation. -fn append_msgpack(var_data: &mut Vec, value: &Value) -> Result<(), ()> { - let bytes = zerompk::to_msgpack_vec(value).map_err(|_| ())?; - var_data.extend_from_slice(&bytes); - Ok(()) -} + #[test] + fn unconvertible_decimal_sources_error() { + let decimal = ColumnType::Decimal(None); + assert_invalid_value(encode_one(decimal, Value::String("12.x".into())), decimal); + assert_invalid_value(encode_one(decimal, Value::Float(f64::NAN)), decimal); + assert_invalid_value(encode_one(decimal, Value::Float(f64::INFINITY)), decimal); + } -#[cfg(test)] -mod tests { - use nodedb_types::columnar::ColumnDef; - use nodedb_types::datetime::NdbDateTime; + #[test] + fn vector_with_wrong_shape_errors() { + let vector = ColumnType::Vector(3); + let short = Value::Array(vec![Value::Float(1.0), Value::Float(2.0)]); + assert_invalid_value(encode_one(vector, short), vector); + let long = Value::Array(vec![Value::Float(0.5); 4]); + assert_invalid_value(encode_one(vector, long), vector); + let non_numeric = Value::Array(vec![ + Value::Float(1.0), + Value::String("two".into()), + Value::Float(3.0), + ]); + assert_invalid_value(encode_one(vector, non_numeric), vector); + assert_invalid_value(encode_one(vector, Value::Bytes(vec![0; 8])), vector); + assert_invalid_value(encode_one(vector, Value::Bytes(vec![0; 13])), vector); + assert!(encode_one(vector, Value::Bytes(vec![0; 12])).is_ok()); + } - use super::*; + #[test] + fn malformed_record_reference_errors() { + for text in ["users", ":42", "users:"] { + assert_invalid_value( + encode_one(ColumnType::Record, Value::String(text.into())), + ColumnType::Record, + ); + } + } + + #[test] + fn invalid_value_error_reaches_encode_bitemporal() { + let schema = + StrictSchema::new_bitemporal(vec![ColumnDef::required("id", ColumnType::Uuid)]) + .unwrap(); + let err = TupleEncoder::new(&schema) + .encode_bitemporal(0, 0, 0, &[Value::String("bad".into())]) + .unwrap_err(); + assert!(matches!(err, StrictError::InvalidValue { ref column, .. } if column == "id")); + } fn crm_schema() -> StrictSchema { StrictSchema::new(vec![ @@ -420,10 +330,9 @@ mod tests { ColumnDef::nullable("email", ColumnType::String), ColumnDef::required( "balance", - ColumnType::Decimal { - precision: 18, - scale: 4, - }, + ColumnType::Decimal(Some( + nodedb_types::columnar::DecimalTypmod::new(18, 4).expect("valid typmod"), + )), ), ColumnDef::nullable("active", ColumnType::Bool), ]) diff --git a/nodedb-strict/src/encode/value.rs b/nodedb-strict/src/encode/value.rs new file mode 100644 index 000000000..40543c30a --- /dev/null +++ b/nodedb-strict/src/encode/value.rs @@ -0,0 +1,241 @@ +// SPDX-License-Identifier: Apache-2.0 + +//! Per-column value encoding for the Binary Tuple encoder. +//! +//! Every value either encodes in full or returns an error. A coercion +//! source that does not convert to the column type is an error, never a +//! zeroed or empty field. + +use nodedb_types::columnar::{ColumnDef, ColumnType}; +use nodedb_types::value::Value; + +use crate::error::StrictError; + +/// Encode a fixed-size value into `dst`. +/// +/// Handles both native Value types and SQL coercion sources. +pub(super) fn encode_fixed( + dst: &mut [u8], + col: &ColumnDef, + value: &Value, +) -> Result<(), StrictError> { + match (&col.column_type, value) { + (ColumnType::Int64, Value::Integer(v)) => { + dst[..8].copy_from_slice(&v.to_le_bytes()); + } + (ColumnType::Float64, Value::Float(v)) => { + dst[..8].copy_from_slice(&v.to_le_bytes()); + } + (ColumnType::Float64, Value::Integer(v)) => { + dst[..8].copy_from_slice(&(*v as f64).to_le_bytes()); + } + (ColumnType::Bool, Value::Bool(v)) => { + dst[0] = u8::from(*v); + } + (ColumnType::Timestamp, Value::NaiveDateTime(dt)) + | (ColumnType::Timestamptz | ColumnType::SystemTimestamp, Value::DateTime(dt)) => { + dst[..8].copy_from_slice(&dt.micros.to_le_bytes()); + } + ( + ColumnType::Timestamp + | ColumnType::Timestamptz + | ColumnType::SystemTimestamp + | ColumnType::Duration, + Value::Integer(micros), + ) => { + dst[..8].copy_from_slice(µs.to_le_bytes()); + } + (ColumnType::Timestamp | ColumnType::Timestamptz, Value::String(s)) => { + let dt = nodedb_types::NdbDateTime::parse(s).ok_or_else(|| { + invalid_value( + col, + format!("'{s}' does not parse as an ISO 8601 timestamp"), + ) + })?; + dst[..8].copy_from_slice(&dt.micros.to_le_bytes()); + } + (ColumnType::Ulid, Value::Ulid(s) | Value::String(s)) => { + let id = ulid::Ulid::from_string(s) + .map_err(|e| invalid_value(col, format!("'{s}' does not parse as a ULID: {e}")))?; + dst[..16].copy_from_slice(&id.to_bytes()); + } + (ColumnType::Duration, Value::Duration(duration)) => { + dst[..8].copy_from_slice(&duration.micros.to_le_bytes()); + } + (ColumnType::Duration, Value::String(s)) => { + let duration = nodedb_types::NdbDuration::parse(s) + .ok_or_else(|| invalid_value(col, format!("'{s}' does not parse as a duration")))?; + dst[..8].copy_from_slice(&duration.micros.to_le_bytes()); + } + (ColumnType::Decimal { .. }, Value::Decimal(d)) => { + dst[..16].copy_from_slice(&d.serialize()); + } + (ColumnType::Decimal { .. }, Value::String(s)) => { + let d: rust_decimal::Decimal = s.parse().map_err(|e| { + invalid_value(col, format!("'{s}' does not parse as a decimal: {e}")) + })?; + dst[..16].copy_from_slice(&d.serialize()); + } + (ColumnType::Decimal { .. }, Value::Float(f)) => { + let d = rust_decimal::Decimal::try_from(*f).map_err(|e| { + invalid_value(col, format!("{f} has no decimal representation: {e}")) + })?; + dst[..16].copy_from_slice(&d.serialize()); + } + (ColumnType::Decimal { .. }, Value::Integer(i)) => { + dst[..16].copy_from_slice(&rust_decimal::Decimal::from(*i).serialize()); + } + (ColumnType::Uuid, Value::Uuid(s) | Value::String(s)) => { + let parsed = uuid::Uuid::parse_str(s) + .map_err(|e| invalid_value(col, format!("'{s}' does not parse as a UUID: {e}")))?; + dst[..16].copy_from_slice(parsed.as_bytes()); + } + (ColumnType::Vector(dim), Value::Array(arr)) => { + check_vector_len(col, *dim, arr.len())?; + for (i, v) in arr.iter().enumerate() { + let f = match v { + Value::Float(f) => *f as f32, + Value::Integer(n) => *n as f32, + other => { + return Err(invalid_value( + col, + format!("element {i} is {other:?}, not a number"), + )); + } + }; + dst[i * 4..(i + 1) * 4].copy_from_slice(&f.to_le_bytes()); + } + } + (ColumnType::Vector(dim), Value::Bytes(b)) => { + if b.len() % 4 != 0 { + return Err(invalid_value( + col, + format!("{} bytes is not a whole number of f32 values", b.len()), + )); + } + check_vector_len(col, *dim, b.len() / 4)?; + dst[..b.len()].copy_from_slice(b); + } + _ => return Err(type_mismatch(col)), + } + Ok(()) +} + +/// Encode a variable-length value, appending to `var_data`. +/// +/// Handles both native Value types and SQL coercion sources. +pub(super) fn encode_variable( + var_data: &mut Vec, + col: &ColumnDef, + value: &Value, +) -> Result<(), StrictError> { + match (&col.column_type, value) { + (ColumnType::String, Value::String(s)) + | (ColumnType::Geometry, Value::String(s)) + | (ColumnType::SparseVector, Value::String(s)) => { + // Geometry text is WKT or GeoJSON. Sparse vector text is a + // `'{id: weight}'` literal parsed at index-build time. + var_data.extend_from_slice(s.as_bytes()); + } + (ColumnType::Bytes | ColumnType::SparseVector, Value::Bytes(b)) => { + var_data.extend_from_slice(b); + } + (ColumnType::Geometry, Value::Geometry(g)) => { + let json = sonic_rs::to_vec(g).map_err(|e| { + invalid_value(col, format!("geometry does not serialize to GeoJSON: {e}")) + })?; + var_data.extend_from_slice(&json); + } + (ColumnType::Json, Value::String(s)) => { + // A JSON text stores as the value it spells. Any other text + // stores as a JSON string. + let parsed = sonic_rs::from_str::(s) + .ok() + .map(Value::from); + append_json(var_data, col, parsed.as_ref().unwrap_or(value))?; + } + (ColumnType::Json, value) => { + append_json(var_data, col, value)?; + } + // Typed variable columns use tagged NodeDB MessagePack so their Value + // variant survives storage. Coercion inputs are converted first. + (ColumnType::Array, Value::Array(_)) + | (ColumnType::Set, Value::Set(_)) + | (ColumnType::Regex, Value::Regex(_)) + | (ColumnType::Range, Value::Range { .. }) + | (ColumnType::Record, Value::Record { .. }) => { + append_msgpack(var_data, col, value)?; + } + (ColumnType::Set, Value::Array(items)) => { + append_msgpack(var_data, col, &Value::Set(items.clone()))?; + } + (ColumnType::Regex, Value::String(pattern)) => { + append_msgpack(var_data, col, &Value::Regex(pattern.clone()))?; + } + (ColumnType::Record, Value::String(reference)) => { + let (table, id) = reference + .split_once(':') + .filter(|(table, id)| !table.is_empty() && !id.is_empty()) + .ok_or_else(|| { + invalid_value(col, format!("'{reference}' is not a 'table:id' reference")) + })?; + append_msgpack( + var_data, + col, + &Value::Record { + table: table.to_owned(), + id: id.to_owned(), + }, + )?; + } + _ => return Err(type_mismatch(col)), + } + Ok(()) +} + +/// Append the untagged MessagePack form a JSON column stores. +fn append_json(var_data: &mut Vec, col: &ColumnDef, value: &Value) -> Result<(), StrictError> { + let bytes = nodedb_types::value_to_msgpack(value) + .map_err(|e| invalid_value(col, format!("value does not serialize to MessagePack: {e}")))?; + var_data.extend_from_slice(&bytes); + Ok(()) +} + +/// Append a lossless NodeDB MessagePack representation. +fn append_msgpack( + var_data: &mut Vec, + col: &ColumnDef, + value: &Value, +) -> Result<(), StrictError> { + let bytes = zerompk::to_msgpack_vec(value) + .map_err(|e| invalid_value(col, format!("value does not serialize to MessagePack: {e}")))?; + var_data.extend_from_slice(&bytes); + Ok(()) +} + +/// Error unless a vector holds exactly `dim` elements. +fn check_vector_len(col: &ColumnDef, dim: u32, len: usize) -> Result<(), StrictError> { + if usize::try_from(dim).is_ok_and(|d| d == len) { + Ok(()) + } else { + Err(invalid_value( + col, + format!("expected {dim} elements, got {len}"), + )) + } +} + +fn type_mismatch(col: &ColumnDef) -> StrictError { + StrictError::TypeMismatch { + column: col.name.clone(), + expected: col.column_type, + } +} + +fn invalid_value(col: &ColumnDef, detail: String) -> StrictError { + StrictError::InvalidValue { + column: col.name.clone(), + expected: col.column_type, + detail, + } +} diff --git a/nodedb-strict/src/error.rs b/nodedb-strict/src/error.rs index f1a67e641..11b8ea798 100644 --- a/nodedb-strict/src/error.rs +++ b/nodedb-strict/src/error.rs @@ -19,6 +19,15 @@ pub enum StrictError { expected: ColumnType, }, + /// A value of an accepted variant does not convert to the column type. + /// A text that does not parse as a UUID is one example. + #[error("column '{column}': invalid {expected} value: {detail}")] + InvalidValue { + column: String, + expected: ColumnType, + detail: String, + }, + /// A non-nullable column received a null value with no default. #[error("column '{0}' is NOT NULL and has no default")] NullViolation(String), diff --git a/nodedb-types/src/columnar/column_def.rs b/nodedb-types/src/columnar/column_def.rs index 8a1e6790e..3fd9e3369 100644 --- a/nodedb-types/src/columnar/column_def.rs +++ b/nodedb-types/src/columnar/column_def.rs @@ -8,6 +8,8 @@ use std::fmt; use serde::{Deserialize, Serialize}; use super::column_type::ColumnType; +use super::float_width::FloatWidth; +use super::int_width::IntWidth; /// Column-level modifiers that designate special engine roles. /// @@ -72,6 +74,16 @@ pub struct ColumnDef { /// sub-schema for tuples written under older versions. #[serde(default = "default_added_at_version")] pub added_at_version: u32, + /// The declared width of an `Int64` column. Storage stays 8 bytes. + /// `SMALLINT` and `INTEGER` bound every stored value and narrow the wire + /// type. `None` is `BIGINT`. + #[serde(default, skip_serializing_if = "Option::is_none")] + pub int_width: Option, + /// The declared width of a `Float64` column. Storage stays 8 bytes. + /// `REAL` refuses a finite value past `f32` and narrows the wire type. + /// `None` is `DOUBLE PRECISION`. + #[serde(default, skip_serializing_if = "Option::is_none")] + pub float_width: Option, } fn default_added_at_version() -> u32 { @@ -90,6 +102,8 @@ impl ColumnDef { generated_expr: None, generated_deps: Vec::new(), added_at_version: 1, + int_width: None, + float_width: None, } } @@ -104,6 +118,36 @@ impl ColumnDef { generated_expr: None, generated_deps: Vec::new(), added_at_version: 1, + int_width: None, + float_width: None, + } + } + + /// Record the numeric width `declared` names. `declared` is the DDL type + /// text of this column, modifiers included. + /// + /// Only an `Int64` column takes an integer width, and only a `Float64` + /// column takes a float width. Every other column keeps no width. + pub fn with_declared_width(mut self, declared: &str) -> Self { + self.int_width = match self.column_type { + ColumnType::Int64 => IntWidth::from_declared_type(declared), + _ => None, + }; + self.float_width = match self.column_type { + ColumnType::Float64 => FloatWidth::from_declared_type(declared), + _ => None, + }; + self + } + + /// The SQL type name this column declares: its numeric width when one is + /// set, else the name of its column type. The name parses back to the + /// same column type and width. + pub fn declared_type_name(&self) -> String { + match (self.int_width, self.float_width) { + (Some(width), _) => width.pg_type_name().to_ascii_uppercase(), + (None, Some(width)) => width.pg_type_name().to_ascii_uppercase(), + (None, None) => self.column_type.to_string(), } } @@ -131,7 +175,7 @@ impl ColumnDef { impl fmt::Display for ColumnDef { fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { - write!(f, "{} {}", self.name, self.column_type)?; + write!(f, "{} {}", self.name, self.declared_type_name())?; if !self.nullable { write!(f, " NOT NULL")?; } @@ -144,3 +188,39 @@ impl fmt::Display for ColumnDef { Ok(()) } } + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn declared_width_lands_only_on_its_numeric_family() { + let small = ColumnDef::nullable("n", ColumnType::Int64).with_declared_width("SMALLINT"); + assert_eq!(small.int_width, Some(IntWidth::I16)); + assert_eq!(small.float_width, None); + + let real = ColumnDef::nullable("r", ColumnType::Float64).with_declared_width("REAL"); + assert_eq!(real.float_width, Some(FloatWidth::F32)); + assert_eq!(real.int_width, None); + + let text = ColumnDef::nullable("t", ColumnType::String).with_declared_width("SMALLINT"); + assert_eq!((text.int_width, text.float_width), (None, None)); + } + + /// The declared name parses back to the column type the column holds. + #[test] + fn declared_type_name_parses_back_to_the_same_column() { + for declared in ["SMALLINT", "INTEGER", "BIGINT", "REAL", "DOUBLE PRECISION"] { + let column_type: ColumnType = declared.parse().expect("declared type parses"); + let column = ColumnDef::nullable("c", column_type).with_declared_width(declared); + assert_eq!(column.declared_type_name(), declared); + let reparsed = ColumnDef::nullable("c", column_type) + .with_declared_width(&column.declared_type_name()); + assert_eq!(reparsed, column); + } + assert_eq!( + ColumnDef::nullable("c", ColumnType::Int64).declared_type_name(), + "BIGINT" + ); + } +} diff --git a/nodedb/src/control/planner/catalog_adapter/type_convert.rs b/nodedb/src/control/planner/catalog_adapter/type_convert.rs index 2a21d4f31..62382febb 100644 --- a/nodedb/src/control/planner/catalog_adapter/type_convert.rs +++ b/nodedb/src/control/planner/catalog_adapter/type_convert.rs @@ -5,6 +5,29 @@ use nodedb_sql::types::{ColumnInfo, EngineType, SqlDataType}; use nodedb_types::columnar::{FloatWidth, IntWidth}; +/// The declared key column of a document collection, per +/// `nodedb_types::declared_key` over the primary key +/// [`convert_collection_type`] resolves. `None` for a collection keyed by the +/// implicit `id` or `_rowid`, and for every non-document collection: a KV row +/// names its key by its own rule, and a columnar row is not a sparse row. +/// +/// The one source the Data-Plane register config and the Control-Plane write +/// admission both read, so a scan row and a write image name the identity +/// column alike. +pub(crate) fn document_declared_key( + stored: &crate::control::security::catalog::StoredCollection, +) -> Option { + match &stored.collection_type { + nodedb_types::CollectionType::Document(_) => { + let (_, _, primary_key) = convert_collection_type(stored); + nodedb_types::declared_key(primary_key.as_deref()).map(str::to_string) + } + nodedb_types::CollectionType::KeyValue(_) | nodedb_types::CollectionType::Columnar(_) => { + None + } + } +} + /// Convert a StoredCollection to engine type, columns, and primary key. pub(crate) fn convert_collection_type( stored: &crate::control::security::catalog::StoredCollection, @@ -12,38 +35,13 @@ pub(crate) fn convert_collection_type( use nodedb_types::CollectionType; use nodedb_types::columnar::DocumentMode; - // Declared numeric widths, resolved once per collection from the raw DDL - // type strings the catalog records in `fields` for *every* engine. - // - // Strict and KV columns are typed by a resolved `ColumnType`, which - // deliberately has one `Int64` variant for every declared integer width - // and one `Float64` variant for every declared float width (nodedb stores - // all integers as i64 and all floats as f64). `fields` is therefore the - // only surviving record of what the author actually wrote, and it is - // populated for strict/KV exactly as it is for schemaless/columnar — see - // `ddl::neutral::collection::create::build`, which fills it from the raw - // column list before the typed schema is built. Resolving from it here - // keeps declared-width fidelity uniform across all engines without - // widening any persisted structure. - let declared_int = declared_widths(&stored.fields, IntWidth::from_declared_type); - let declared_float = declared_widths(&stored.fields, FloatWidth::from_declared_type); - + // Strict and KV columns take their declared numeric width from the typed + // schema, the same width the Data Plane enforces on write. Schemaless and + // columnar-family columns resolve it from the declared text in `fields`, + // the text their write rule is built from. match &stored.collection_type { CollectionType::Document(DocumentMode::Strict(schema)) => { - let columns = schema - .columns - .iter() - .map(|c| ColumnInfo { - name: c.name.clone(), - data_type: convert_column_type(&c.column_type), - nullable: c.nullable, - is_primary_key: c.primary_key, - default: c.default.clone(), - raw_type: None, - int_width: lookup_width(&declared_int, &c.name), - float_width: lookup_width(&declared_float, &c.name), - }) - .collect(); + let columns = schema.columns.iter().map(schema_column_info).collect(); let pk = schema .columns .iter() @@ -94,16 +92,7 @@ pub(crate) fn convert_collection_type( .schema .columns .iter() - .map(|c| ColumnInfo { - name: c.name.clone(), - data_type: convert_column_type(&c.column_type), - nullable: c.nullable, - is_primary_key: c.primary_key, - default: c.default.clone(), - raw_type: None, - int_width: lookup_width(&declared_int, &c.name), - float_width: lookup_width(&declared_float, &c.name), - }) + .map(schema_column_info) .collect(); let pk = config .schema @@ -214,40 +203,22 @@ fn declared_default(type_str: &str) -> Option { default_expr } -/// Resolve the declared width of every catalog field `resolve` recognizes, -/// keyed by column name. +/// The planner-facing column a typed strict or KV schema column is. /// -/// Generic over the width family so the integer and float passes share one -/// implementation: `resolve` is `IntWidth::from_declared_type` or -/// `FloatWidth::from_declared_type`, each of which is the single source of -/// truth for its own keyword set. -/// -/// Fields `resolve` does not recognize are dropped rather than stored as -/// `None`, so the result is usually empty and the common case costs one -/// allocation of zero capacity. -fn declared_widths( - fields: &[(String, String)], - resolve: fn(&str) -> Option, -) -> Vec<(&str, W)> { - fields - .iter() - .filter_map(|(name, type_str)| resolve(type_str).map(|w| (name.as_str(), w))) - .collect() -} - -/// Look up a column's declared width by name, case-insensitively to match the -/// rest of this module's column-name comparisons. -/// -/// `None` means either "not a column of this numeric family" or "the catalog -/// has no record of this column's declared type" — for example a column added -/// by `ALTER ADD COLUMN`, whose declared width was never recorded in `fields`. -/// Both degrade to the widest wire type of the family (`BIGINT` / -/// `double precision`), which is the only lossless fallback. -fn lookup_width(widths: &[(&str, W)], column: &str) -> Option { - widths - .iter() - .find(|(name, _)| name.eq_ignore_ascii_case(column)) - .map(|(_, w)| *w) +/// The declared numeric width comes from the schema column, the width the +/// Data Plane enforces on write. An absent width is the widest wire type of +/// its family (`BIGINT` / `double precision`). +fn schema_column_info(column: &nodedb_types::columnar::ColumnDef) -> ColumnInfo { + ColumnInfo { + name: column.name.clone(), + data_type: convert_column_type(&column.column_type), + nullable: column.nullable, + is_primary_key: column.primary_key, + default: column.default.clone(), + raw_type: None, + int_width: column.int_width, + float_width: column.float_width, + } } fn convert_column_type(ct: &nodedb_types::columnar::ColumnType) -> SqlDataType { @@ -257,17 +228,20 @@ fn convert_column_type(ct: &nodedb_types::columnar::ColumnType) -> SqlDataType { ColumnType::Float64 => SqlDataType::Float64, ColumnType::String => SqlDataType::String, ColumnType::Bool => SqlDataType::Bool, - ColumnType::Bytes | ColumnType::Geometry | ColumnType::Json => SqlDataType::Bytes, + ColumnType::Bytes => SqlDataType::Bytes, + // A structured column reads back as its JSON text. + ColumnType::Json + | ColumnType::Array + | ColumnType::Set + | ColumnType::Range + | ColumnType::Record => SqlDataType::Json, + ColumnType::Geometry => SqlDataType::Geometry, ColumnType::Timestamp | ColumnType::SystemTimestamp => SqlDataType::Timestamp, ColumnType::Timestamptz => SqlDataType::Timestamptz, - ColumnType::Decimal { .. } => SqlDataType::Decimal, - ColumnType::Uuid | ColumnType::Ulid | ColumnType::Regex | ColumnType::SparseVector => { - SqlDataType::String - } + ColumnType::Decimal(typmod) => SqlDataType::Decimal(*typmod), + ColumnType::Uuid => SqlDataType::Uuid, + ColumnType::Ulid | ColumnType::Regex | ColumnType::SparseVector => SqlDataType::String, ColumnType::Duration => SqlDataType::Int64, - ColumnType::Array | ColumnType::Set | ColumnType::Range | ColumnType::Record => { - SqlDataType::Bytes - } ColumnType::Vector(dim) => SqlDataType::Vector(*dim as usize), // ColumnType is #[non_exhaustive]; unknown types surface as Bytes // until the planner learns about them. @@ -378,6 +352,22 @@ mod tests { assert_eq!(parse_type_str("int2"), SqlDataType::Int64); } + /// A declared `DECIMAL(p,s)` field carries its typmod to the planner, so + /// the schemaless and key-value write paths fit values to it. + #[test] + fn parse_type_str_keeps_the_decimal_typmod() { + let typmod = nodedb_types::columnar::DecimalTypmod::new(5, 2).expect("valid typmod"); + assert_eq!( + parse_type_str("DECIMAL(5, 2) NOT NULL"), + SqlDataType::Decimal(Some(typmod)) + ); + assert_eq!( + parse_type_str("NUMERIC(5,2)"), + SqlDataType::Decimal(Some(typmod)) + ); + assert_eq!(parse_type_str("DECIMAL"), SqlDataType::Decimal(None)); + } + /// Every float spelling `FloatWidth::from_declared_type` recognizes must /// also resolve to `SqlDataType::Float64` here — `FLOAT4`/`FLOAT8` were /// rejected by DDL entirely, and a spelling this function does not list @@ -406,45 +396,39 @@ mod tests { } } - /// Declared float widths must be recovered for *every* engine, from the - /// same `fields` entries the integer widths come from. A strict collection - /// is the case that motivated this: its typed schema collapses `REAL` and - /// `DOUBLE` to one `Float64` column type, so `fields` is the only record of - /// what was declared. + /// A strict column advertises the numeric width its schema column + /// declares, the width the Data Plane enforces on write. The catalog + /// `fields` text plays no part. #[test] - fn declared_float_widths_are_recovered_for_strict_columns() { + fn strict_columns_take_their_width_from_the_schema() { use nodedb_types::columnar::{ - ColumnDef, ColumnType, DocumentMode, FloatWidth, StrictSchema, + ColumnDef, ColumnType, DocumentMode, FloatWidth, IntWidth, StrictSchema, }; - let schema = StrictSchema::new( - ["r", "d", "f"] - .into_iter() - .map(|name| ColumnDef::nullable(name, ColumnType::Float64)) - .collect(), - ) - .expect("three nullable float columns are a valid strict schema"); + let schema = StrictSchema::new(vec![ + ColumnDef::nullable("r", ColumnType::Float64).with_declared_width("REAL"), + ColumnDef::nullable("d", ColumnType::Float64).with_declared_width("DOUBLE"), + ColumnDef::nullable("f", ColumnType::Float64).with_declared_width("FLOAT"), + ColumnDef::nullable("s", ColumnType::Int64).with_declared_width("SMALLINT"), + ]) + .expect("nullable numeric columns are a valid strict schema"); let mut stored = StoredCollection::new(1, "coll", "owner"); stored.collection_type = CollectionType::Document(DocumentMode::Strict(schema)); - stored.fields = vec![ - ("r".to_string(), "REAL".to_string()), - ("d".to_string(), "DOUBLE".to_string()), - ("f".to_string(), "FLOAT".to_string()), - ]; let (_, columns, _) = convert_collection_type(&stored); let width_of = |name: &str| { - columns + let column = columns .iter() .find(|c| c.name == name) - .unwrap_or_else(|| panic!("column {name} must be present")) - .float_width + .unwrap_or_else(|| panic!("column {name} must be present")); + (column.int_width, column.float_width) }; - assert_eq!(width_of("r"), Some(FloatWidth::F32)); - assert_eq!(width_of("d"), Some(FloatWidth::F64)); + assert_eq!(width_of("r"), (None, Some(FloatWidth::F32))); + assert_eq!(width_of("d"), (None, Some(FloatWidth::F64))); // Bare FLOAT is double precision, not single. - assert_eq!(width_of("f"), Some(FloatWidth::F64)); + assert_eq!(width_of("f"), (None, Some(FloatWidth::F64))); + assert_eq!(width_of("s"), (Some(IntWidth::I16), None)); } /// A columnar (or spatial, which shares the same non-timeseries @@ -514,4 +498,28 @@ mod tests { ); assert_eq!(pk.as_deref(), Some("id")); } + + /// A strict or KV structured column is `Json`, so it advertises `json` + /// and reads back as JSON text. A `BYTEA` column stays `Bytes`. + #[test] + fn structured_schema_columns_are_json() { + use nodedb_types::columnar::ColumnType; + for structured in [ + ColumnType::Json, + ColumnType::Array, + ColumnType::Set, + ColumnType::Range, + ColumnType::Record, + ] { + assert_eq!( + super::convert_column_type(&structured), + SqlDataType::Json, + "{structured}" + ); + } + assert_eq!( + super::convert_column_type(&ColumnType::Bytes), + SqlDataType::Bytes + ); + } } diff --git a/nodedb/src/control/planner/sql_plan_convert/dml/insert/schema.rs b/nodedb/src/control/planner/sql_plan_convert/dml/insert/schema.rs index 81b236e94..72bf0dacc 100644 --- a/nodedb/src/control/planner/sql_plan_convert/dml/insert/schema.rs +++ b/nodedb/src/control/planner/sql_plan_convert/dml/insert/schema.rs @@ -34,22 +34,18 @@ pub(crate) fn build_columnar_schema( let mut cols = Vec::with_capacity(column_schema.len()); let mut has_id = false; for (name, type_str) in column_schema { - // `type_str` may contain SQL modifiers such as `NOT NULL` or `PRIMARY KEY` - // (e.g. "BIGINT NOT NULL"). Strip everything after the first token so that - // `ColumnType::from_str` receives the bare type name (e.g. "BIGINT"). - let bare_type = type_str - .split_whitespace() - .next() - .unwrap_or(type_str.as_str()); - let col_type = bare_type - .parse::() - .unwrap_or(ColumnType::String); - if name == identity_column { + // `type_str` carries SQL modifiers such as `NOT NULL` or `PRIMARY KEY` + // ("BIGINT NOT NULL", "DECIMAL(10, 2) NOT NULL"). The declared-type + // resolver reads the leading type token and keeps a spaced parameter + // list whole. + let col_type = ColumnType::from_declared_type(type_str).unwrap_or(ColumnType::String); + let col = if name == identity_column { has_id = true; - cols.push(ColumnDef::required(name.clone(), col_type).with_primary_key()); + ColumnDef::required(name.clone(), col_type).with_primary_key() } else { - cols.push(ColumnDef::nullable(name.clone(), col_type)); - } + ColumnDef::nullable(name.clone(), col_type) + }; + cols.push(col.with_declared_width(type_str)); } // No column carries the identity: synthesize it. if !has_id { @@ -76,3 +72,39 @@ pub(in super::super) fn build_schema_bytes( .map(|schema| zerompk::to_msgpack_vec(&schema).unwrap_or_default()) .unwrap_or_default() } + +#[cfg(test)] +mod tests { + use nodedb_types::columnar::{DecimalTypmod, FloatWidth, IntWidth}; + + use super::*; + + #[test] + fn columnar_schema_keeps_declared_widths_and_typmods() { + let fields: Vec<(String, String)> = [ + ("id", "BIGINT PRIMARY KEY"), + ("s", "SMALLINT NOT NULL"), + ("r", "REAL"), + ("d", "DECIMAL(10, 2) NOT NULL"), + ] + .into_iter() + .map(|(name, declared)| (name.to_string(), declared.to_string())) + .collect(); + let schema = build_columnar_schema(&fields, "id").expect("valid schema"); + let column = |name: &str| { + schema + .columns + .iter() + .find(|c| c.name == name) + .cloned() + .expect("column present") + }; + assert_eq!(column("id").int_width, Some(IntWidth::I64)); + assert_eq!(column("s").int_width, Some(IntWidth::I16)); + assert_eq!(column("r").float_width, Some(FloatWidth::F32)); + assert_eq!( + column("d").column_type, + ColumnType::Decimal(Some(DecimalTypmod::new(10, 2).expect("valid typmod"))) + ); + } +} diff --git a/nodedb/src/control/server/pgwire/handler/shape_encode/array_text.rs b/nodedb/src/control/server/pgwire/handler/shape_encode/array_text.rs new file mode 100644 index 000000000..253bae3ff --- /dev/null +++ b/nodedb/src/control/server/pgwire/handler/shape_encode/array_text.rs @@ -0,0 +1,310 @@ +// SPDX-License-Identifier: BUSL-1.1 + +//! The PostgreSQL text form of an array cell: `{1,2.5,NaN,NULL}`. +//! +//! Each dimension is a `{...}` list of comma-separated elements. SQL NULL is +//! the unquoted word `NULL`. An element whose text is empty, is `NULL` in any +//! case, or holds a brace, comma, double quote, backslash or whitespace is +//! double-quoted, with each `"` and `\` escaped by a backslash. A nested +//! list is one more dimension. Every list of one dimension holds the same +//! number of elements, as PostgreSQL requires. + +use std::error::Error; + +use bytes::{BufMut, BytesMut}; +use nodedb_types::Value; +use nodedb_types::columnar::FloatWidth; +use nodedb_types::error::NodeDbError; +use nodedb_types::value::non_finite_float_text; +use pgwire::api::Type; +use pgwire::error::PgWireResult; +use pgwire::types::ToSqlText; +use pgwire::types::format::FormatOptions; +use postgres_types::{IsNull, ToSql, accepts, to_sql_checked}; + +use crate::control::server::pgwire::numeric_narrow::checked_narrow_f32; +use crate::control::server::pgwire::types::error_map::shape_error_to_pg; +use crate::control::server::response_shape::cell::shape_mismatch; + +/// The text of a `float4[]` (`width` `F32`) or `float8[]` (`width` `F64`) +/// cell of column `column`. The cell is an array of floats, integers, NULLs +/// or nested arrays, or a vector. Any other shape is an error naming the +/// column. +pub(super) fn float_array_text(column: &str, v: &Value, width: FloatWidth) -> PgWireResult { + let mut out = String::new(); + match v { + Value::Array(items) => write_list(column, items, width, &mut out)?, + Value::Vector(floats) => { + out.push('{'); + for (i, f) in floats.iter().enumerate() { + if i > 0 { + out.push(','); + } + push_element(&mut out, &float_text(f64::from(*f), width)?); + } + out.push('}'); + } + other => { + return Err(shape_error_to_pg(&shape_mismatch( + column, "an array", other, + ))); + } + } + Ok(out) +} + +/// A rendered `{...}` array literal as pgwire encodes it. +/// +/// pgwire double-quotes a string under an array type as one element, so the +/// literal travels through this type, which writes its bytes unchanged. An +/// array column always travels in the text format, and the binary encoding +/// is an error naming that rule. +#[derive(Debug)] +pub(super) struct PgArrayLiteral(pub(super) String); + +impl ToSqlText for PgArrayLiteral { + fn to_sql_text( + &self, + _ty: &Type, + out: &mut BytesMut, + _format_options: &FormatOptions, + ) -> Result> { + out.put_slice(self.0.as_bytes()); + Ok(IsNull::No) + } +} + +impl ToSql for PgArrayLiteral { + fn to_sql( + &self, + ty: &Type, + _out: &mut BytesMut, + ) -> Result> { + Err( + format!("a {ty} cell has no binary encoding: array columns travel in the text format") + .into(), + ) + } + + accepts!(FLOAT4_ARRAY, FLOAT8_ARRAY); + + to_sql_checked!(); +} + +/// Append one `{...}` list. A list holds only nested lists of one length, +/// or only scalars and NULLs. +fn write_list( + column: &str, + items: &[Value], + width: FloatWidth, + out: &mut String, +) -> PgWireResult<()> { + check_rectangular(column, items)?; + out.push('{'); + for (i, item) in items.iter().enumerate() { + if i > 0 { + out.push(','); + } + match item { + Value::Null => out.push_str("NULL"), + Value::Array(inner) => write_list(column, inner, width, out)?, + Value::Float(f) => push_element(out, &float_text(*f, width)?), + Value::Integer(n) => push_element(out, &float_text(*n as f64, width)?), + other => return Err(shape_error_to_pg(&shape_mismatch(column, "a float", other))), + } + } + out.push('}'); + Ok(()) +} + +/// A list's items are all nested lists of one length, or none is a list. +fn check_rectangular(column: &str, items: &[Value]) -> PgWireResult<()> { + let mut lists = items.iter().map(|item| match item { + Value::Array(inner) => Some(inner.len()), + _ => None, + }); + let Some(first) = lists.next() else { + return Ok(()); + }; + if lists.all(|len| len == first) { + return Ok(()); + } + Err(shape_error_to_pg(&NodeDbError::serialization( + "cell", + format!( + "column \"{column}\" holds a ragged array: every list of one dimension must hold \ + the same number of elements" + ), + ))) +} + +/// The PostgreSQL text of one float element. A non-finite float is `NaN`, +/// `Infinity` or `-Infinity`. A finite float is its shortest round-trip +/// decimal at `width`. A finite value beyond `f32` range under `F32` is an +/// out-of-range error. +fn float_text(f: f64, width: FloatWidth) -> PgWireResult { + if let Some(text) = non_finite_float_text(f) { + return Ok(text.to_owned()); + } + Ok(match width { + FloatWidth::F64 => f.to_string(), + FloatWidth::F32 => checked_narrow_f32(f)?.to_string(), + }) +} + +/// Append one element's text, double-quoted when PostgreSQL quotes it. +pub(super) fn push_element(out: &mut String, text: &str) { + if !needs_quotes(text) { + out.push_str(text); + return; + } + out.push('"'); + for ch in text.chars() { + if matches!(ch, '"' | '\\') { + out.push('\\'); + } + out.push(ch); + } + out.push('"'); +} + +/// Whether an element's text needs double quotes to read back as itself. +fn needs_quotes(text: &str) -> bool { + text.is_empty() + || text.eq_ignore_ascii_case("NULL") + || text.chars().any(|ch| { + matches!( + ch, + '{' | '}' | ',' | '"' | '\\' | ' ' | '\t' | '\n' | '\r' | '\u{b}' | '\u{c}' + ) + }) +} + +#[cfg(test)] +mod tests { + use std::sync::Arc; + + use pgwire::error::PgWireError; + + use super::*; + + fn floats(values: &[f64]) -> Value { + Value::Array(values.iter().copied().map(Value::Float).collect()) + } + + fn text_of(v: &Value, width: FloatWidth) -> String { + float_array_text("c", v, width).expect("encodes") + } + + fn message_of(err: PgWireError) -> String { + let PgWireError::UserError(info) = err else { + panic!("expected a UserError, got {err:?}"); + }; + info.message.clone() + } + + #[test] + fn float_arrays_render_postgres_text() { + assert_eq!( + text_of(&floats(&[1.0, 2.5, -0.125]), FloatWidth::F64), + "{1,2.5,-0.125}" + ); + assert_eq!(text_of(&floats(&[]), FloatWidth::F64), "{}"); + assert_eq!( + text_of(&Value::Array(vec![Value::Integer(3)]), FloatWidth::F64), + "{3}" + ); + // `float4` elements render the shortest text of the `f32` value. + assert_eq!(text_of(&floats(&[0.1]), FloatWidth::F32), "{0.1}"); + assert_eq!( + text_of(&Value::Vector(Arc::from([0.5f32, -2.0])), FloatWidth::F32), + "{0.5,-2}" + ); + } + + #[test] + fn non_finite_and_null_elements_render_postgres_words() { + let v = Value::Array(vec![ + Value::Float(f64::NAN), + Value::Null, + Value::Float(f64::INFINITY), + Value::Float(f64::NEG_INFINITY), + ]); + assert_eq!( + text_of(&v, FloatWidth::F64), + "{NaN,NULL,Infinity,-Infinity}" + ); + assert_eq!( + text_of(&v, FloatWidth::F32), + "{NaN,NULL,Infinity,-Infinity}" + ); + } + + #[test] + fn nested_arrays_are_extra_dimensions() { + let v = Value::Array(vec![floats(&[1.0, 2.0]), floats(&[3.0, 4.0])]); + assert_eq!(text_of(&v, FloatWidth::F64), "{{1,2},{3,4}}"); + } + + #[test] + fn ragged_and_mixed_arrays_are_refused() { + let ragged = Value::Array(vec![floats(&[1.0, 2.0]), floats(&[3.0])]); + let message = + message_of(float_array_text("emb", &ragged, FloatWidth::F64).expect_err("ragged")); + assert!( + message.contains("column \"emb\" holds a ragged array"), + "{message}" + ); + + let mixed = Value::Array(vec![Value::Float(1.0), floats(&[2.0])]); + assert!(float_array_text("emb", &mixed, FloatWidth::F64).is_err()); + } + + #[test] + fn non_float_elements_and_non_arrays_are_refused() { + let text_element = Value::Array(vec![Value::String("x".into())]); + let message = + message_of(float_array_text("emb", &text_element, FloatWidth::F64).expect_err("text")); + assert!( + message.contains("holds text where a float is required"), + "{message}" + ); + + let message = message_of( + float_array_text("emb", &Value::String("[1]".into()), FloatWidth::F64) + .expect_err("not an array"), + ); + assert!( + message.contains("holds text where an array is required"), + "{message}" + ); + } + + #[test] + fn float4_overflow_is_refused() { + assert!(float_array_text("c", &floats(&[1e39]), FloatWidth::F32).is_err()); + assert_eq!( + text_of(&floats(&[1e39]), FloatWidth::F64), + format!("{{{}}}", 1e39f64) + ); + } + + #[test] + fn elements_are_quoted_as_postgres_quotes_them() { + let quoted = |text: &str| { + let mut out = String::new(); + push_element(&mut out, text); + out + }; + assert_eq!(quoted("1.5"), "1.5"); + assert_eq!(quoted("abc"), "abc"); + assert_eq!(quoted(""), "\"\""); + assert_eq!(quoted("NULL"), "\"NULL\""); + assert_eq!(quoted("null"), "\"null\""); + assert_eq!(quoted("a,b"), "\"a,b\""); + assert_eq!(quoted("{x}"), "\"{x}\""); + assert_eq!(quoted("two words"), "\"two words\""); + assert_eq!(quoted("say \"hi\""), "\"say \\\"hi\\\"\""); + assert_eq!(quoted("back\\slash"), "\"back\\\\slash\""); + } +} diff --git a/nodedb/src/control/server/shared/ddl/neutral/collection/alter/alter_type.rs b/nodedb/src/control/server/shared/ddl/neutral/collection/alter/alter_type.rs index f9d2ee1b7..f199d3a92 100644 --- a/nodedb/src/control/server/shared/ddl/neutral/collection/alter/alter_type.rs +++ b/nodedb/src/control/server/shared/ddl/neutral/collection/alter/alter_type.rs @@ -36,7 +36,7 @@ pub(super) async fn alter_collection_alter_column_type( let tenant_id = identity.tenant_id; let new_type = nodedb_types::columnar::ColumnType::from_str(new_type_str) - .map_err(|e| err("42601", format!("invalid type '{new_type_str}': {e}")))?; + .map_err(|e| err(e.sqlstate(), format!("invalid type '{new_type_str}': {e}")))?; let (coll, mut schema) = load_strict_collection( state, @@ -48,7 +48,7 @@ pub(super) async fn alter_collection_alter_column_type( let col = schema .columns - .iter() + .iter_mut() .find(|c| c.name.eq_ignore_ascii_case(column_name)) .ok_or_else(|| { err( @@ -62,7 +62,7 @@ pub(super) async fn alter_collection_alter_column_type( // An alias change resolves to the identical `ColumnType`: every integer // width parses to `Int64`, and only the declared spelling differs. A // parameter change — `VECTOR(384)` to `VECTOR(768)`, `DECIMAL(10,2)` to - // `DECIMAL(38,10)` — shares the discriminant but rewrites every stored + // `DECIMAL(28,10)` — shares the discriminant but rewrites every stored // value, so full equality is the gate. if col.column_type != new_type { return Err(err( @@ -74,14 +74,27 @@ pub(super) async fn alter_collection_alter_column_type( ), )); } - // The resolved type already matches; the alias change lands in - // `retype_field`, which records the declared spelling. + // An alias change moves the declared numeric width. A narrower width + // bounds values that existing rows can already exceed, so it needs every + // stored row checked. Only an equal or wider width is accepted. + let retyped = col.clone().with_declared_width(new_type_str); + if narrows_declared_width(col, &retyped) { + return Err(err( + "0A000", + format!( + "type change from {} to {} narrows the column and requires every \ + stored row checked; only a change to an equal or wider type is supported", + col.declared_type_name(), + retyped.declared_type_name() + ), + )); + } + *col = retyped; schema.version = schema.version.saturating_add(1); let mut updated = coll; write_schema_back(&mut updated, schema); - // The declared spelling, not the resolved `ColumnType`, is what carries - // the integer width — see `retype_field`. + // The catalog keeps the declared spelling the schema width came from. retype_field(&mut updated, column_name, new_type_str); persist_schema_change(state, &updated).await?; @@ -94,3 +107,40 @@ pub(super) async fn alter_collection_alter_column_type( Ok(status("ALTER COLLECTION")) } + +/// Whether `to` declares a narrower integer or float width than `from`. An +/// absent width is the widest of its family. +fn narrows_declared_width( + from: &nodedb_types::columnar::ColumnDef, + to: &nodedb_types::columnar::ColumnDef, +) -> bool { + use nodedb_types::columnar::{FloatWidth, IntWidth}; + let int = |c: &nodedb_types::columnar::ColumnDef| c.int_width.unwrap_or(IntWidth::I64); + let float = |c: &nodedb_types::columnar::ColumnDef| c.float_width.unwrap_or(FloatWidth::F64); + int(to) < int(from) || float(to) < float(from) +} + +#[cfg(test)] +mod tests { + use nodedb_types::columnar::{ColumnDef, ColumnType}; + + use super::narrows_declared_width; + + fn int(declared: &str) -> ColumnDef { + ColumnDef::nullable("n", ColumnType::Int64).with_declared_width(declared) + } + + fn float(declared: &str) -> ColumnDef { + ColumnDef::nullable("f", ColumnType::Float64).with_declared_width(declared) + } + + #[test] + fn only_a_narrower_width_is_a_narrowing() { + assert!(narrows_declared_width(&int("BIGINT"), &int("SMALLINT"))); + assert!(narrows_declared_width(&int("INT"), &int("INT2"))); + assert!(narrows_declared_width(&float("DOUBLE"), &float("REAL"))); + assert!(!narrows_declared_width(&int("SMALLINT"), &int("BIGINT"))); + assert!(!narrows_declared_width(&int("INT"), &int("INTEGER"))); + assert!(!narrows_declared_width(&float("REAL"), &float("FLOAT8"))); + } +} diff --git a/nodedb/src/control/server/shared/ddl/neutral/collection/alter/strict_schema.rs b/nodedb/src/control/server/shared/ddl/neutral/collection/alter/strict_schema.rs index 35803e6eb..4e674e4c3 100644 --- a/nodedb/src/control/server/shared/ddl/neutral/collection/alter/strict_schema.rs +++ b/nodedb/src/control/server/shared/ddl/neutral/collection/alter/strict_schema.rs @@ -64,14 +64,9 @@ pub(super) fn write_schema_back( /// Retype a column's entry in `coll.fields`, the catalog's record of the /// *declared* type string each column was created with. /// -/// `fields` is not redundant with the strict schema: `ColumnType` collapses -/// every integer width onto one `Int64` variant, so the declared spelling is -/// the only surviving record of how wide the author said the column was. That -/// spelling drives the column's advertised wire OID and the range accepted on -/// write, which is exactly why `ALTER COLUMN TYPE` — whose only supported use -/// *is* an alias change such as `INT` → `BIGINT` — has to update it. Leaving -/// it stale will make the alter a silent no-op for the case it exists to -/// serve, and will keep rejecting writes the new type allows. +/// The strict schema column carries the declared numeric width, and catalog +/// introspection reads the width from this spelling. `ALTER COLUMN TYPE` +/// updates both, so the two report the same width. pub(super) fn retype_field(coll: &mut StoredCollection, column: &str, new_type: &str) { for (name, type_str) in coll.fields.iter_mut() { if name.eq_ignore_ascii_case(column) { diff --git a/nodedb/src/control/server/shared/ddl/neutral/collection/dml/indexed_vector_fields.rs b/nodedb/src/control/server/shared/ddl/neutral/collection/dml/indexed_vector_fields.rs index f9a4d928c..be440cd2f 100644 --- a/nodedb/src/control/server/shared/ddl/neutral/collection/dml/indexed_vector_fields.rs +++ b/nodedb/src/control/server/shared/ddl/neutral/collection/dml/indexed_vector_fields.rs @@ -8,6 +8,7 @@ //! - A strict collection indexes every `VECTOR(n)` column. //! - Any other collection indexes each field that has its own vector index. //! When no field has one, a default-field vector index covers `embedding`. +//! - A schemaless collection also indexes every declared `VECTOR(n)` column. //! //! The `{ ... }` handler also sends a vector insert for numeric-array //! fields, so a field no index covers is still searchable. It skips the @@ -71,5 +72,28 @@ pub(super) fn indexed_vector_fields( if named.is_empty() && has_default { named.insert(DEFAULT_VECTOR_FIELD.to_string()); } + if collection_type.is_some_and(CollectionType::is_schemaless) { + let stored = state + .credentials + .catalog() + .get_collection(database_id, tenant_id, collection) + .map_err(|e| { + DdlError::from_error_in_context( + &format!("read declared columns of \"{collection}\" for INSERT"), + &e, + ) + })?; + if let Some(stored) = stored + && stored.vector_primary.is_none() + { + named.extend( + crate::control::server::shared::ddl::schema_validation::extract_vector_fields( + &stored.fields, + ) + .into_iter() + .map(|(field, _dim, _metric)| field), + ); + } + } Ok(named) } diff --git a/nodedb/src/control/server/shared/ddl/neutral/collection/dml/insert.rs b/nodedb/src/control/server/shared/ddl/neutral/collection/dml/insert.rs index ffcf28728..9596702aa 100644 --- a/nodedb/src/control/server/shared/ddl/neutral/collection/dml/insert.rs +++ b/nodedb/src/control/server/shared/ddl/neutral/collection/dml/insert.rs @@ -365,7 +365,8 @@ async fn insert_parsed( /// The `(field, sql_type)` pairs a schemaless write contributes to the /// collection's inferred projection. `id` is the document key, never a -/// projected column. +/// projected column. An integer infers `BIGINT`: the value is an `i64`, and +/// `INT` declares the 32-bit width that range-checks every later write. fn inferred_field_types( fields: &std::collections::HashMap, ) -> Vec<(String, String)> { @@ -375,7 +376,7 @@ fn inferred_field_types( .map(|(name, value)| { let sql_type = match value { nodedb_types::Value::Float(_) => "FLOAT", - nodedb_types::Value::Integer(_) => "INT", + nodedb_types::Value::Integer(_) => "BIGINT", nodedb_types::Value::Bool(_) => "BOOL", _ => "TEXT", }; diff --git a/nodedb/src/control/server/shared/ddl/neutral/collection/helpers.rs b/nodedb/src/control/server/shared/ddl/neutral/collection/helpers.rs index 94af343a0..62dc4ac45 100644 --- a/nodedb/src/control/server/shared/ddl/neutral/collection/helpers.rs +++ b/nodedb/src/control/server/shared/ddl/neutral/collection/helpers.rs @@ -79,7 +79,8 @@ pub(crate) fn parse_origin_column_def(s: &str) -> crate::Result Vec<(String, usize)> { + if !coll.collection_type.is_schemaless() || coll.vector_primary.is_some() { + return Vec::new(); + } + crate::control::server::shared::ddl::schema_validation::extract_vector_fields(&coll.fields) + .into_iter() + .map(|(name, dim, _metric)| (name, dim)) + .collect() +} + +/// The declared numeric columns of a schemaless or KV collection, which the +/// Data Plane re-types on every write. +/// +/// A schemaless collection reads them from its declared text. A KV +/// collection reads them from its typed schema, the source its wire types +/// come from. Empty for every other collection. Strict and columnar re-type +/// through their typed schema. A vector-primary collection stores tagged +/// metadata sidecars, not field maps. A CRDT collection stores the state its +/// replicas converged on, and a refusal there would split them. +/// +/// The primary key is left out. The engines key a row on the literal as +/// written, and a lookup renders its own literal the same way, so a re-typed +/// key would no longer match its row. +fn declared_columns( + coll: &StoredCollection, +) -> Vec { + use nodedb_physical::physical_plan::DeclaredColumn; + match &coll.collection_type { + nodedb_types::CollectionType::Document(nodedb_types::DocumentMode::Schemaless) => { + if coll.vector_primary.is_some() || coll.crdt { + return Vec::new(); + } + let primary_key = coll + .declared_primary_key + .as_deref() + .unwrap_or(nodedb_types::DEFAULT_IDENTITY_COLUMN); + coll.fields + .iter() + .filter(|(name, _)| !primary_key.eq_ignore_ascii_case(name)) + .filter_map(|(name, declared)| DeclaredColumn::from_declared(name, declared)) + .collect() + } + nodedb_types::CollectionType::KeyValue(config) => config + .schema + .columns + .iter() + .filter(|column| !column.primary_key) + .filter_map(DeclaredColumn::from_column_def) + .collect(), + nodedb_types::CollectionType::Document(nodedb_types::DocumentMode::Strict(_)) + | nodedb_types::CollectionType::Columnar(_) => Vec::new(), + } +} + /// Build the `CollectionConfig` a `DocumentOp::Register` will install in /// `doc_configs`, straight from the durable catalog — storage mode, /// enforcement options, generated columns, and secondary indexes. @@ -334,6 +394,9 @@ pub(crate) fn build_doc_config_from_stored( conflict_policy: coll.conflict_policy.clone(), timeseries: build_timeseries_schema(coll), vector_primary: coll.vector_primary.clone().map(Box::new), + vector_fields: declared_vector_fields(coll), + declared_columns: declared_columns(coll), + declared_key: crate::control::planner::catalog_adapter::document_declared_key(coll), } } @@ -357,6 +420,9 @@ async fn dispatch_register_from_stored_inner( conflict_policy: config.conflict_policy.clone(), timeseries: config.timeseries.clone(), vector_primary: config.vector_primary.clone(), + vector_fields: config.vector_fields.clone(), + declared_columns: config.declared_columns.clone(), + declared_key: config.declared_key.clone(), }, ); @@ -443,4 +509,28 @@ mod tests { ); assert!(config.vector_primary.is_none()); } + + /// A declared key reaches the Data Plane config. The implicit `id` key + /// names no declared key. + #[test] + fn a_declared_key_reaches_the_doc_config() { + let mut coll = StoredCollection::new(1, "items", "owner"); + coll.declared_primary_key = Some("sku".to_string()); + let config = build_doc_config_from_stored( + &EmptyCatalog, + crate::types::TenantId::new(coll.tenant_id), + &coll, + &[], + ); + assert_eq!(config.declared_key.as_deref(), Some("sku")); + + let plain = StoredCollection::new(1, "docs", "owner"); + let config = build_doc_config_from_stored( + &EmptyCatalog, + crate::types::TenantId::new(plain.tenant_id), + &plain, + &[], + ); + assert_eq!(config.declared_key, None); + } } diff --git a/nodedb/src/control/server/shared/ddl/neutral/convert/typeguard_columns.rs b/nodedb/src/control/server/shared/ddl/neutral/convert/typeguard_columns.rs index 85fc6e48a..e2cb18497 100644 --- a/nodedb/src/control/server/shared/ddl/neutral/convert/typeguard_columns.rs +++ b/nodedb/src/control/server/shared/ddl/neutral/convert/typeguard_columns.rs @@ -46,17 +46,17 @@ pub(super) fn typeguards_to_column_defs( ColumnDef::required(guard.field.clone(), ct) } else { ColumnDef::nullable(guard.field.clone(), ct) - }; + } + .with_declared_width(&guard.type_expr); // A guard carries either DEFAULT or VALUE, never both. Strict schema // has one materialization slot, so both land on the column `DEFAULT`. if let Some(expr) = guard.default_expr.clone().or(guard.value_expr.clone()) { - // The resolved type's own spelling stands in for the declaration: - // a guard names no numeric width, so the canonical name resolves - // to the same width-less type the column will carry. + // The column's declared name stands in for the declaration. It + // resolves to the type and numeric width the column carries. validate_column_default( &DeclaredColumn { name: &col.name, - declared_type: &col.column_type.to_string(), + declared_type: &col.declared_type_name(), primary_key: col.primary_key, }, &expr, diff --git a/nodedb/src/data/executor/core_loop/declared_columns.rs b/nodedb/src/data/executor/core_loop/declared_columns.rs new file mode 100644 index 000000000..e7afe4b95 --- /dev/null +++ b/nodedb/src/data/executor/core_loop/declared_columns.rs @@ -0,0 +1,40 @@ +// SPDX-License-Identifier: BUSL-1.1 + +//! The declared numeric columns a schemaless document or KV write re-types, +//! read from the collection's registered config. + +use nodedb_physical::physical_plan::DeclaredColumn; + +use crate::types::{DatabaseId, TenantId}; + +use super::CoreLoop; + +impl CoreLoop { + /// The declared numeric columns of the collection `config_key` names. + /// + /// Empty when the collection declares none, is strict or columnar, or is + /// not registered on this core. + pub(in crate::data::executor) fn declared_columns( + &self, + config_key: &(DatabaseId, TenantId, String), + ) -> &[DeclaredColumn] { + self.doc_configs + .get(config_key) + .map(|config| config.declared_columns.as_slice()) + .unwrap_or(&[]) + } + + /// [`Self::declared_columns`] for a collection named by its parts. + pub(in crate::data::executor) fn declared_columns_of( + &self, + database_id: u64, + tid: u64, + collection: &str, + ) -> &[DeclaredColumn] { + self.declared_columns(&( + DatabaseId::new(database_id), + TenantId::new(tid), + collection.to_string(), + )) + } +} diff --git a/nodedb/src/data/executor/core_loop/mod.rs b/nodedb/src/data/executor/core_loop/mod.rs index 2e3c71899..cf3575947 100644 --- a/nodedb/src/data/executor/core_loop/mod.rs +++ b/nodedb/src/data/executor/core_loop/mod.rs @@ -9,11 +9,13 @@ pub(in crate::data::executor) mod checkpoint_floors; mod columnar_schema_seed; pub(in crate::data::executor) mod commit_pending; mod crdt_dead_letters; +mod declared_columns; mod decode_stored; pub(in crate::data::executor) mod deferred; mod doc_config_seed; pub(in crate::data::executor) mod event_emit; mod event_emit_engines; +pub(in crate::data::executor) mod event_image; pub(in crate::data::executor) mod event_outlet; pub(in crate::data::executor) use event_emit_engines::KvWriteEvent; pub(in crate::data::executor) mod fail_stop; diff --git a/nodedb/src/data/executor/dispatch/document.rs b/nodedb/src/data/executor/dispatch/document.rs index 9e3efe0f0..4c278828f 100644 --- a/nodedb/src/data/executor/dispatch/document.rs +++ b/nodedb/src/data/executor/dispatch/document.rs @@ -356,31 +356,7 @@ impl CoreLoop { ) } - DocumentOp::Register { - collection, - indexes, - crdt_enabled, - storage_mode, - enforcement, - bitemporal, - conflict_policy, - timeseries, - vector_primary, - } => self.execute_register_document_collection( - task, - super::super::handlers::document::write::RegisterDocumentCollectionParams { - tid, - collection: collection.as_str(), - indexes, - crdt_enabled: *crdt_enabled, - storage_mode, - enforcement, - bitemporal: *bitemporal, - conflict_policy: conflict_policy.as_deref(), - timeseries: timeseries.as_deref(), - vector_primary: vector_primary.as_deref(), - }, - ), + DocumentOp::Register { .. } => self.dispatch_register(task, tid, op), DocumentOp::IndexLookup { collection, diff --git a/nodedb/src/data/executor/dispatch/document_register.rs b/nodedb/src/data/executor/dispatch/document_register.rs new file mode 100644 index 000000000..acb02f8c1 --- /dev/null +++ b/nodedb/src/data/executor/dispatch/document_register.rs @@ -0,0 +1,62 @@ +// SPDX-License-Identifier: BUSL-1.1 + +//! Dispatch for `DocumentOp::Register`, the DDL op that installs a document +//! collection's config on this core. + +use crate::bridge::envelope::{ErrorCode, Response}; +use nodedb_physical::physical_plan::DocumentOp; + +use crate::data::executor::core_loop::CoreLoop; +use crate::data::executor::handlers::document::write::RegisterDocumentCollectionParams; +use crate::data::executor::task::ExecutionTask; + +impl CoreLoop { + /// Install the collection config a `Register` op carries. + pub(super) fn dispatch_register( + &mut self, + task: &ExecutionTask, + tid: u64, + op: &DocumentOp, + ) -> Response { + let DocumentOp::Register { + collection, + indexes, + crdt_enabled, + storage_mode, + enforcement, + bitemporal, + conflict_policy, + timeseries, + vector_primary, + vector_fields, + declared_columns, + declared_key, + } = op + else { + return self.response_error( + task, + ErrorCode::Internal { + detail: "dispatch_register: plan is not Register".into(), + }, + ); + }; + self.execute_register_document_collection( + task, + RegisterDocumentCollectionParams { + tid, + collection: collection.as_str(), + indexes, + crdt_enabled: *crdt_enabled, + storage_mode, + enforcement, + bitemporal: *bitemporal, + conflict_policy: conflict_policy.as_deref(), + timeseries: timeseries.as_deref(), + vector_primary: vector_primary.as_deref(), + vector_fields, + declared_columns, + declared_key: declared_key.as_deref(), + }, + ) + } +} diff --git a/nodedb/src/data/executor/dispatch/kv.rs b/nodedb/src/data/executor/dispatch/kv.rs index 57a5a01b9..97978bf73 100644 --- a/nodedb/src/data/executor/dispatch/kv.rs +++ b/nodedb/src/data/executor/dispatch/kv.rs @@ -1,7 +1,8 @@ // SPDX-License-Identifier: BUSL-1.1 //! Dispatch for KvOp variants: engine pressure check, refusal of an unbound -//! row write, then delegation to execute_kv. +//! row write, declared-column re-typing of incoming row bodies, then +//! delegation to execute_kv. use crate::bridge::envelope::Response; use nodedb_physical::physical_plan::KvOp; @@ -44,6 +45,12 @@ impl CoreLoop { { return self.response_error(task, refusal); } - self.execute_kv(task, did, tid, op) + // Every row body the op supplies whole holds its declared numeric + // columns' values before any handler sees it. + let coerced = match self.coerce_kv_op_bodies(did, tid, op) { + Ok(coerced) => coerced, + Err(e) => return self.response_error(task, e), + }; + self.execute_kv(task, did, tid, coerced.as_ref().unwrap_or(op)) } } diff --git a/nodedb/src/data/executor/dispatch/mod.rs b/nodedb/src/data/executor/dispatch/mod.rs index 6b61a6644..237ecb855 100644 --- a/nodedb/src/data/executor/dispatch/mod.rs +++ b/nodedb/src/data/executor/dispatch/mod.rs @@ -10,6 +10,7 @@ pub mod crdt; pub mod document; mod document_admit; mod document_dml; +mod document_register; mod execute; pub mod graph; pub mod kv; diff --git a/nodedb/src/data/executor/handlers/columnar_resolved_mutation.rs b/nodedb/src/data/executor/handlers/columnar_resolved_mutation.rs index 572e4964a..aa0d280b1 100644 --- a/nodedb/src/data/executor/handlers/columnar_resolved_mutation.rs +++ b/nodedb/src/data/executor/handlers/columnar_resolved_mutation.rs @@ -31,6 +31,7 @@ use tracing::debug; use crate::bridge::envelope::{ErrorCode, Response}; use crate::data::executor::core_loop::CoreLoop; +use crate::data::executor::handlers::columnar_write::coerce_columnar_row; use crate::data::executor::handlers::transaction::undo::UndoEntry; use crate::data::executor::task::ExecutionTask; @@ -89,10 +90,21 @@ impl CoreLoop { } } + // Every shipped post-image meets the declared column rule before the + // policy decides it and before the first row changes. The Control + // Plane computed these values, so a value past a declared width + // refuses the statement whole. + let mut rows = rows.to_vec(); + for (_, new_row) in &mut rows { + if let Err(e) = coerce_columnar_row(&schema, new_row) { + return self.response_error(task, e); + } + } + // The gate stays on every write path even though `DecidedEarlierInRequest` // makes this a no-op — a single path that skips it entirely is a hole // future callers can fall into. - for (_pk, new_row) in rows { + for (_pk, new_row) in &rows { if let Err(error) = crate::data::executor::handlers::rls_write_gate::admit_columnar_row( rls_write_check, new_row, @@ -106,8 +118,16 @@ impl CoreLoop { let row_count_before = engine.memtable().row_count(); let mut undo_log = undo_log; - let outcome = - self.apply_columnar_update_rows(task, &key, &schema, rows, undo_log.as_deref_mut()); + let outcome = match self.apply_columnar_update_rows( + task, + &key, + &schema, + &rows, + undo_log.as_deref_mut(), + ) { + Ok(outcome) => outcome, + Err(e) => return self.response_error(task, e), + }; let affected = outcome.affected; if let Some(log) = undo_log { @@ -132,12 +152,7 @@ impl CoreLoop { let result = serde_json::json!({ "affected": affected }); match super::super::response_codec::encode_json_as_msgpack(&result) { Ok(payload) => self.response_with_payload(task, payload), - Err(e) => self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ), + Err(e) => self.response_error(task, ErrorCode::from(e)), } } @@ -194,7 +209,11 @@ impl CoreLoop { } let mut undo_log = undo_log; - let outcome = self.apply_columnar_delete_pks(&key, &schema, pks, undo_log.as_deref_mut()); + let outcome = + match self.apply_columnar_delete_pks(&key, &schema, pks, undo_log.as_deref_mut()) { + Ok(outcome) => outcome, + Err(e) => return self.response_error(task, e), + }; let affected = outcome.affected; if let Some(log) = undo_log { @@ -212,12 +231,7 @@ impl CoreLoop { let result = serde_json::json!({ "affected": affected }); match super::super::response_codec::encode_json_as_msgpack(&result) { Ok(payload) => self.response_with_payload(task, payload), - Err(e) => self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ), + Err(e) => self.response_error(task, ErrorCode::from(e)), } } } @@ -349,6 +363,7 @@ mod tests { engine .scan_memtable_rows() .map(|r| { + let r = r.expect("read"); let id = match &r[id_idx] { Value::Integer(n) => *n, other => panic!("expected integer id, got {other:?}"), @@ -537,4 +552,132 @@ mod tests { "rejected update must not mutate the row" ); } + + /// Seed rows into a collection whose `v` column is declared `SMALLINT`. + fn insert_smallint_rows(core: &mut CoreLoop, task: &ExecutionTask, rows: Vec) { + use nodedb_types::columnar::{ColumnDef, ColumnType, ColumnarSchema}; + let schema = ColumnarSchema::new(vec![ + ColumnDef::required("id", ColumnType::Int64).with_primary_key(), + ColumnDef::nullable("v", ColumnType::Int64).with_declared_width("SMALLINT"), + ]) + .expect("valid schema"); + let schema_bytes = zerompk::to_msgpack_vec(&schema).expect("encode schema"); + let payload = + nodedb_types::value_to_msgpack(&Value::Array(rows)).expect("encode insert payload"); + let resp = core.execute_columnar_insert( + task, + crate::data::executor::handlers::columnar_write::ColumnarInsertParams { + collection: COLLECTION, + payload: &payload, + format: "msgpack", + intent: nodedb_physical::physical_plan::ColumnarInsertIntent::Insert, + on_conflict_updates: &[], + surrogates: &[], + schema_bytes: &schema_bytes, + provenance: None, + rls_write_check: &RlsWriteCheck::already_decided_elsewhere(), + returning: None, + rls_filters: &[], + spatial_undo: None, + }, + ); + assert_eq!( + resp.status, + crate::bridge::envelope::Status::Ok, + "seed insert failed: {:?}", + resp.error_code + ); + } + + fn is_out_of_range(resp: &crate::bridge::envelope::Response) -> bool { + matches!( + resp.error_code.as_deref(), + Some(ErrorCode::NumericValueOutOfRange { .. }) + ) + } + + /// A SMALLINT column refuses a value past its width on INSERT. + #[test] + fn declared_smallint_refuses_an_insert_past_its_width() { + let mut h = make_core(); + let task = task_for(&h); + insert_smallint_rows(&mut h.core, &task, vec![row(1, 1)]); + + let payload = + nodedb_types::value_to_msgpack(&Value::Array(vec![row(2, 40000)])).expect("encode"); + let resp = h.core.execute_columnar_insert( + &task, + crate::data::executor::handlers::columnar_write::ColumnarInsertParams { + collection: COLLECTION, + payload: &payload, + format: "msgpack", + intent: nodedb_physical::physical_plan::ColumnarInsertIntent::Insert, + on_conflict_updates: &[], + surrogates: &[], + schema_bytes: &[], + provenance: None, + rls_write_check: &RlsWriteCheck::already_decided_elsewhere(), + returning: None, + rls_filters: &[], + spatial_undo: None, + }, + ); + assert!(is_out_of_range(&resp), "{:?}", resp.error_code); + assert_eq!(scan_ids(&mut h.core), vec![(1, 1)]); + } + + /// A computed post-image past the declared width, here `v + 39999` over a + /// stored `1`, refuses the whole resolved update and changes no row. + #[test] + fn resolved_update_past_a_declared_width_changes_nothing() { + let mut h = make_core(); + let task = task_for(&h); + insert_smallint_rows(&mut h.core, &task, vec![row(1, 1), row(2, 2)]); + + let rows = vec![ + (Value::Integer(1), resolved_row(&h.core, 1, 40000)), + (Value::Integer(2), resolved_row(&h.core, 2, 3)), + ]; + let resp = h.core.execute_columnar_resolved_update( + &task, + COLLECTION, + &rows, + &RlsWriteCheck::decided_earlier_in_request(), + None, + ); + assert!(is_out_of_range(&resp), "{:?}", resp.error_code); + + let mut ids = scan_ids(&mut h.core); + ids.sort(); + assert_eq!(ids, vec![(1, 1), (2, 2)]); + } + + /// A predicate update whose post-image is past the declared width is + /// refused before the first row changes. + #[test] + fn predicate_update_past_a_declared_width_changes_nothing() { + let mut h = make_core(); + let task = task_for(&h); + insert_smallint_rows(&mut h.core, &task, vec![row(1, 1), row(2, 2)]); + + let updates = vec![( + "v".to_string(), + nodedb_physical::physical_plan::UpdateValue::Literal( + nodedb_types::value_to_msgpack(&Value::Integer(40000)).expect("encode"), + ), + )]; + let resp = h.core.execute_columnar_update( + &task, + COLLECTION, + &[], + &updates, + &RlsWriteCheck::decided_earlier_in_request(), + None, + ); + assert!(is_out_of_range(&resp), "{:?}", resp.error_code); + + let mut ids = scan_ids(&mut h.core); + ids.sort(); + assert_eq!(ids, vec![(1, 1), (2, 2)]); + } } diff --git a/nodedb/src/data/executor/handlers/columnar_write/mod.rs b/nodedb/src/data/executor/handlers/columnar_write/mod.rs index 331ca3724..c8369722e 100644 --- a/nodedb/src/data/executor/handlers/columnar_write/mod.rs +++ b/nodedb/src/data/executor/handlers/columnar_write/mod.rs @@ -3,8 +3,9 @@ //! Columnar base insert handler. //! //! Writes rows to `nodedb-columnar`'s `MutationEngine`. Accepts msgpack payload -//! (array of objects). Creates the engine on first insert with schema inferred -//! from the first row. +//! (array of objects). Creates the engine on first insert with the catalog +//! schema the plan carries. Only a plan with no schema infers one from the +//! first row. pub mod flush; pub mod geometry_index; @@ -13,12 +14,13 @@ pub mod insert; pub mod read_prior; pub mod row_ingest; pub mod schema; -pub mod spatial; pub(in crate::data::executor) use geometry_index::GeometryIndexDelta; pub(in crate::data::executor) use geometry_remove::{RemovedSpatialEntry, schema_has_geometry}; pub(in crate::data::executor) use insert::ColumnarInsertParams; -pub(in crate::data::executor) use schema::{ndb_field_to_value, row_values_to_object}; +pub(in crate::data::executor) use schema::{ + coerce_columnar_row, ndb_field_to_value, row_values_to_object, +}; // `ensure_columnar_engine_schema` is an inherent `CoreLoop` method (defined // in `schema.rs`), called via `self.` — no re-export needed. // `flush_columnar_memtable_if_needed`, `index_columnar_geometry_columns`, and diff --git a/nodedb/src/data/executor/handlers/columnar_write/schema.rs b/nodedb/src/data/executor/handlers/columnar_write/schema.rs index e34fad710..5f2362236 100644 --- a/nodedb/src/data/executor/handlers/columnar_write/schema.rs +++ b/nodedb/src/data/executor/handlers/columnar_write/schema.rs @@ -29,6 +29,9 @@ impl CoreLoop { /// from "not yet created" for every read path, and the durable insert /// path already treats engine creation as idempotent /// (`if !self.columnar_engines.contains_key(...)`). + /// + /// Fails when `schema_bytes` is present but does not decode. A sampled + /// row never stands in for a declared schema. pub(in crate::data::executor) fn ensure_columnar_engine_schema( &mut self, engine_key: &(DatabaseId, TenantId, String), @@ -36,33 +39,32 @@ impl CoreLoop { bitemporal: bool, first_row: &Value, schema_bytes: &[u8], - ) -> ColumnarSchema { + ) -> crate::Result { if let Some(engine) = self.columnar_engines.get(engine_key) { - return engine.schema().clone(); + return Ok(engine.schema().clone()); } - let flush_threshold = self.query_tuning.columnar_flush_threshold; - let engine = self - .columnar_engines - .entry(engine_key.clone()) - .or_insert_with(|| { - let base_schema = if !schema_bytes.is_empty() { - zerompk::from_msgpack::(schema_bytes) - .unwrap_or_else(|_| infer_schema_from_value(first_row)) - } else { - infer_schema_from_value(first_row) - }; - let schema = if bitemporal { - prepend_bitemporal_columns(base_schema) - } else { - base_schema - }; - nodedb_columnar::MutationEngine::with_flush_threshold( - collection.to_string(), - schema, - flush_threshold, - ) - }); - engine.schema().clone() + let base_schema = if schema_bytes.is_empty() { + infer_schema_from_value(first_row) + } else { + zerompk::from_msgpack::(schema_bytes).map_err(|e| { + crate::Error::Serialization { + format: "msgpack".to_string(), + detail: format!("columnar schema of '{collection}' does not decode: {e}"), + } + })? + }; + let schema = if bitemporal { + prepend_bitemporal_columns(base_schema) + } else { + base_schema + }; + let engine = nodedb_columnar::MutationEngine::with_flush_threshold( + collection.to_string(), + schema.clone(), + self.query_tuning.columnar_flush_threshold, + ); + self.columnar_engines.insert(engine_key.clone(), engine); + Ok(schema) } } @@ -82,73 +84,44 @@ pub(in crate::data::executor) fn row_values_to_object( nodedb_types::Value::Object(map) } -/// Coerce a `nodedb_types::Value` field to match the column type. +/// Coerce a `nodedb_types::Value` field to the value `column` stores. +/// +/// Every column converts by the strict document coercion rule, declared +/// numeric width included, so columnar and strict collections accept and +/// refuse the same values. A refusal names the column. Text the type cannot +/// read is `InvalidTextRepresentation`, a value of the wrong kind is +/// `DatatypeMismatch`, and a value past the type's range or the declared +/// width is `NumericValueOutOfRange`. The memtable stores every shape the +/// coercion yields. /// -/// Returns `Err` if a millisecond timestamp value overflows `i64` microseconds. +/// A `SYSTEM_TIMESTAMP` column is the one exception. The columnar write path +/// assigns no value to it, so it stores the instant the row carries, in the +/// two forms the planner's type guard admits: `DateTime` and `Integer`. pub(in crate::data::executor) fn ndb_field_to_value( val: Option<&Value>, - col_type: &ColumnType, + column: &ColumnDef, ) -> crate::Result { - let Some(val) = val else { - return Ok(Value::Null); - }; - let v = match (col_type, val) { - (_, Value::Null) => Value::Null, - (ColumnType::Int64, Value::Integer(_)) => val.clone(), - (ColumnType::Int64, Value::Float(f)) => Value::Integer(*f as i64), - (ColumnType::Int64, Value::String(s)) => { - s.parse::().map(Value::Integer).unwrap_or(Value::Null) + match (&column.column_type, val) { + (_, None | Some(Value::Null)) => Ok(Value::Null), + (ColumnType::SystemTimestamp, Some(v @ (Value::DateTime(_) | Value::Integer(_)))) => { + Ok(v.clone()) } - (ColumnType::Float64, Value::Float(_)) => val.clone(), - (ColumnType::Float64, Value::Integer(n)) => Value::Float(*n as f64), - (ColumnType::Float64, Value::String(s)) => { - s.parse::().map(Value::Float).unwrap_or(Value::Null) - } - (ColumnType::Bool, Value::Bool(_)) => val.clone(), - (ColumnType::String, Value::String(_)) => val.clone(), - (ColumnType::Timestamp, Value::Integer(n)) => { - Value::NaiveDateTime(nodedb_types::NdbDateTime::from_millis(*n).map_err(|e| { - crate::Error::BadRequest { - detail: format!("timestamp coercion: {e}"), - } - })?) - } - (ColumnType::Timestamp, Value::Float(f)) => { - Value::NaiveDateTime(nodedb_types::NdbDateTime::from_millis(*f as i64).map_err( - |e| crate::Error::BadRequest { - detail: format!("timestamp coercion: {e}"), - }, - )?) - } - (ColumnType::Timestamp, Value::String(s)) => nodedb_types::datetime::NdbDateTime::parse(s) - .map(Value::NaiveDateTime) - .unwrap_or_else(|| Value::String(s.clone())), - (ColumnType::Timestamptz, Value::Integer(n)) => { - Value::DateTime(nodedb_types::NdbDateTime::from_millis(*n).map_err(|e| { - crate::Error::BadRequest { - detail: format!("timestamptz coercion: {e}"), - } - })?) - } - (ColumnType::Timestamptz, Value::Float(f)) => Value::DateTime( - nodedb_types::NdbDateTime::from_millis(*f as i64).map_err(|e| { - crate::Error::BadRequest { - detail: format!("timestamptz coercion: {e}"), - } - })?, - ), - (ColumnType::Timestamptz, Value::String(s)) => { - nodedb_types::datetime::NdbDateTime::parse(s) - .map(Value::DateTime) - .unwrap_or_else(|| Value::String(s.clone())) - } - (ColumnType::Uuid, Value::String(_)) => val.clone(), - // Fallback: integers as floats, strings as strings. - (ColumnType::Float64, _) => Value::Null, - (ColumnType::Int64, _) => Value::Null, - _ => val.clone(), - }; - Ok(v) + (_, Some(val)) => crate::data::executor::strict_format::coerce_value(val, column), + } +} + +/// Coerce every cell of the schema-ordered `row` to the value its column +/// stores, in place. An UPDATE runs this on each post-image before the +/// engine stores it, the rule an INSERT applies to each new row. +pub(in crate::data::executor) fn coerce_columnar_row( + schema: &ColumnarSchema, + row: &mut [Value], +) -> crate::Result<()> { + for (column, cell) in schema.columns.iter().zip(row.iter_mut()) { + let coerced = ndb_field_to_value(Some(&*cell), column)?; + *cell = coerced; + } + Ok(()) } /// Infer a columnar schema from a `nodedb_types::Value::Object` (first row). @@ -175,6 +148,9 @@ pub(in crate::data::executor) fn infer_schema_from_value(row: &Value) -> Columna Value::Bool(_) => ColumnType::Bool, Value::DateTime(_) => ColumnType::Timestamptz, Value::NaiveDateTime(_) => ColumnType::Timestamp, + Value::Geometry(_) => ColumnType::Geometry, + Value::Bytes(_) => ColumnType::Bytes, + Value::Object(_) | Value::Array(_) => ColumnType::Json, _ => ColumnType::String, }; let lower = key.to_lowercase(); @@ -208,8 +184,164 @@ pub(in crate::data::executor) fn prepend_bitemporal_columns( ColumnarSchema::new(cols).expect("bitemporal columnar schema must be valid") } -/// Infer a columnar schema from a JSON object — used by the spatial insert path. -pub(super) fn infer_schema_from_json(row: &serde_json::Value) -> ColumnarSchema { - let ndb: Value = row.clone().into(); - infer_schema_from_value(&ndb) +#[cfg(test)] +mod tests { + use std::collections::HashMap; + + use super::*; + + fn column_type(schema: &ColumnarSchema, name: &str) -> ColumnType { + schema + .columns + .iter() + .find(|c| c.name == name) + .map(|c| c.column_type) + .expect("column inferred") + } + + #[test] + fn field_coercion_matches_strict_for_every_column_type() { + let coerce = + |v: Value, t: ColumnType| ndb_field_to_value(Some(&v), &ColumnDef::nullable("c", t)); + + assert_eq!( + coerce(Value::Integer(5), ColumnType::String).expect("string"), + Value::String("5".into()) + ); + assert_eq!( + coerce(Value::String("7".into()), ColumnType::Int64).expect("int"), + Value::Integer(7) + ); + assert!(matches!( + coerce(Value::Float(1.5), ColumnType::Int64), + Err(crate::Error::DatatypeMismatch { ref detail }) if detail.contains("'c'") + )); + assert!(matches!( + coerce(Value::Integer(1), ColumnType::Geometry), + Err(crate::Error::DatatypeMismatch { .. }) + )); + assert_eq!( + coerce( + Value::Array(vec![Value::Float(0.5), Value::Float(1.25)]), + ColumnType::Vector(2) + ) + .expect("vector"), + Value::Bytes( + [0.5f32, 1.25f32] + .iter() + .flat_map(|f| f.to_le_bytes()) + .collect() + ) + ); + assert!(matches!( + coerce(Value::Array(vec![Value::Float(0.5)]), ColumnType::Vector(2)), + Err(crate::Error::DataException { .. }) + )); + assert_eq!( + coerce(Value::Integer(9), ColumnType::SystemTimestamp).expect("system ts"), + Value::Integer(9) + ); + assert_eq!( + ndb_field_to_value(None, &ColumnDef::nullable("c", ColumnType::Int64)).expect("absent"), + Value::Null + ); + } + + /// A columnar `SMALLINT` or `REAL` column refuses a value past its + /// declared width, on a new row and on an UPDATE post-image alike. + #[test] + fn declared_width_is_enforced_on_insert_and_update_post_images() { + let small = ColumnDef::nullable("v", ColumnType::Int64).with_declared_width("SMALLINT"); + let real = ColumnDef::nullable("r", ColumnType::Float64).with_declared_width("REAL"); + for (value, column) in [(Value::Integer(40000), &small), (Value::Float(1e39), &real)] { + let err = ndb_field_to_value(Some(&value), column).expect_err("past the width"); + assert!( + matches!(err, crate::Error::NumericValueOutOfRange { .. }), + "{value:?}: {err:?}" + ); + } + + let schema = ColumnarSchema::new(vec![ + ColumnDef::required("id", ColumnType::Int64).with_primary_key(), + small.clone(), + real.clone(), + ]) + .expect("valid schema"); + // A computed `SET v = v + 39999` over a stored `1`. + let mut post_image = vec![Value::Integer(1), Value::Integer(40000), Value::Float(1.5)]; + let err = coerce_columnar_row(&schema, &mut post_image).expect_err("past smallint"); + assert!( + matches!(err, crate::Error::NumericValueOutOfRange { .. }), + "{err:?}" + ); + + let mut fits = vec![Value::Integer(1), Value::Integer(32767), Value::Integer(2)]; + coerce_columnar_row(&schema, &mut fits).expect("fits"); + assert_eq!( + fits, + vec![Value::Integer(1), Value::Integer(32767), Value::Float(2.0)] + ); + } + + /// Every value `ndb_field_to_value` yields is a value the memtable holds. + #[test] + fn coerced_fields_append_to_the_memtable() { + let schema = ColumnarSchema::new(vec![ + ColumnDef::required("s", ColumnType::String), + ColumnDef::required("u", ColumnType::Uuid), + ColumnDef::required("v", ColumnType::Vector(2)), + ColumnDef::required("d", ColumnType::Duration), + ColumnDef::required("b", ColumnType::Bytes), + ]) + .expect("valid schema"); + let input = [ + Value::Integer(5), + Value::String("67e55044-10b1-426f-9247-bb680e5fe0c8".into()), + Value::Array(vec![Value::Float(0.5), Value::Float(1.25)]), + Value::Duration(nodedb_types::NdbDuration::from_micros(10)), + Value::String("AQI=".into()), + ]; + let row: Vec = schema + .columns + .iter() + .zip(input.iter()) + .map(|(col, v)| ndb_field_to_value(Some(v), col)) + .collect::>() + .expect("coerce"); + let mut mt = nodedb_columnar::ColumnarMemtable::new(&schema); + mt.append_row(&row).expect("append"); + assert_eq!( + mt.get_row(0).expect("read"), + Some(vec![ + Value::String("5".into()), + Value::Uuid("67e55044-10b1-426f-9247-bb680e5fe0c8".into()), + Value::Array(vec![Value::Float(0.5), Value::Float(1.25)]), + Value::Integer(10), + Value::Bytes(vec![1, 2]), + ]) + ); + } + + #[test] + fn nested_and_geometry_fields_infer_columns_that_hold_them() { + let row = Value::Object(HashMap::from([ + ("id".to_string(), Value::String("a".into())), + ( + "geom".to_string(), + Value::Geometry(nodedb_types::geometry::Geometry::point(1.0, 2.0)), + ), + ( + "emb".to_string(), + Value::Array(vec![Value::Float(0.5), Value::Float(1.5)]), + ), + ("meta".to_string(), Value::Object(HashMap::new())), + ("raw".to_string(), Value::Bytes(vec![1, 2])), + ])); + let schema = infer_schema_from_value(&row); + assert_eq!(column_type(&schema, "id"), ColumnType::String); + assert_eq!(column_type(&schema, "geom"), ColumnType::Geometry); + assert_eq!(column_type(&schema, "emb"), ColumnType::Json); + assert_eq!(column_type(&schema, "meta"), ColumnType::Json); + assert_eq!(column_type(&schema, "raw"), ColumnType::Bytes); + } } diff --git a/nodedb/src/data/executor/handlers/document/write/register.rs b/nodedb/src/data/executor/handlers/document/write/register.rs index 78c925c03..3d0d67b88 100644 --- a/nodedb/src/data/executor/handlers/document/write/register.rs +++ b/nodedb/src/data/executor/handlers/document/write/register.rs @@ -6,7 +6,7 @@ use tracing::{debug, warn}; -use crate::bridge::envelope::Response; +use crate::bridge::envelope::{ErrorCode, Response}; use crate::data::executor::core_loop::CoreLoop; use crate::data::executor::task::ExecutionTask; @@ -32,6 +32,15 @@ pub(in crate::data::executor) struct RegisterDocumentCollectionParams<'a> { /// The read path decodes this collection's sparse rows as `zerompk` /// TAGGED sidecars solely on the strength of this marker. pub vector_primary: Option<&'a nodedb_types::VectorPrimaryConfig>, + /// Declared `VECTOR(n)` columns of a schemaless collection, as + /// `(column, n)`. Empty for every other collection. + pub vector_fields: &'a [(String, usize)], + /// Declared numeric columns of a schemaless or KV collection. Empty for + /// every other collection. + pub declared_columns: &'a [nodedb_physical::physical_plan::DeclaredColumn], + /// The collection's declared key column. `None` when rows key by the + /// implicit `id` or `_rowid`. + pub declared_key: Option<&'a str>, } impl CoreLoop { @@ -56,6 +65,9 @@ impl CoreLoop { conflict_policy, timeseries, vector_primary, + vector_fields, + declared_columns, + declared_key, } = params; let mode_label = match storage_mode { nodedb_physical::physical_plan::StorageMode::Schemaless => "document_schemaless", @@ -90,43 +102,38 @@ impl CoreLoop { conflict_policy: conflict_policy.map(str::to_string), timeseries: timeseries.map(|ts| Box::new(ts.clone())), vector_primary: vector_primary.map(|vp| Box::new(vp.clone())), + vector_fields: vector_fields.to_vec(), + declared_columns: declared_columns.to_vec(), + declared_key: declared_key.map(str::to_string), }; - let config_key = ( - task.request.database_id, - crate::types::TenantId::new(tid), - collection.to_string(), - ); - self.doc_configs.insert(config_key, config); - // Rehydrate the durable CRDT conflict-resolution policy (if any) into // this core's `PolicyRegistry`. Runs on every `Register` — live DDL // apply AND boot rehydration replay — so `ALTER COLLECTION ... SET ON - // CONFLICT ...` survives a restart instead of silently reverting to - // `CollectionPolicy::ephemeral()`. + // CONFLICT ...` survives a restart. A policy that does not apply + // refuses the register before the collection config is installed. if let Some(policy_json) = conflict_policy { - match self.get_crdt_engine(task.request.database_id, crate::types::TenantId::new(tid)) { - Ok(engine) => { - if let Err(e) = engine.set_collection_policy(collection, policy_json) { - warn!( - core = self.core_id, - %collection, - error = %e, - "failed to rehydrate persisted conflict policy on register" - ); - } - } - Err(e) => { - warn!( - core = self.core_id, - %collection, - error = %e, - "failed to create CRDT engine for conflict policy rehydration" - ); - } + let applied = self + .get_crdt_engine(task.request.database_id, crate::types::TenantId::new(tid)) + .and_then(|engine| engine.set_collection_policy(collection, policy_json)); + if let Err(e) = applied { + warn!( + core = self.core_id, + %collection, + error = %e, + "persisted conflict policy did not apply on register" + ); + return self.response_error(task, ErrorCode::from(e)); } } + let config_key = ( + task.request.database_id, + crate::types::TenantId::new(tid), + collection.to_string(), + ); + self.doc_configs.insert(config_key, config); + self.response_ok(task) } } diff --git a/nodedb/src/data/executor/handlers/kv/conflict_merge.rs b/nodedb/src/data/executor/handlers/kv/conflict_merge.rs index 1c506b2a6..290c7ae28 100644 --- a/nodedb/src/data/executor/handlers/kv/conflict_merge.rs +++ b/nodedb/src/data/executor/handlers/kv/conflict_merge.rs @@ -1,6 +1,7 @@ // SPDX-License-Identifier: BUSL-1.1 -//! The one `INSERT ... ON CONFLICT (key) DO UPDATE SET` merge for a KV row. +//! The one `INSERT ... ON CONFLICT (key) DO UPDATE SET` post-image for a KV +//! row. //! //! The live handler, the transaction resolver, statement staging, and WAL //! replay all compute the post-image here, so a staged value, its durable @@ -9,21 +10,35 @@ //! `kv_body_to_row` and re-encodes in the existing body's shape, so a raw //! row stays raw and RESP `GET` keeps returning the bare value. -use nodedb_physical::physical_plan::UpdateValue; +use nodedb_physical::physical_plan::{DeclaredColumn, UpdateValue}; use nodedb_query::msgpack_scan::{KvBodyError, kv_body_to_row, row_to_kv_body}; +use nodedb_types::Value; +use crate::data::executor::handlers::kv::declared_body::coerce_kv_body; use crate::data::executor::handlers::upsert::apply_on_conflict_updates; - -/// Apply `updates` to the stored `existing` body, with `incoming` as the -/// `EXCLUDED` row, and return the merged body in `existing`'s shape. +use crate::data::executor::strict_format::coerce_declared_row; + +/// The body the upsert stores. With no `existing` row it is `incoming`. +/// Otherwise it is `updates` applied to `existing`, with `incoming` as the +/// `EXCLUDED` row, in `existing`'s shape. +/// +/// Every `declared` column of the stored body holds the value its type +/// stores, in both branches. pub(in crate::data::executor) fn merge_kv_conflict_body( - existing: &[u8], + existing: Option<&[u8]>, incoming: &[u8], updates: &[(String, UpdateValue)], + declared: &[DeclaredColumn], ) -> crate::Result> { + let Some(existing) = existing else { + return Ok(coerce_kv_body(incoming, declared)?.into_owned()); + }; let (existing_row, shape) = kv_body_to_row(existing).map_err(KvBodyError::from)?; let (excluded_row, _) = kv_body_to_row(incoming).map_err(KvBodyError::from)?; - let merged = apply_on_conflict_updates(existing_row, &excluded_row, updates)?; + let mut merged = apply_on_conflict_updates(existing_row, &excluded_row, updates)?; + if let Value::Object(map) = &mut merged { + coerce_declared_row(map, declared)?; + } Ok(row_to_kv_body(&merged, shape)?) } @@ -31,7 +46,7 @@ pub(in crate::data::executor) fn merge_kv_conflict_body( mod tests { use std::collections::HashMap; - use nodedb_physical::physical_plan::UpdateValue; + use nodedb_physical::physical_plan::{DeclaredColumn, UpdateValue}; use nodedb_query::SqlExpr; use nodedb_types::Value; @@ -59,9 +74,13 @@ mod tests { #[test] fn raw_body_overwritten_from_excluded_stays_raw() { - let merged = - merge_kv_conflict_body(b"first", b"second-longer-value", &set_value_from_excluded()) - .expect("merge"); + let merged = merge_kv_conflict_body( + Some(b"first".as_slice()), + b"second-longer-value", + &set_value_from_excluded(), + &[], + ) + .expect("merge"); assert_eq!(merged, b"second-longer-value".to_vec()); } @@ -69,21 +88,24 @@ mod tests { fn raw_single_byte_body_keeps_its_shape() { // 0x31 is a msgpack fixint; the merge must still treat it as the // string "1" and write back "2" as one raw byte. - let merged = merge_kv_conflict_body(b"1", b"2", &set_value_from_excluded()).expect("merge"); + let merged = + merge_kv_conflict_body(Some(b"1".as_slice()), b"2", &set_value_from_excluded(), &[]) + .expect("merge"); assert_eq!(merged, b"2".to_vec()); } #[test] fn raw_body_with_literal_value_assignment_stays_raw() { let updates = vec![("value".to_string(), literal(Value::String("lit".into())))]; - let merged = merge_kv_conflict_body(b"first", b"ignored", &updates).expect("merge"); + let merged = merge_kv_conflict_body(Some(b"first".as_slice()), b"ignored", &updates, &[]) + .expect("merge"); assert_eq!(merged, b"lit".to_vec()); } #[test] fn raw_body_refuses_a_typed_column_assignment() { let updates = vec![("n".to_string(), literal(Value::Integer(1)))]; - let err = merge_kv_conflict_body(b"first", b"second", &updates) + let err = merge_kv_conflict_body(Some(b"first".as_slice()), b"second", &updates, &[]) .expect_err("a raw row cannot grow a typed column"); assert!( matches!(err, crate::Error::BadRequest { .. }), @@ -95,9 +117,13 @@ mod tests { #[test] fn map_body_merges_and_stays_a_map() { let updates = vec![("mana".to_string(), literal(Value::Integer(5)))]; - let merged = - merge_kv_conflict_body(&map_body(&[("hp", 10)]), &map_body(&[("hp", 1)]), &updates) - .expect("merge"); + let merged = merge_kv_conflict_body( + Some(map_body(&[("hp", 10)]).as_slice()), + &map_body(&[("hp", 1)]), + &updates, + &[], + ) + .expect("merge"); let row = nodedb_types::value_from_msgpack(&merged).expect("map body"); assert_eq!(row.get("hp"), Some(&Value::Integer(10))); assert_eq!(row.get("mana"), Some(&Value::Integer(5))); @@ -109,10 +135,47 @@ mod tests { "hp".to_string(), UpdateValue::Expr(SqlExpr::ExcludedColumn("hp".to_string())), )]; - let merged = - merge_kv_conflict_body(&map_body(&[("hp", 10)]), &map_body(&[("hp", 1)]), &updates) - .expect("merge"); + let merged = merge_kv_conflict_body( + Some(map_body(&[("hp", 10)]).as_slice()), + &map_body(&[("hp", 1)]), + &updates, + &[], + ) + .expect("merge"); let row = nodedb_types::value_from_msgpack(&merged).expect("map body"); assert_eq!(row.get("hp"), Some(&Value::Integer(1))); } + + /// A computed assignment past a declared `SMALLINT` is refused, and the + /// insert branch re-types the incoming row. + #[test] + fn declared_columns_are_retyped_in_both_branches() { + let declared = [DeclaredColumn::from_declared("hp", "SMALLINT").expect("smallint")]; + let updates = vec![( + "hp".to_string(), + UpdateValue::Expr(SqlExpr::BinaryOp { + left: Box::new(SqlExpr::Column("hp".to_string())), + op: nodedb_query::BinaryOp::Add, + right: Box::new(SqlExpr::Literal(Value::Integer(39999))), + }), + )]; + let err = merge_kv_conflict_body( + Some(map_body(&[("hp", 10)]).as_slice()), + &map_body(&[("hp", 1)]), + &updates, + &declared, + ) + .expect_err("39999 + 10 is past smallint"); + assert!( + matches!(err, crate::Error::NumericValueOutOfRange { .. }), + "{err:?}" + ); + + let err = merge_kv_conflict_body(None, &map_body(&[("hp", 40000)]), &updates, &declared) + .expect_err("the insert row is past smallint"); + assert!( + matches!(err, crate::Error::NumericValueOutOfRange { .. }), + "{err:?}" + ); + } } diff --git a/nodedb/src/data/executor/handlers/kv/crud/write_upsert.rs b/nodedb/src/data/executor/handlers/kv/crud/write_upsert.rs index 80a8b3802..0e129d757 100644 --- a/nodedb/src/data/executor/handlers/kv/crud/write_upsert.rs +++ b/nodedb/src/data/executor/handlers/kv/crud/write_upsert.rs @@ -48,18 +48,14 @@ impl CoreLoop { let now_ms = self.kv_ttl_now_ms(task); let existing_bytes = self.kv_engine.get(did, tid, collection, key, now_ms); - let stored_bytes: Vec = match &existing_bytes { - None => value.to_vec(), - Some(existing_raw) => { - match super::super::conflict_merge::merge_kv_conflict_body( - existing_raw, - value, - updates, - ) { - Ok(b) => b, - Err(e) => return self.response_error(task, e), - } - } + let stored_bytes: Vec = match super::super::conflict_merge::merge_kv_conflict_body( + existing_bytes.as_deref(), + value, + updates, + self.declared_columns_of(did, tid, collection), + ) { + Ok(b) => b, + Err(e) => return self.response_error(task, e), }; // `stored_bytes` is whichever body this op actually persists — the diff --git a/nodedb/src/data/executor/handlers/kv/declared_body.rs b/nodedb/src/data/executor/handlers/kv/declared_body.rs new file mode 100644 index 000000000..c1a9ae638 --- /dev/null +++ b/nodedb/src/data/executor/handlers/kv/declared_body.rs @@ -0,0 +1,277 @@ +// SPDX-License-Identifier: BUSL-1.1 + +//! Re-typing a KV row body to its collection's declared numeric columns. +//! +//! A body keeps its shape: a map stays a map, and a raw `value` body stays +//! raw. The rule per value is `strict_format::coerce_declared`, the one the +//! schemaless document path runs. + +use std::borrow::Cow; + +use nodedb_physical::physical_plan::{DeclaredColumn, KvOp}; +use nodedb_query::msgpack_scan::{KvBodyError, kv_body_to_row, row_to_kv_body}; +use nodedb_types::Value; + +use crate::data::executor::core_loop::CoreLoop; +use crate::data::executor::strict_format::coerce_declared; +use crate::engine::kv::AtomicError; + +/// `body` with every declared column re-typed. Borrowed when no value +/// changes. +pub(in crate::data::executor) fn coerce_kv_body<'a>( + body: &'a [u8], + declared: &[DeclaredColumn], +) -> crate::Result> { + if declared.is_empty() { + return Ok(Cow::Borrowed(body)); + } + let (row, shape) = kv_body_to_row(body).map_err(KvBodyError::from)?; + let Value::Object(mut map) = row else { + return Ok(Cow::Borrowed(body)); + }; + let mut changed = false; + for column in declared { + let Some(slot) = map.get_mut(&column.name) else { + continue; + }; + let coerced = coerce_declared(slot, column)?; + if coerced != *slot { + *slot = coerced; + changed = true; + } + } + if !changed { + return Ok(Cow::Borrowed(body)); + } + Ok(Cow::Owned(row_to_kv_body(&Value::Object(map), shape)?)) +} + +/// The image a KV write stores for the image it computed: `None` when every +/// declared column already holds the value it stores, else the fitted image. +/// +/// INCR, INCRBYFLOAT, CAS, and GETSET compute the bytes they store, and +/// TRANSFER_ITEM moves a row into a collection with its own declared columns. +/// Each one runs this on those bytes before it stores them, on every path: +/// autocommit, transaction staging, resolve, and WAL replay. +pub(in crate::data::executor) fn fit_kv_image( + image: &[u8], + declared: &[DeclaredColumn], +) -> crate::Result>> { + Ok(match coerce_kv_body(image, declared)? { + Cow::Owned(bytes) => Some(bytes), + Cow::Borrowed(_) => None, + }) +} + +/// The atomic gate of a live write: fit the computed `image` to `declared`, +/// then decide the image it stores against the write policy. +pub(in crate::data::executor) fn fit_and_admit_kv_image( + image: &[u8], + declared: &[DeclaredColumn], + rls_write_check: &nodedb_types::RlsWriteCheck, + key: &[u8], + tid: u64, + collection: &str, +) -> Result>, AtomicError> { + let fitted = + fit_kv_image(image, declared).map_err(|error| AtomicError::Declared(Box::new(error)))?; + super::rls::admit_kv_row( + rls_write_check, + fitted.as_deref().unwrap_or(image), + key, + tid, + collection, + ) + .map_err(|error| AtomicError::Rejected(Box::new(error)))?; + Ok(fitted) +} + +/// The atomic gate of a WAL redo: fit the computed `image` to `declared`. +/// +/// A redo re-applies a write whose policy verdict was reached when it was +/// first accepted, so no policy is decided here. The declared rule is part of +/// the computation, so the redo stores the bytes the live write stored, and a +/// value the live write refused is refused again. +pub(in crate::data::executor) fn fit_replayed_kv_image( + image: &[u8], + declared: &[DeclaredColumn], +) -> Result>, AtomicError> { + fit_kv_image(image, declared).map_err(|error| AtomicError::Declared(Box::new(error))) +} + +impl CoreLoop { + /// `op` with every row body it supplies whole re-typed to its + /// collection's declared numeric columns. `None` when no body changes. + /// + /// Covers `Put`, `Insert`, `InsertIfAbsent`, and `BatchPut`. A field + /// merge re-types its merged row in `merge_field_updates`, a conflict + /// upsert re-types both its branches in `merge_kv_conflict_body`, and an + /// atomic or a transfer fits the image it computes. + pub(in crate::data::executor) fn coerce_kv_op_bodies( + &self, + did: u64, + tid: u64, + op: &KvOp, + ) -> crate::Result> { + let collection = match op { + KvOp::Put { collection, .. } + | KvOp::Insert { collection, .. } + | KvOp::InsertIfAbsent { collection, .. } + | KvOp::BatchPut { collection, .. } => collection.as_str(), + // These ops carry no row body supplied whole. A field merge and a + // conflict upsert re-type the row they compute, in the merge. An + // atomic and a transfer fit the image they compute through + // `fit_kv_image` or `compute_transfer`. + KvOp::Get { .. } + | KvOp::InsertOnConflictUpdate { .. } + | KvOp::Delete { .. } + | KvOp::Scan { .. } + | KvOp::Expire { .. } + | KvOp::Persist { .. } + | KvOp::GetTtl { .. } + | KvOp::BatchGet { .. } + | KvOp::RegisterIndex { .. } + | KvOp::DropIndex { .. } + | KvOp::FieldGet { .. } + | KvOp::FieldSet { .. } + | KvOp::Truncate { .. } + | KvOp::Incr { .. } + | KvOp::IncrFloat { .. } + | KvOp::Cas { .. } + | KvOp::GetSet { .. } + | KvOp::Transfer { .. } + | KvOp::TransferItem { .. } + | KvOp::RegisterSortedIndex { .. } + | KvOp::DropSortedIndex { .. } + | KvOp::SortedIndexRank { .. } + | KvOp::SortedIndexTopK { .. } + | KvOp::SortedIndexRange { .. } + | KvOp::SortedIndexCount { .. } + | KvOp::SortedIndexScore { .. } + | KvOp::SortedIndexTxnRead { .. } + | KvOp::MaterializeScan { .. } + | KvOp::ResolveWrite(_) + | KvOp::ResolvedWrite { .. } + | KvOp::PredicateUpdate { .. } + | KvOp::PredicateDelete { .. } => return Ok(None), + }; + let declared = self.declared_columns_of(did, tid, collection); + if declared.is_empty() { + return Ok(None); + } + + if let KvOp::BatchPut { entries, .. } = op { + let mut coerced: Vec>> = Vec::with_capacity(entries.len()); + for (_, value) in entries { + coerced.push(match coerce_kv_body(value, declared)? { + Cow::Owned(bytes) => Some(bytes), + Cow::Borrowed(_) => None, + }); + } + if coerced.iter().all(Option::is_none) { + return Ok(None); + } + let mut rewritten = op.clone(); + if let KvOp::BatchPut { entries, .. } = &mut rewritten { + for ((_, value), bytes) in entries.iter_mut().zip(coerced) { + if let Some(bytes) = bytes { + *value = bytes; + } + } + } + return Ok(Some(rewritten)); + } + + let (KvOp::Put { value, .. } + | KvOp::Insert { value, .. } + | KvOp::InsertIfAbsent { value, .. }) = op + else { + return Ok(None); + }; + let Cow::Owned(bytes) = coerce_kv_body(value, declared)? else { + return Ok(None); + }; + let mut rewritten = op.clone(); + if let KvOp::Put { value, .. } + | KvOp::Insert { value, .. } + | KvOp::InsertIfAbsent { value, .. } = &mut rewritten + { + *value = bytes; + } + Ok(Some(rewritten)) + } +} + +#[cfg(test)] +mod tests { + use std::collections::HashMap; + + use super::*; + + fn declared(declared: &str) -> Vec { + vec![DeclaredColumn::from_declared("v", declared).expect("numeric declaration")] + } + + fn map_body(value: Value) -> Vec { + let mut map = HashMap::new(); + map.insert("v".to_string(), value); + map.insert("note".to_string(), Value::String("x".into())); + nodedb_types::value_to_msgpack(&Value::Object(map)).expect("encode body") + } + + #[test] + fn map_body_is_retyped_and_stays_a_map() { + let body = map_body(Value::String("1.005".into())); + let coerced = coerce_kv_body(&body, &declared("DECIMAL(5,2)")).expect("fits"); + let row = nodedb_types::value_from_msgpack(&coerced).expect("map body"); + assert_eq!(row.get("v"), Some(&Value::String("1.01".into()))); + assert_eq!(row.get("note"), Some(&Value::String("x".into()))); + } + + #[test] + fn unchanged_body_is_borrowed() { + let body = map_body(Value::String("1.50".into())); + assert!(matches!( + coerce_kv_body(&body, &declared("DECIMAL(5,2)")).expect("fits"), + Cow::Borrowed(_) + )); + } + + #[test] + fn raw_value_body_stays_raw() { + let declared = + vec![DeclaredColumn::from_declared("value", "DECIMAL(5,2)").expect("decimal")]; + let coerced = coerce_kv_body(b"1.005", &declared).expect("fits"); + assert_eq!(coerced.as_ref(), b"1.01"); + } + + #[test] + fn computed_image_is_fitted_or_refused() { + let declared = + vec![DeclaredColumn::from_declared("value", "DECIMAL(5,2)").expect("decimal")]; + assert_eq!( + fit_kv_image(b"1.505", &declared).expect("fits"), + Some(b"1.51".to_vec()) + ); + assert_eq!(fit_kv_image(b"1.50", &declared).expect("unchanged"), None); + let err = fit_replayed_kv_image(b"1000.99", &declared).expect_err("past precision"); + assert!( + matches!( + err, + AtomicError::Declared(ref e) + if matches!(**e, crate::Error::NumericValueOutOfRange { .. }) + ), + "{err:?}" + ); + } + + #[test] + fn out_of_range_body_is_refused() { + let err = coerce_kv_body(&map_body(Value::Integer(40000)), &declared("SMALLINT")) + .expect_err("past smallint"); + assert!( + matches!(err, crate::Error::NumericValueOutOfRange { .. }), + "{err:?}" + ); + } +} diff --git a/nodedb/src/data/executor/handlers/kv/field_compute.rs b/nodedb/src/data/executor/handlers/kv/field_compute.rs index 182880aec..8f94a56c2 100644 --- a/nodedb/src/data/executor/handlers/kv/field_compute.rs +++ b/nodedb/src/data/executor/handlers/kv/field_compute.rs @@ -7,10 +7,12 @@ //! mirrors the `nodedb_physical::kv_atomic::compute` / `stage_kv_atomic` split for //! `Incr`/`Cas`/etc. +use nodedb_physical::physical_plan::DeclaredColumn; use nodedb_query::msgpack_scan::{KvBodyError, KvBodyShape, kv_body_to_row, row_to_kv_body}; use nodedb_types::Value; use crate::bridge::envelope::ErrorCode; +use crate::data::executor::strict_format::coerce_declared_row; /// Result of merging field updates into a KV document body. #[derive(Debug)] @@ -27,10 +29,14 @@ pub(in crate::data::executor) struct FieldSetComputation { /// A raw scalar body (the single-`value` SQL form, RESP `SET`) is not a /// hash: a field set against it is a type mismatch, the verdict Redis gives /// `HSET` on a string key. It is never silently replaced by a map. +/// +/// Every `declared` column of the merged row holds the value its type +/// stores, so a value out of the declared range fails the write. pub(in crate::data::executor) fn merge_field_updates( collection: &str, current: Option<&[u8]>, updates: &[(String, Vec)], + declared: &[DeclaredColumn], ) -> Result { let mut doc = match current { None => std::collections::HashMap::new(), @@ -69,6 +75,7 @@ pub(in crate::data::executor) fn merge_field_updates( } doc.insert(field.clone(), new_value); } + coerce_declared_row(&mut doc, declared).map_err(ErrorCode::from)?; let new_value = row_to_kv_body(&Value::Object(doc), KvBodyShape::Map) .map_err(|e| ErrorCode::from(crate::Error::from(e)))?; @@ -93,7 +100,8 @@ mod tests { #[test] fn merges_into_empty_document() { - let result = merge_field_updates("c", None, &[("score".to_string(), int(42))]).unwrap(); + let result = + merge_field_updates("c", None, &[("score".to_string(), int(42))], &[]).unwrap(); assert_eq!(result.fields_added, 1); assert_eq!( decode(&result.new_value).get("score"), @@ -109,7 +117,8 @@ mod tests { nodedb_types::value_to_msgpack(&Value::Object(m)).unwrap() }; let result = - merge_field_updates("c", Some(&existing), &[("score".to_string(), int(2))]).unwrap(); + merge_field_updates("c", Some(&existing), &[("score".to_string(), int(2))], &[]) + .unwrap(); assert_eq!(result.fields_added, 0); assert_eq!( decode(&result.new_value).get("score"), @@ -119,14 +128,29 @@ mod tests { #[test] fn empty_value_bytes_set_null() { - let result = merge_field_updates("c", None, &[("f".to_string(), Vec::new())]).unwrap(); + let result = merge_field_updates("c", None, &[("f".to_string(), Vec::new())], &[]).unwrap(); assert_eq!(decode(&result.new_value).get("f"), Some(&Value::Null)); } + #[test] + fn declared_column_is_retyped_and_range_checked() { + let declared = [DeclaredColumn::from_declared("n", "SMALLINT").expect("smallint")]; + let result = merge_field_updates("c", None, &[("n".to_string(), int(7))], &declared) + .expect("fits smallint"); + assert_eq!(decode(&result.new_value).get("n"), Some(&Value::Integer(7))); + + let err = merge_field_updates("c", None, &[("n".to_string(), int(40000))], &declared) + .expect_err("past smallint"); + assert!( + matches!(err, ErrorCode::NumericValueOutOfRange { .. }), + "{err:?}" + ); + } + #[test] fn raw_body_is_a_type_mismatch_not_replaced_by_a_map() { for body in [b"first".as_slice(), b"1".as_slice()] { - let err = merge_field_updates("c", Some(body), &[("f".to_string(), int(1))]) + let err = merge_field_updates("c", Some(body), &[("f".to_string(), int(1))], &[]) .expect_err("HSET on a bare-value key must be refused"); match err { ErrorCode::TypeMismatch { collection, .. } => assert_eq!(collection, "c"), @@ -138,7 +162,7 @@ mod tests { #[test] fn corrupt_map_body_is_an_error_not_an_empty_document() { // fixmap header claiming one entry, then nothing. - let err = merge_field_updates("c", Some(&[0x81]), &[("f".to_string(), int(1))]) + let err = merge_field_updates("c", Some(&[0x81]), &[("f".to_string(), int(1))], &[]) .expect_err("a truncated body must not merge onto an empty map"); assert!(!matches!(err, ErrorCode::TypeMismatch { .. }), "{err:?}"); } diff --git a/nodedb/src/data/executor/handlers/kv/mod.rs b/nodedb/src/data/executor/handlers/kv/mod.rs index 72f260c03..5d2ad40c3 100644 --- a/nodedb/src/data/executor/handlers/kv/mod.rs +++ b/nodedb/src/data/executor/handlers/kv/mod.rs @@ -6,6 +6,7 @@ pub(in crate::data::executor) mod atomic; pub(in crate::data::executor) mod batch; pub(in crate::data::executor) mod conflict_merge; pub(in crate::data::executor) mod crud; +pub(in crate::data::executor) mod declared_body; mod dispatch; mod dispatch_scan; mod dispatch_transfer; diff --git a/nodedb/src/data/executor/handlers/kv/resolve/atomic_ops.rs b/nodedb/src/data/executor/handlers/kv/resolve/atomic_ops.rs index 0b981364a..4ff2a57aa 100644 --- a/nodedb/src/data/executor/handlers/kv/resolve/atomic_ops.rs +++ b/nodedb/src/data/executor/handlers/kv/resolve/atomic_ops.rs @@ -14,8 +14,10 @@ use crate::data::executor::core_loop::CoreLoop; use crate::data::executor::handlers::kv::atomic::{ KvAtomicCtx, atomic_error_code, incr_float_reply, }; +use crate::data::executor::handlers::kv::declared_body::fit_kv_image; use crate::data::executor::handlers::kv::rls::admit_kv_row; use crate::data::executor::response_codec; +use crate::engine::kv::fitted_counter_f64; /// Render a stored body for the `current_value` / `old_value` slot of an /// atomic's reply, exactly as the live handlers do. @@ -47,8 +49,12 @@ impl CoreLoop { } let now_ms = self.kv_ttl_now_ms(task); let current = self.kv_resolve_read(did, tid, collection, key, now_ms); - let (new_value, new_bytes) = compute::incr(current.as_deref(), delta, shape) + let (new_value, computed) = compute::incr(current.as_deref(), delta, shape) .map_err(|e| atomic_error_code(e.into(), collection))?; + // The declared column rule fits the computed row, as the live engine + // gate does. It never changes an integer it accepts. + let declared = self.declared_columns_of(did, tid, collection); + let new_bytes = fit_kv_image(&computed, declared)?.unwrap_or(computed); admit_kv_row(rls_write_check, &new_bytes, key, tid, collection)?; let expire_at_ms = if ttl_ms > 0 { @@ -94,8 +100,18 @@ impl CoreLoop { } let now_ms = self.kv_read_now_ms(); let current = self.kv_resolve_read(did, tid, collection, key, now_ms); - let (new_value, new_bytes) = compute::incr_float(current.as_deref(), delta, shape) + let (computed_value, computed) = compute::incr_float(current.as_deref(), delta, shape) .map_err(|e| atomic_error_code(e.into(), collection))?; + // The declared column rule fits the computed row, as the live engine + // gate does, and the reply carries the value the fitted row stores. + let declared = self.declared_columns_of(did, tid, collection); + let (new_value, new_bytes) = match fit_kv_image(&computed, declared)? { + None => (computed_value, computed), + Some(fitted) => ( + fitted_counter_f64(computed_value, &computed, &fitted), + fitted, + ), + }; admit_kv_row(rls_write_check, &new_bytes, key, tid, collection)?; let response_payload = @@ -137,13 +153,19 @@ impl CoreLoop { } let now_ms = self.kv_read_now_ms(); let current = self.kv_resolve_read(did, tid, collection, key, now_ms); - let (matches, write_bytes) = compute::cas(current.as_deref(), expected, new_value) + let (matches, computed) = compute::cas(current.as_deref(), expected, new_value) .map_err(|e| atomic_error_code(e.into(), collection))?; - // Decided on the image the swap stores, same as `execute_kv_cas`: a - // swap into a typed row stores the row, not `new_value` itself. - if matches { - admit_kv_row(rls_write_check, &write_bytes, key, tid, collection)?; - } + // Fitted and decided on the image the swap stores, same as + // `execute_kv_cas`: a swap into a typed row stores the row, not + // `new_value` itself. + let write_bytes = if matches { + let declared = self.declared_columns_of(did, tid, collection); + let fitted = fit_kv_image(&computed, declared)?.unwrap_or(computed); + admit_kv_row(rls_write_check, &fitted, key, tid, collection)?; + fitted + } else { + computed + }; let response_payload = response_codec::encode_json_as_msgpack(&serde_json::json!({ "success": matches, @@ -192,9 +214,12 @@ impl CoreLoop { } let now_ms = self.kv_read_now_ms(); let old = self.kv_resolve_read(did, tid, collection, key, now_ms); - let write_bytes = compute::getset(old.as_deref(), new_value) + let computed = compute::getset(old.as_deref(), new_value) .map_err(|e| atomic_error_code(e.into(), collection))?; - // Decided on the image the write stores, same as `execute_kv_getset`. + // Fitted and decided on the image the write stores, same as + // `execute_kv_getset`. + let declared = self.declared_columns_of(did, tid, collection); + let write_bytes = fit_kv_image(&computed, declared)?.unwrap_or(computed); admit_kv_row(rls_write_check, &write_bytes, key, tid, collection)?; let disclosable_old = match &old { @@ -202,9 +227,7 @@ impl CoreLoop { Ok(true) => old.as_deref(), Ok(false) => None, Err(e) => { - return Err(ErrorCode::Internal { - detail: e.to_string(), - }); + return Err(ErrorCode::from(e)); } }, None => None, diff --git a/nodedb/src/data/executor/handlers/kv/resolve/predicate_ops.rs b/nodedb/src/data/executor/handlers/kv/resolve/predicate_ops.rs index 526deb231..a86e3445f 100644 --- a/nodedb/src/data/executor/handlers/kv/resolve/predicate_ops.rs +++ b/nodedb/src/data/executor/handlers/kv/resolve/predicate_ops.rs @@ -36,9 +36,11 @@ impl CoreLoop { let now_ms = current_ms(); let matched = self.kv_predicate_matches(did, tid, collection, filters, now_ms)?; + let declared = self.declared_columns_of(did, tid, collection); let mut writes: Vec<(Vec, Vec, Vec)> = Vec::with_capacity(matched.len()); for (key, body) in matched { - let computed = merge_field_updates(collection, Some(body.as_slice()), updates)?; + let computed = + merge_field_updates(collection, Some(body.as_slice()), updates, declared)?; admit_kv_row(rls_write_check, &computed.new_value, &key, tid, collection)?; writes.push((key, body, computed.new_value)); } diff --git a/nodedb/src/data/executor/handlers/kv/resolve/transfer_ops.rs b/nodedb/src/data/executor/handlers/kv/resolve/transfer_ops.rs index 8b5ee0196..74d4b74a6 100644 --- a/nodedb/src/data/executor/handlers/kv/resolve/transfer_ops.rs +++ b/nodedb/src/data/executor/handlers/kv/resolve/transfer_ops.rs @@ -10,6 +10,7 @@ use nodedb_physical::physical_plan::KvResolveOutcome; use super::context::{ResolveResult, ResolvedPut, delete_mutation, put_mutation}; use crate::bridge::envelope::ErrorCode; use crate::data::executor::core_loop::CoreLoop; +use crate::data::executor::handlers::kv::declared_body::fit_kv_image; use crate::data::executor::handlers::kv::rls::admit_kv_row; use crate::data::executor::handlers::kv::transfer::{TransferItemParams, TransferParams}; use crate::data::executor::handlers::kv::transfer_compute::{TransferError, compute_transfer}; @@ -50,8 +51,10 @@ impl CoreLoop { dest_bytes.as_deref().filter(|b| !b.is_empty()), field, amount, + self.declared_columns_of(did, tid, collection), ) .map_err(|e| match e { + TransferError::Declared(error) => ErrorCode::from(error), TransferError::TypeMismatch(detail) => ErrorCode::TypeMismatch { collection: collection.to_string(), detail, @@ -107,9 +110,9 @@ impl CoreLoop { "source_key": String::from_utf8_lossy(source_key), "dest_key": String::from_utf8_lossy(dest_key), "field": field, - "amount": amount, - "source_balance": computed.source_balance_after, - "dest_balance": computed.dest_balance_after, + "amount": computed.amount.to_json(), + "source_balance": computed.source_balance_after.to_json(), + "dest_balance": computed.dest_balance_after.to_json(), }))?; Ok(KvResolveOutcome { mutations, @@ -141,6 +144,13 @@ impl CoreLoop { else { return Err(ErrorCode::NotFound); }; + // The row arriving at the destination meets the destination's + // declared numeric columns. + let dest_data = fit_kv_image( + &item_data, + self.declared_columns_of(did, tid, dest_collection), + )? + .unwrap_or_else(|| item_data.clone()); admit_kv_row( source_rls_write_check, &item_data, @@ -150,7 +160,7 @@ impl CoreLoop { )?; admit_kv_row( dest_rls_write_check, - &item_data, + &dest_data, dest_key, tid, dest_collection, @@ -169,11 +179,11 @@ impl CoreLoop { Ok(KvResolveOutcome { mutations: vec![ - delete_mutation(source_collection, item_key, Some(item_data.clone())), + delete_mutation(source_collection, item_key, Some(item_data)), put_mutation(ResolvedPut { collection: dest_collection, key: dest_key, - value: item_data, + value: dest_data, ttl_ms: 0, expire_at_ms: 0, surrogate, diff --git a/nodedb/src/data/executor/handlers/kv/resolve/write_ops.rs b/nodedb/src/data/executor/handlers/kv/resolve/write_ops.rs index cb3dcfe6b..9b65b2e5a 100644 --- a/nodedb/src/data/executor/handlers/kv/resolve/write_ops.rs +++ b/nodedb/src/data/executor/handlers/kv/resolve/write_ops.rs @@ -56,10 +56,12 @@ impl CoreLoop { let now_ms = self.kv_ttl_now_ms(task); let existing_bytes = self.kv_resolve_read(did, tid, collection, key, now_ms); - let stored_bytes: Vec = match &existing_bytes { - None => value.to_vec(), - Some(existing_raw) => merge_kv_conflict_body(existing_raw, value, updates)?, - }; + let stored_bytes = merge_kv_conflict_body( + existing_bytes.as_deref(), + value, + updates, + self.declared_columns_of(did, tid, collection), + )?; admit_kv_row(rls_write_check, &stored_bytes, key, tid, collection)?; @@ -241,6 +243,7 @@ impl CoreLoop { collection, current.as_deref(), updates, + self.declared_columns_of(did, tid, collection), )?; admit_kv_row(rls_write_check, &computed.new_value, key, tid, collection)?; diff --git a/nodedb/src/data/executor/handlers/point/apply_put/stored_body.rs b/nodedb/src/data/executor/handlers/point/apply_put/stored_body.rs index ac68a12e3..2df229c23 100644 --- a/nodedb/src/data/executor/handlers/point/apply_put/stored_body.rs +++ b/nodedb/src/data/executor/handlers/point/apply_put/stored_body.rs @@ -91,6 +91,10 @@ impl CoreLoop { doc_format::canonicalize_document_for_storage(value) }; + // A declared numeric column holds the value its type stores. Runs + // after the generated columns, so a generated value is re-typed too. + let value = strict_format::coerce_declared_body(value, self.declared_columns(config_key))?; + // Strict (Binary Tuple) pipeline: inject an auto-generated `_rowid` // from the surrogate if the schema declares one and the client // payload lacks it, then encode into Binary Tuple. @@ -106,15 +110,19 @@ impl CoreLoop { let value = match link_after { Some(head) => { // A chained row carries its identity inside the linked - // contents: a minted row gets `id` from its surrogate here, - // as a strict row gets `_rowid` below. A copy under a new - // surrogate then keeps both its id and its link. - // A body that is not an object is left for the link to - // refuse. + // contents: a minted row gets its identity column from its + // surrogate here, as a strict row gets `_rowid` below. A + // copy under a new surrogate then keeps both its identity + // and its link. A declared-key row holds its key and gains + // no `id`. A body that is not an object is left for the + // link to refuse. let value = if nodedb_query::msgpack_scan::map_header(&value, 0).is_some() { nodedb_query::msgpack_scan::inject_str_field( &value, - nodedb_types::DEFAULT_IDENTITY_COLUMN, + config + .declared_key + .as_deref() + .unwrap_or(nodedb_types::DEFAULT_IDENTITY_COLUMN), &surrogate.as_u32().to_string(), ) } else { diff --git a/nodedb/src/data/executor/handlers/point/apply_put/vector/fields.rs b/nodedb/src/data/executor/handlers/point/apply_put/vector/fields.rs index 0a17aacbf..5cdfc23bc 100644 --- a/nodedb/src/data/executor/handlers/point/apply_put/vector/fields.rs +++ b/nodedb/src/data/executor/handlers/point/apply_put/vector/fields.rs @@ -1,10 +1,14 @@ // SPDX-License-Identifier: BUSL-1.1 //! Which fields of a collection carry vectors: strict-schema `Vector(dim)` -//! columns and schemaless fields registered via `vector_params`. +//! columns, schemaless fields registered via `vector_params`, and the +//! declared `VECTOR(n)` columns of a schemaless collection. use crate::data::executor::core_loop::CoreLoop; +/// The field a collection-level vector index (no field name) covers. +pub(super) const DEFAULT_VECTOR_FIELD: &str = "embedding"; + impl CoreLoop { /// Strict-schema `Vector(dim)` column names + dims declared on /// `collection`, or empty if the collection has no strict schema / no @@ -54,11 +58,11 @@ impl CoreLoop { .unwrap_or_default() } - /// Schemaless vector field names registered via `vector_params` for - /// `collection` (named-field entries `"{collection}:{field}"`, plus the - /// bare `"{collection}"` key defaulting to `"embedding"`). Shared by the - /// put path's schemaless indexing branch and the delete cleanup's exact - /// key construction. + /// Schemaless vector field names of `collection`: the named-field + /// `vector_params` entries `"{collection}:{field}"`, the bare + /// `"{collection}"` key defaulting to `"embedding"`, and the collection's + /// declared `VECTOR(n)` columns. Shared by the put path's schemaless + /// indexing branch and the delete cleanup's exact key construction. pub(in crate::data::executor) fn schemaless_vector_field_names( &self, database_id: u64, @@ -79,13 +83,43 @@ impl CoreLoop { .map(|k| k.2[field_prefix.len()..].to_string()) .collect(); if names.is_empty() && self.vector_params.contains_key(&bare_key) { - names.push("embedding".to_string()); + names.push(DEFAULT_VECTOR_FIELD.to_string()); + } + if let Some(config) = self.doc_configs.get(&bare_key) { + for (field, _dim) in &config.vector_fields { + if !names.contains(field) { + names.push(field.clone()); + } + } } names } + /// The width of `field` when `collection` declares it a `VECTOR(n)` + /// column of a schemaless collection, else `None`. + pub(in crate::data::executor) fn declared_schemaless_vector_dim( + &self, + database_id: u64, + tid: u64, + collection: &str, + field: &str, + ) -> Option { + let config_key = ( + nodedb_types::DatabaseId::new(database_id), + crate::types::TenantId::new(tid), + collection.to_string(), + ); + self.doc_configs + .get(&config_key)? + .vector_fields + .iter() + .find(|(name, _dim)| name == field) + .map(|(_name, dim)| *dim) + } + /// Whether `collection` has any vector fields — strict-schema `Vector(dim)` - /// columns OR schemaless fields registered via `vector_params`. Combines + /// columns OR schemaless fields registered via `vector_params` or declared + /// as `VECTOR(n)` columns. Combines /// `strict_vector_fields` + `schemaless_vector_field_names` into the single /// gate check callers need before deciding whether to pay for HNSW /// maintenance at all. Callers that loop over many rows (bulk update/ diff --git a/nodedb/src/data/executor/handlers/point/apply_put/vector/mod.rs b/nodedb/src/data/executor/handlers/point/apply_put/vector/mod.rs index 909439731..fc4f199f0 100644 --- a/nodedb/src/data/executor/handlers/point/apply_put/vector/mod.rs +++ b/nodedb/src/data/executor/handlers/point/apply_put/vector/mod.rs @@ -1,8 +1,9 @@ // SPDX-License-Identifier: BUSL-1.1 //! HNSW vector-index side-effects for `apply_point_put`: index declared -//! strict-schema `Vector(dim)` columns and schemaless `vector_params` -//! fields, and soft-delete a document's prior vector nodes. +//! strict-schema `Vector(dim)` columns, schemaless `vector_params` fields and +//! declared schemaless `VECTOR(n)` columns, and soft-delete a document's prior +//! vector nodes. mod fields; mod put; diff --git a/nodedb/src/data/executor/handlers/point/update/post_image.rs b/nodedb/src/data/executor/handlers/point/update/post_image.rs index d277e3759..005b1ede1 100644 --- a/nodedb/src/data/executor/handlers/point/update/post_image.rs +++ b/nodedb/src/data/executor/handlers/point/update/post_image.rs @@ -9,6 +9,7 @@ use crate::bridge::envelope::ErrorCode; use crate::data::executor::core_loop::CoreLoop; use crate::data::executor::doc_format; +use crate::data::executor::handlers::identity_guard::{IdentitySnapshot, assigns_identity}; use crate::data::executor::strict_format; use crate::types::{DatabaseId, TenantId}; use nodedb_physical::physical_plan::UpdateValue; @@ -140,8 +141,13 @@ impl CoreLoop { } } - // Fast path: non-strict, no generated columns, all literal — merge at binary level. - if !is_strict && !has_generated && !has_expr { + // Fast path: non-strict, no generated columns, all literal, identity + // untouched — merge at binary level. + if !is_strict + && !has_generated + && !has_expr + && !assigns_identity(None, declared_primary_key, updates) + { let base_mp = doc_format::json_to_msgpack(current_bytes); let update_pairs: Vec<(&str, &[u8])> = updates .iter() @@ -150,29 +156,37 @@ impl CoreLoop { UpdateValue::Expr(_) => None, }) .collect(); - return Ok(PointUpdateBody::Msgpack( - nodedb_query::msgpack_scan::merge_fields(&base_mp, &update_pairs), - )); + let merged = nodedb_query::msgpack_scan::merge_fields(&base_mp, &update_pairs); + let merged = + strict_format::coerce_declared_body(merged, self.declared_columns(config_key)) + .map_err(ErrorCode::from)?; + return Ok(PointUpdateBody::Msgpack(merged)); } - // Strict, generated, or expression RHS: decode → mutate → re-encode. - let mut doc = if is_strict { - if let Some(config) = self.doc_configs.get(config_key) - && let nodedb_physical::physical_plan::StorageMode::Strict { ref schema } = - config.storage_mode - { - match strict_format::binary_tuple_to_json(current_bytes, schema) { - Some(v) => v, - None => { - return Err(ErrorCode::Internal { - detail: "failed to decode Binary Tuple for update".into(), - }); - } + // Strict, generated, expression RHS, or identity assignment: + // decode → mutate → re-encode. + let strict_schema = if is_strict { + match self.doc_configs.get(config_key).map(|c| &c.storage_mode) { + Some(nodedb_physical::physical_plan::StorageMode::Strict { schema }) => { + Some(schema) + } + _ => { + return Err(ErrorCode::Internal { + detail: "strict config missing during update".into(), + }); + } + } + } else { + None + }; + let mut doc = if let Some(schema) = strict_schema { + match strict_format::binary_tuple_to_json(current_bytes, schema) { + Some(v) => v, + None => { + return Err(ErrorCode::Internal { + detail: "failed to decode Binary Tuple for update".into(), + }); } - } else { - return Err(ErrorCode::Internal { - detail: "strict config missing during update".into(), - }); } } else { match doc_format::decode_document(current_bytes) { @@ -185,6 +199,9 @@ impl CoreLoop { } }; + let identity = + IdentitySnapshot::capture(strict_schema, declared_primary_key, updates, &doc); + // Expressions evaluate against the pre-update snapshot, so later // assignments don't observe earlier ones — matches PostgreSQL. let eval_doc: nodedb_types::Value = doc.clone().into(); @@ -223,6 +240,9 @@ impl CoreLoop { detail: format!("primary key '{pk}' cannot be NULL or omitted"), }); } + identity + .check_unchanged(&config_key.2, &doc) + .map_err(ErrorCode::from)?; // Recompute generated columns. if has_generated @@ -235,6 +255,11 @@ impl CoreLoop { return Err(e); } + // A declared numeric column holds the value its type stores, whether + // the assignment was a literal or computed. + strict_format::coerce_declared_doc(&mut doc, self.declared_columns(config_key)) + .map_err(ErrorCode::from)?; + Ok(PointUpdateBody::Document(doc)) } } diff --git a/nodedb/src/data/executor/handlers/transaction/stage_write/body.rs b/nodedb/src/data/executor/handlers/transaction/stage_write/body.rs index d8344b4b6..c54996a7b 100644 --- a/nodedb/src/data/executor/handlers/transaction/stage_write/body.rs +++ b/nodedb/src/data/executor/handlers/transaction/stage_write/body.rs @@ -16,6 +16,7 @@ use nodedb_types::{RowIdentity, StorageKey, Surrogate}; use crate::data::executor::core_loop::CoreLoop; use crate::data::executor::handlers::generated; +use crate::data::executor::handlers::identity_guard::IdentitySnapshot; use crate::data::executor::handlers::merge_helpers::check_declared_pk_not_null; use crate::data::executor::{doc_format, strict_format}; use crate::types::TenantId; @@ -88,6 +89,10 @@ impl CoreLoop { doc_format::canonicalize_document_for_storage(value) }; + // A declared numeric column holds the value its type stores, as on + // the durable path. + let value = strict_format::coerce_declared_body(value, self.declared_columns(&config_key))?; + let bitemporal = self.is_bitemporal(database_id, tid, collection); let sys_from_ms = self.bitemporal_now_ms(); @@ -181,6 +186,9 @@ impl CoreLoop { })?, }; + let identity = + IdentitySnapshot::capture(strict_schema.as_ref(), declared_primary_key, updates, &doc); + // Expressions evaluate against the pre-update snapshot (PostgreSQL // semantics): a later assignment observing a column updated earlier in // the same statement still sees the pre-statement value. @@ -213,6 +221,7 @@ impl CoreLoop { if strict_schema.is_none() { check_declared_pk_not_null(collection, &doc, declared_primary_key)?; } + identity.check_unchanged(collection, &doc)?; // Recompute generated columns after the patch. if let Some(config) = self.doc_configs.get(&config_key) @@ -222,6 +231,10 @@ impl CoreLoop { .map_err(crate::Error::DataPlane)?; } + // A declared numeric column holds the value its type stores, whether + // the assignment was a literal or computed. + strict_format::coerce_declared_doc(&mut doc, self.declared_columns(&config_key))?; + match strict_schema.as_ref() { Some(schema) => { let ndb_val: nodedb_types::Value = doc.into(); diff --git a/nodedb/src/data/executor/handlers/transaction/stage_write/stage_columnar_resolved_dml.rs b/nodedb/src/data/executor/handlers/transaction/stage_write/stage_columnar_resolved_dml.rs index 8988839b9..4b7c3d478 100644 --- a/nodedb/src/data/executor/handlers/transaction/stage_write/stage_columnar_resolved_dml.rs +++ b/nodedb/src/data/executor/handlers/transaction/stage_write/stage_columnar_resolved_dml.rs @@ -44,6 +44,7 @@ use nodedb_types::{RowIdentity, value_to_pk_string}; use crate::bridge::envelope::{ErrorCode, Response}; use crate::bridge::scan_filter::ScanFilter; use crate::data::executor::core_loop::CoreLoop; +use crate::data::executor::handlers::columnar_write::coerce_columnar_row; use crate::data::executor::task::ExecutionTask; use crate::types::{TenantId, TxnId}; @@ -144,18 +145,27 @@ impl CoreLoop { for (surrogate, row) in &matched { by_pk.insert(encode_pk(&row[pk_idx]), *surrogate); } - let mut resolved: Vec<(u32, RowIdentity, &Vec)> = Vec::with_capacity(rows.len()); + let mut resolved: Vec<(u32, RowIdentity, Vec)> = Vec::with_capacity(rows.len()); for (pk, new_row) in rows { let identity = match resolved_pk_identity(collection, pk) { Ok(identity) => identity, Err(e) => return self.response_error(task, e), }; match by_pk.get(&encode_pk(pk)) { - Some(surrogate) => resolved.push((*surrogate, identity, new_row)), + Some(surrogate) => resolved.push((*surrogate, identity, new_row.clone())), None => return self.response_error(task, ErrorCode::OllpRetryRequired), } } + // Every shipped post-image meets the declared column rule before the + // policy decides it and before the first put stages, as the durable + // apply does. + for (_, _, new_row) in &mut resolved { + if let Err(e) = coerce_columnar_row(&schema, new_row) { + return self.response_error(task, e); + } + } + // The gate stays on every write path even though `DecidedEarlierInRequest` // makes this a no-op — mirrors `execute_columnar_resolved_update`. if let Err(response) = self.stage_admit_columnar_rows( @@ -176,7 +186,7 @@ impl CoreLoop { } } for (surrogate, identity, new_row) in resolved { - let body = match nodedb_types::value_to_msgpack(&Value::Array(new_row.clone())) { + let body = match nodedb_types::value_to_msgpack(&Value::Array(new_row)) { Ok(b) => b, Err(e) => { return self.response_error( diff --git a/nodedb/src/data/executor/handlers/transaction/stage_write/stage_kv.rs b/nodedb/src/data/executor/handlers/transaction/stage_write/stage_kv.rs index cfd8f0fd1..0a880f5bc 100644 --- a/nodedb/src/data/executor/handlers/transaction/stage_write/stage_kv.rs +++ b/nodedb/src/data/executor/handlers/transaction/stage_write/stage_kv.rs @@ -51,6 +51,13 @@ impl CoreLoop { { return self.response_error(task, refusal); } + // A staged row body holds its declared numeric columns' values, as an + // autocommit write does. + let coerced = match self.coerce_kv_op_bodies(task.request.database_id.as_u64(), tid, op) { + Ok(coerced) => coerced, + Err(e) => return self.response_error(task, e), + }; + let op = coerced.as_ref().unwrap_or(op); match op { KvOp::Put { collection, diff --git a/nodedb/src/data/executor/handlers/transaction/stage_write/stage_kv_atomic.rs b/nodedb/src/data/executor/handlers/transaction/stage_write/stage_kv_atomic.rs index 309c104e7..7e4748344 100644 --- a/nodedb/src/data/executor/handlers/transaction/stage_write/stage_kv_atomic.rs +++ b/nodedb/src/data/executor/handlers/transaction/stage_write/stage_kv_atomic.rs @@ -45,9 +45,11 @@ use super::stage_kv::kv_row_identity; use crate::bridge::envelope::Response; use crate::data::executor::core_loop::CoreLoop; use crate::data::executor::handlers::kv::atomic::incr_float_reply; +use crate::data::executor::handlers::kv::declared_body::fit_kv_image; use crate::data::executor::handlers::transaction::overlay::StagedTtl; use crate::data::executor::response_codec; use crate::data::executor::task::ExecutionTask; +use crate::engine::kv::fitted_counter_f64; use crate::types::TxnId; /// FNV-1a 32-bit hash, used only to derive a stable, collection-local overlay @@ -196,7 +198,11 @@ impl CoreLoop { ) -> Response { let current = self.resolve_kv_current(ctx, key); match atomic_compute::incr(current.as_deref(), delta, shape) { - Ok((new_i64, new_bytes)) => { + Ok((new_i64, computed)) => { + let new_bytes = match self.stage_fit_kv_image(ctx, computed) { + Ok(bytes) => bytes, + Err(e) => return self.response_error(ctx.task, e), + }; if let Err(e) = self.stage_admit_kv_image(ctx, &new_bytes, rls_write_check) { return self.response_error(ctx.task, e); } @@ -219,7 +225,15 @@ impl CoreLoop { ) -> Response { let current = self.resolve_kv_current(ctx, key); match atomic_compute::incr_float(current.as_deref(), delta, shape) { - Ok((new_f64, new_bytes)) => { + Ok((computed_f64, computed)) => { + let declared = self.declared_columns(&ctx.coll_key); + let (new_f64, new_bytes) = match fit_kv_image(&computed, declared) { + Ok(None) => (computed_f64, computed), + Ok(Some(fitted)) => { + (fitted_counter_f64(computed_f64, &computed, &fitted), fitted) + } + Err(e) => return self.response_error(ctx.task, e), + }; if let Err(e) = self.stage_admit_kv_image(ctx, &new_bytes, rls_write_check) { return self.response_error(ctx.task, e); } @@ -233,6 +247,15 @@ impl CoreLoop { } } + /// The image a staged KV write stores for the image it computed: fitted + /// to the collection's declared numeric columns, the rule the durable + /// apply runs. A value past a declared column is refused here, at + /// statement time. + fn stage_fit_kv_image(&self, ctx: &StageCtx<'_>, image: Vec) -> crate::Result> { + let declared = self.declared_columns(&ctx.coll_key); + Ok(fit_kv_image(&image, declared)?.unwrap_or(image)) + } + /// Decide one staged KV image against the compiled write policy, naming /// the row by the overlay's own doc-id so a rejection reports the same /// identity the overlay filed it under. @@ -273,6 +296,10 @@ impl CoreLoop { }; if matches { + let write_bytes = match self.stage_fit_kv_image(ctx, write_bytes) { + Ok(bytes) => bytes, + Err(e) => return self.response_error(ctx.task, e), + }; if let Err(e) = self.stage_admit_kv_image(ctx, &write_bytes, rls_write_check) { return self.response_error(ctx.task, e); } @@ -308,6 +335,10 @@ impl CoreLoop { Ok(bytes) => bytes, Err(e) => return self.response_atomic_error(ctx.task, ctx.collection, e.into()), }; + let write_bytes = match self.stage_fit_kv_image(ctx, write_bytes) { + Ok(bytes) => bytes, + Err(e) => return self.response_error(ctx.task, e), + }; if let Err(e) = self.stage_admit_kv_image(ctx, &write_bytes, rls_write_check) { return self.response_error(ctx.task, e); } diff --git a/nodedb/src/data/executor/handlers/transaction/stage_write/stage_kv_conflict.rs b/nodedb/src/data/executor/handlers/transaction/stage_write/stage_kv_conflict.rs index c33f96ae6..0bf3dc9e1 100644 --- a/nodedb/src/data/executor/handlers/transaction/stage_write/stage_kv_conflict.rs +++ b/nodedb/src/data/executor/handlers/transaction/stage_write/stage_kv_conflict.rs @@ -34,12 +34,19 @@ impl CoreLoop { rls_write_check: &nodedb_types::RlsWriteCheck, ) -> Response { let existing = self.resolve_kv_current(ctx, key); - let (stored_bytes, op) = match &existing { - None => (value.to_vec(), "insert"), - Some(existing_raw) => match merge_kv_conflict_body(existing_raw, value, updates) { - Ok(b) => (b, "update"), - Err(e) => return self.response_error(ctx.task, e), - }, + let op = if existing.is_some() { + "update" + } else { + "insert" + }; + let stored_bytes = match merge_kv_conflict_body( + existing.as_deref(), + value, + updates, + self.declared_columns_of(ctx.database_id, ctx.tid, ctx.collection), + ) { + Ok(b) => b, + Err(e) => return self.response_error(ctx.task, e), }; // Staging is where an in-transaction statement's row image is produced, diff --git a/nodedb/src/data/executor/handlers/transaction/stage_write/stage_kv_predicate.rs b/nodedb/src/data/executor/handlers/transaction/stage_write/stage_kv_predicate.rs index 7bdb60ed6..c8eb39688 100644 --- a/nodedb/src/data/executor/handlers/transaction/stage_write/stage_kv_predicate.rs +++ b/nodedb/src/data/executor/handlers/transaction/stage_write/stage_kv_predicate.rs @@ -91,8 +91,12 @@ impl CoreLoop { let mut affected = 0usize; for row in matched { - let computed = match merge_field_updates(collection, Some(row.body.as_slice()), updates) - { + let computed = match merge_field_updates( + collection, + Some(row.body.as_slice()), + updates, + self.declared_columns(&coll_key), + ) { Ok(c) => c, Err(e) => return self.response_error(task, e), }; diff --git a/nodedb/src/data/executor/handlers/transaction/stage_write/stage_kv_transfer.rs b/nodedb/src/data/executor/handlers/transaction/stage_write/stage_kv_transfer.rs index 94734aeb4..9c3d0a8a7 100644 --- a/nodedb/src/data/executor/handlers/transaction/stage_write/stage_kv_transfer.rs +++ b/nodedb/src/data/executor/handlers/transaction/stage_write/stage_kv_transfer.rs @@ -27,6 +27,7 @@ use nodedb_physical::physical_plan::KvOp; use super::context::StageCtx; use crate::bridge::envelope::{ErrorCode, Response}; use crate::data::executor::core_loop::CoreLoop; +use crate::data::executor::handlers::kv::declared_body::fit_kv_image; use crate::data::executor::handlers::kv::field_compute::merge_field_updates; use crate::data::executor::handlers::kv::transfer_compute::{TransferError, compute_transfer}; use crate::data::executor::response_codec; @@ -50,7 +51,8 @@ struct StageTransfer<'a> { source_key: &'a [u8], dest_key: &'a [u8], field: &'a str, - amount: f64, + /// The amount, typed by the field it moves. + amount: nodedb_physical::physical_plan::TransferAmount, /// Compiled RLS write predicate for the collection both rows live in. rls_write_check: &'a nodedb_types::RlsWriteCheck, } @@ -168,7 +170,12 @@ impl CoreLoop { Err(e) => self.response_error(ctx.task, e), }; } - let computed = match merge_field_updates(ctx.collection, current.as_deref(), updates) { + let computed = match merge_field_updates( + ctx.collection, + current.as_deref(), + updates, + self.declared_columns_of(ctx.database_id, ctx.tid, ctx.collection), + ) { Ok(c) => c, Err(e) => return self.response_error(ctx.task, e), }; @@ -206,8 +213,16 @@ impl CoreLoop { let dest_ctx = self.kv_atomic_stage_ctx(task, cx.tid, cx.txn_id, collection, dest_key); let dest_bytes = self.resolve_kv_current(&dest_ctx, dest_key); - let computed = match compute_transfer(&source_bytes, dest_bytes.as_deref(), field, amount) { + let declared = self.declared_columns(&source_ctx.coll_key); + let computed = match compute_transfer( + &source_bytes, + dest_bytes.as_deref(), + field, + amount, + declared, + ) { Ok(c) => c, + Err(TransferError::Declared(e)) => return self.response_error(task, e), Err(TransferError::TypeMismatch(detail)) => { return self.response_error( task, @@ -253,9 +268,9 @@ impl CoreLoop { "source_key": src_str, "dest_key": dst_str, "field": field, - "amount": amount, - "source_balance": computed.source_balance_after, - "dest_balance": computed.dest_balance_after, + "amount": computed.amount.to_json(), + "source_balance": computed.source_balance_after.to_json(), + "dest_balance": computed.dest_balance_after.to_json(), })) { Ok(payload) => self.response_with_payload(task, payload), Err(e) => self.response_error(task, e), @@ -284,16 +299,23 @@ impl CoreLoop { return self.response_error(task, ErrorCode::NotFound); }; let dest_ctx = self.kv_atomic_stage_ctx(task, cx.tid, cx.txn_id, dest_collection, dest_key); + // The row arriving at the destination meets the destination's + // declared numeric columns. + let fitted = match fit_kv_image(&item_bytes, self.declared_columns(&dest_ctx.coll_key)) { + Ok(fitted) => fitted, + Err(e) => return self.response_error(task, e), + }; - // The same bytes are two different images to two independent policies: - // the row leaving the source and the row arriving at the destination. - // Both are decided before the source is tombstoned, so a rejected move - // never removes the row it could not deliver. + // The row leaving the source and the row arriving at the destination + // are two images to two independent policies. Both are decided before + // the source is tombstoned, so a rejected move never removes the row + // it could not deliver. if let Err(e) = self.stage_admit_kv_image(&source_ctx, &item_bytes, source_rls_write_check) { return self.response_error(task, e); } - if let Err(e) = self.stage_admit_kv_image(&dest_ctx, &item_bytes, dest_rls_write_check) { + let dest_bytes = fitted.unwrap_or(item_bytes); + if let Err(e) = self.stage_admit_kv_image(&dest_ctx, &dest_bytes, dest_rls_write_check) { return self.response_error(task, e); } @@ -303,7 +325,7 @@ impl CoreLoop { &source_ctx.document_id, ); - if let Err(e) = self.stage_put_capped(&dest_ctx, item_bytes) { + if let Err(e) = self.stage_put_capped(&dest_ctx, dest_bytes) { return self.response_error(task, e); } diff --git a/nodedb/src/data/executor/handlers/transaction/stage_write/stage_upsert.rs b/nodedb/src/data/executor/handlers/transaction/stage_write/stage_upsert.rs index 44aad66cf..d61074251 100644 --- a/nodedb/src/data/executor/handlers/transaction/stage_write/stage_upsert.rs +++ b/nodedb/src/data/executor/handlers/transaction/stage_write/stage_upsert.rs @@ -144,10 +144,15 @@ impl CoreLoop { strict_format::value_to_binary_tuple(&merged, schema, ctx.collection) } } else { - nodedb_types::value_to_msgpack(&merged).map_err(|e| crate::Error::Serialization { - format: "msgpack".into(), - detail: format!("staged upsert merge: {e}"), - }) + let body = nodedb_types::value_to_msgpack(&merged).map_err(|e| { + crate::Error::Serialization { + format: "msgpack".into(), + detail: format!("staged upsert merge: {e}"), + } + })?; + // A declared numeric column holds the value its type stores, as + // on the durable path. + strict_format::coerce_declared_body(body, self.declared_columns(&config_key)) } } diff --git a/nodedb/src/data/executor/strict_format/coerce.rs b/nodedb/src/data/executor/strict_format/coerce.rs index 4ebba44fd..dfb7f50d2 100644 --- a/nodedb/src/data/executor/strict_format/coerce.rs +++ b/nodedb/src/data/executor/strict_format/coerce.rs @@ -2,11 +2,125 @@ //! Value coercion and JSON conversion for strict document columns. -use nodedb_types::columnar::ColumnType; +use nodedb_types::columnar::{ColumnDef, ColumnType, FloatWidth, IntWidth}; use nodedb_types::value::Value; -/// Coerce a `nodedb_types::Value` to match a column's declared type. -pub fn coerce_value(val: &Value, col_type: &ColumnType, col_name: &str) -> crate::Result { +/// Coerce a `nodedb_types::Value` to the value `column` stores. +/// +/// The value is re-typed to the column type, then checked against the +/// declared integer or float width. A value out of the declared range is +/// refused with [`crate::Error::NumericValueOutOfRange`] (SQLSTATE 22003). +pub fn coerce_value(val: &Value, column: &ColumnDef) -> crate::Result { + coerce_declared_value( + val, + &column.column_type, + column.int_width, + column.float_width, + &column.name, + ) +} + +/// The declared numeric rule every write path applies: re-type `val` to +/// `col_type`, then check the declared width. Strict, columnar, schemaless +/// and KV columns all run this one rule. +pub(crate) fn coerce_declared_value( + val: &Value, + col_type: &ColumnType, + int_width: Option, + float_width: Option, + col_name: &str, +) -> crate::Result { + let typed = coerce_to_type(val, col_type, col_name)?; + check_declared_width(&typed, int_width, float_width, col_name)?; + Ok(typed) +} + +/// Refuse an integer outside the declared `SMALLINT` / `INTEGER` range, and +/// a finite float that overflows `REAL`. Rounding into `REAL` is accepted. +fn check_declared_width( + value: &Value, + int_width: Option, + float_width: Option, + col_name: &str, +) -> crate::Result<()> { + match value { + Value::Integer(n) => match int_width { + Some(width) if !width.contains(*n) => { + Err(out_of_range(col_name, &n.to_string(), width.pg_type_name())) + } + _ => Ok(()), + }, + Value::Float(f) => match float_width { + Some(width @ FloatWidth::F32) if f.is_finite() && !(*f as f32).is_finite() => { + Err(out_of_range(col_name, &f.to_string(), width.pg_type_name())) + } + _ => Ok(()), + }, + _ => Ok(()), + } +} + +/// A value outside the range of `declared_type`: SQLSTATE `22003`. +fn out_of_range(column: &str, value: &str, declared_type: &str) -> crate::Error { + crate::Error::NumericValueOutOfRange { + detail: format!( + "value {value} is out of range for column '{column}' of type {declared_type}" + ), + } +} + +/// Text that does not parse as `declared_type`: SQLSTATE `22P02`. +fn invalid_text(column: &str, text: &str, declared_type: &str) -> crate::Error { + crate::Error::InvalidTextRepresentation { + detail: format!("column '{column}': cannot parse '{text}' as {declared_type}"), + } +} + +/// A value of a kind `declared_type` does not hold: SQLSTATE `42804`. +fn wrong_kind(column: &str, value: &Value, declared_type: &str) -> crate::Error { + crate::Error::DatatypeMismatch { + detail: format!("column '{column}': expected {declared_type}, got {value:?}"), + } +} + +/// The error for a float with no `i64` image. A finite float with a +/// fraction, such as `2.5`, is the wrong kind for an integer column. NaN, an +/// infinity, and a magnitude past `i64` are out of range. +fn float_to_int_error(column: &str, f: f64) -> crate::Error { + if f.is_finite() && f.fract() != 0.0 { + crate::Error::DatatypeMismatch { + detail: format!("column '{column}': {f} is not a whole number, expected INT"), + } + } else { + out_of_range(column, &f.to_string(), "INT") + } +} + +/// The error for text that does not parse as an integer of `declared_type`. +/// Digits past the `i64` range are out of range. Anything else is text the +/// type cannot read. +fn int_text_error( + column: &str, + text: &str, + declared_type: &str, + error: &std::num::ParseIntError, +) -> crate::Error { + match error.kind() { + std::num::IntErrorKind::PosOverflow | std::num::IntErrorKind::NegOverflow => { + out_of_range(column, text, declared_type) + } + _ => invalid_text(column, text, declared_type), + } +} + +/// Re-type `val` to `col_type`. No declared width is checked. +/// +/// A refusal carries the SQLSTATE PostgreSQL gives the same assignment: +/// text that does not parse is `22P02`, a value past the type's range is +/// `22003`, and a value of the wrong kind is `42804`. A timestamp column +/// refuses text that is not a date-time with `22007`, and an instant past +/// its range with `22008`. +fn coerce_to_type(val: &Value, col_type: &ColumnType, col_name: &str) -> crate::Result { match col_type { ColumnType::Bool => match val { Value::Bool(_) => Ok(val.clone()), @@ -14,46 +128,48 @@ pub fn coerce_value(val: &Value, col_type: &ColumnType, col_name: &str) -> crate Value::String(s) => match s.to_lowercase().as_str() { "true" | "1" | "yes" => Ok(Value::Bool(true)), "false" | "0" | "no" => Ok(Value::Bool(false)), - _ => Err(crate::Error::BadRequest { - detail: format!("column '{col_name}': cannot coerce '{s}' to BOOL"), - }), + _ => Err(invalid_text(col_name, s, "BOOL")), }, - _ => Err(crate::Error::BadRequest { - detail: format!("column '{col_name}': expected BOOL, got {val:?}"), - }), + _ => Err(wrong_kind(col_name, val, "BOOL")), }, ColumnType::Int64 => match val { Value::Integer(_) => Ok(val.clone()), - Value::Float(f) => Ok(Value::Integer(*f as i64)), - Value::String(s) => { - s.parse::() - .map(Value::Integer) - .map_err(|_| crate::Error::BadRequest { - detail: format!("column '{col_name}': cannot parse '{s}' as INT"), - }) - } - _ => Err(crate::Error::BadRequest { - detail: format!("column '{col_name}': expected INT, got {val:?}"), + // A whole float in `i64` range; any other float has no INT image. + Value::Float(f) => float_to_i64(*f) + .map(Value::Integer) + .ok_or_else(|| float_to_int_error(col_name, *f)), + // A whole decimal in `i64` range. A fractional decimal is the + // wrong kind. A whole one past `i64`, such as a `u64` above + // `i64::MAX`, is out of range. + Value::Decimal(d) if !d.is_integer() => Err(crate::Error::DatatypeMismatch { + detail: format!("column '{col_name}': {d} is not a whole number, expected INT"), }), + Value::Decimal(d) => rust_decimal::prelude::ToPrimitive::to_i64(d) + .map(Value::Integer) + .ok_or_else(|| out_of_range(col_name, &d.to_string(), "INT")), + Value::String(s) => s + .parse::() + .map(Value::Integer) + .map_err(|e| int_text_error(col_name, s, "INT", &e)), + _ => Err(wrong_kind(col_name, val, "INT")), }, ColumnType::Float64 => match val { Value::Float(_) => Ok(val.clone()), Value::Integer(n) => Ok(Value::Float(*n as f64)), - Value::String(s) => { - s.parse::() - .map(Value::Float) - .map_err(|_| crate::Error::BadRequest { - detail: format!("column '{col_name}': cannot parse '{s}' as FLOAT"), - }) - } - _ => Err(crate::Error::BadRequest { - detail: format!("column '{col_name}': expected FLOAT, got {val:?}"), - }), + Value::Decimal(d) => rust_decimal::prelude::ToPrimitive::to_f64(d) + .map(Value::Float) + .ok_or_else(|| out_of_range(col_name, &d.to_string(), "FLOAT")), + Value::String(s) => s + .parse::() + .map(Value::Float) + .map_err(|_| invalid_text(col_name, s, "FLOAT")), + _ => Err(wrong_kind(col_name, val, "FLOAT")), }, ColumnType::String | ColumnType::Uuid | ColumnType::Ulid | ColumnType::Regex => match val { Value::String(_) | Value::Uuid(_) | Value::Ulid(_) | Value::Regex(_) => Ok(val.clone()), Value::Integer(n) => Ok(Value::String(n.to_string())), Value::Float(f) => Ok(Value::String(f.to_string())), + Value::Decimal(d) => Ok(Value::String(d.to_string())), Value::Bool(b) => Ok(Value::String(b.to_string())), other => Ok(Value::String(format!("{other:?}"))), }, @@ -64,78 +180,14 @@ pub fn coerce_value(val: &Value, col_type: &ColumnType, col_name: &str) -> crate .unwrap_or_else(|_| s.as_bytes().to_vec()); Ok(Value::Bytes(bytes)) } - _ => Err(crate::Error::BadRequest { - detail: format!("column '{col_name}': expected BYTES, got {val:?}"), - }), - }, - ColumnType::Timestamp => match val { - Value::NaiveDateTime(_) => Ok(val.clone()), - Value::DateTime(dt) => Ok(Value::NaiveDateTime(*dt)), - Value::Integer(ms) => Ok(Value::NaiveDateTime( - nodedb_types::NdbDateTime::from_millis(*ms).map_err(|e| { - crate::Error::BadRequest { - detail: format!("column '{col_name}': {e}"), - } - })?, - )), - Value::Float(f) => Ok(Value::NaiveDateTime( - nodedb_types::NdbDateTime::from_millis(*f as i64).map_err(|e| { - crate::Error::BadRequest { - detail: format!("column '{col_name}': {e}"), - } - })?, - )), - Value::String(s) => { - if let Ok(ms) = s.parse::() { - Ok(Value::NaiveDateTime( - nodedb_types::NdbDateTime::from_millis(ms).map_err(|e| { - crate::Error::BadRequest { - detail: format!("column '{col_name}': {e}"), - } - })?, - )) - } else { - Ok(Value::String(s.clone())) - } - } - _ => Err(crate::Error::BadRequest { - detail: format!("column '{col_name}': expected TIMESTAMP, got {val:?}"), - }), - }, - ColumnType::Timestamptz => match val { - Value::DateTime(_) => Ok(val.clone()), - Value::NaiveDateTime(dt) => Ok(Value::DateTime(*dt)), - Value::Integer(ms) => Ok(Value::DateTime( - nodedb_types::NdbDateTime::from_millis(*ms).map_err(|e| { - crate::Error::BadRequest { - detail: format!("column '{col_name}': {e}"), - } - })?, - )), - Value::Float(f) => Ok(Value::DateTime( - nodedb_types::NdbDateTime::from_millis(*f as i64).map_err(|e| { - crate::Error::BadRequest { - detail: format!("column '{col_name}': {e}"), - } - })?, - )), - Value::String(s) => { - if let Ok(ms) = s.parse::() { - Ok(Value::DateTime( - nodedb_types::NdbDateTime::from_millis(ms).map_err(|e| { - crate::Error::BadRequest { - detail: format!("column '{col_name}': {e}"), - } - })?, - )) - } else { - Ok(Value::String(s.clone())) - } - } - _ => Err(crate::Error::BadRequest { - detail: format!("column '{col_name}': expected TIMESTAMPTZ, got {val:?}"), - }), + _ => Err(wrong_kind(col_name, val, "BYTES")), }, + ColumnType::Timestamp => { + coerce_instant(val, col_name, "TIMESTAMP").map(Value::NaiveDateTime) + } + ColumnType::Timestamptz => { + coerce_instant(val, col_name, "TIMESTAMPTZ").map(Value::DateTime) + } ColumnType::SystemTimestamp => { // Engine-assigned; user-supplied values must not reach coercion. let _ = val; @@ -145,44 +197,47 @@ pub fn coerce_value(val: &Value, col_type: &ColumnType, col_name: &str) -> crate ), }) } - ColumnType::Decimal { .. } => match val { - Value::Decimal(_) => Ok(val.clone()), - Value::Float(f) => rust_decimal::Decimal::try_from(*f) - .map(Value::Decimal) - .map_err(|_| crate::Error::BadRequest { - detail: format!("column '{col_name}': cannot convert {f} to DECIMAL"), - }), - Value::Integer(n) => Ok(Value::Decimal(rust_decimal::Decimal::from(*n))), - Value::String(s) => s - .parse::() - .map(Value::Decimal) - .map_err(|_| crate::Error::BadRequest { - detail: format!("column '{col_name}': cannot parse '{s}' as DECIMAL"), + ColumnType::Decimal(typmod) => { + let d = match val { + Value::Decimal(d) => *d, + Value::Float(f) => rust_decimal::Decimal::try_from(*f) + .map_err(|_| out_of_range(col_name, &f.to_string(), "DECIMAL"))?, + Value::Integer(n) => rust_decimal::Decimal::from(*n), + Value::String(s) => { + s.parse::() + .map_err(|error| match error { + rust_decimal::Error::ExceedsMaximumPossibleValue + | rust_decimal::Error::LessThanMinimumPossibleValue => { + out_of_range(col_name, s, "DECIMAL") + } + _ => invalid_text(col_name, s, "DECIMAL"), + })? + } + _ => return Err(wrong_kind(col_name, val, "DECIMAL")), + }; + match typmod { + Some(typmod) => typmod.fit(d).map(Value::Decimal).map_err(|error| { + crate::Error::NumericValueOutOfRange { + detail: format!("column '{col_name}': {error}"), + } }), - _ => Err(crate::Error::BadRequest { - detail: format!("column '{col_name}': expected DECIMAL, got {val:?}"), - }), - }, + None => Ok(Value::Decimal(d)), + } + } ColumnType::Vector(dim) => match val { Value::Bytes(b) if b.len() == *dim as usize * 4 => Ok(val.clone()), Value::Array(arr) => { - let floats = extract_vector_floats(arr); + let floats = extract_vector_floats(arr, col_name, *dim)?; validate_and_encode_vector(col_name, *dim, &floats) } Value::String(s) => { // UPDATE path may serialize ARRAY literal as string — parse it. match crate::data::executor::vector_string::parse_vector_string(s) { Some(floats) => validate_and_encode_vector(col_name, *dim, &floats), - None => Err(crate::Error::BadRequest { - detail: format!( - "column '{col_name}': expected VECTOR array, got String({s:?})" - ), - }), + None => Err(invalid_text(col_name, s, &format!("VECTOR({dim})"))), } } - _ => Err(crate::Error::BadRequest { - detail: format!("column '{col_name}': expected VECTOR array, got {val:?}"), - }), + _ => Err(wrong_kind(col_name, val, &format!("VECTOR({dim})"))), }, ColumnType::SparseVector => { // Variable-length string-backed: the `'{id: weight}'` literal (or a @@ -190,20 +245,19 @@ pub fn coerce_value(val: &Value, col_type: &ColumnType, col_name: &str) -> crate // index-build time. Schema validation catches genuine mismatches. Ok(val.clone()) } - ColumnType::Geometry => Ok(Value::String(format!("{val:?}"))), + // The tuple encoder stores a native geometry, or WKT / GeoJSON text. + ColumnType::Geometry => match val { + Value::Geometry(_) | Value::String(_) => Ok(val.clone()), + _ => Err(wrong_kind(col_name, val, "GEOMETRY")), + }, ColumnType::Duration => match val { Value::Duration(_) => Ok(val.clone()), Value::Integer(n) => Ok(Value::Integer(*n)), - Value::String(s) => { - s.parse::() - .map(Value::Integer) - .map_err(|_| crate::Error::BadRequest { - detail: format!("column '{col_name}': cannot parse '{s}' as DURATION"), - }) - } - _ => Err(crate::Error::BadRequest { - detail: format!("column '{col_name}': expected DURATION, got {val:?}"), - }), + Value::String(s) => s + .parse::() + .map(Value::Integer) + .map_err(|e| int_text_error(col_name, s, "DURATION", &e)), + _ => Err(wrong_kind(col_name, val, "DURATION")), }, ColumnType::Json | ColumnType::Array @@ -212,11 +266,13 @@ pub fn coerce_value(val: &Value, col_type: &ColumnType, col_name: &str) -> crate | ColumnType::Record => { // Variable-length inline MessagePack column: val is raw bytes — deserialize to Value. if let Value::Bytes(b) = val { - Ok(match nodedb_types::value_from_msgpack(b) { - Ok(v) => v, - Err(e) => { - tracing::warn!(len = b.len(), error = %e, "corrupted msgpack in strict column"); - Value::Null + nodedb_types::value_from_msgpack(b).map_err(|e| { + crate::Error::InvalidTextRepresentation { + detail: format!( + "column '{col_name}': {} bytes are not a MessagePack value of \ + type {col_type}: {e}", + b.len() + ), } }) } else { @@ -229,21 +285,92 @@ pub fn coerce_value(val: &Value, col_type: &ColumnType, col_name: &str) -> crate } } -/// Extract f32 floats from a `Value::Array`. -fn extract_vector_floats(arr: &[Value]) -> Vec { +/// `f` as an `i64` when it is a whole number in `i64` range. NaN and the +/// infinities have no `i64` image. +pub(crate) fn float_to_i64(f: f64) -> Option { + // `i64::MAX as f64` rounds up to 2^63, which is out of range. + const TWO_POW_63: f64 = 9_223_372_036_854_775_808.0; + (f.fract() == 0.0 && (-TWO_POW_63..TWO_POW_63).contains(&f)).then_some(f as i64) +} + +/// The instant a `TIMESTAMP` or `TIMESTAMPTZ` column of type name `kind` +/// stores for `val`. An integer or a float is epoch milliseconds. +fn coerce_instant( + val: &Value, + col_name: &str, + kind: &str, +) -> crate::Result { + match val { + Value::NaiveDateTime(dt) | Value::DateTime(dt) => Ok(*dt), + Value::Integer(ms) => nodedb_types::NdbDateTime::from_millis(*ms) + .map_err(|_| instant_overflow(col_name, &ms.to_string(), kind)), + Value::Float(f) => float_millis_instant(*f, col_name, kind), + Value::String(s) => parse_instant(s, col_name, kind), + _ => Err(wrong_kind(col_name, val, kind)), + } +} + +/// Microseconds per millisecond. +const MICROS_PER_MILLI: f64 = 1_000.0; + +/// A float count of epoch milliseconds as an instant. An instant holds epoch +/// microseconds, so a fractional millisecond is kept to the microsecond. A +/// part below one microsecond rounds half to even, as PostgreSQL's +/// `to_timestamp(double precision)` does. NaN, an infinity, and a value past +/// the microsecond range overflow the instant. +fn float_millis_instant( + ms: f64, + col_name: &str, + kind: &str, +) -> crate::Result { + float_to_i64((ms * MICROS_PER_MILLI).round_ties_even()) + .map(nodedb_types::NdbDateTime::from_micros) + .ok_or_else(|| instant_overflow(col_name, &ms.to_string(), kind)) +} + +/// A time text as an instant: an integer count of epoch milliseconds, or a +/// date-time spelling. A count past the instant range overflows the +/// instant. Any other text is not a date-time. +fn parse_instant(s: &str, col_name: &str, kind: &str) -> crate::Result { + if let Ok(ms) = s.parse::() { + return nodedb_types::NdbDateTime::from_millis(ms) + .map_err(|_| instant_overflow(col_name, s, kind)); + } + nodedb_types::datetime::NdbDateTime::parse(s).ok_or_else(|| { + crate::Error::InvalidDatetimeFormat { + detail: format!("column '{col_name}': cannot parse '{s}' as {kind}"), + } + }) +} + +/// An instant outside the range `kind` holds: SQLSTATE `22008`. +fn instant_overflow(column: &str, value: &str, kind: &str) -> crate::Error { + crate::Error::DatetimeFieldOverflow { + detail: format!("value {value} is out of range for column '{column}' of type {kind}"), + } +} + +/// The `f32` elements of a `Value::Array` bound for a `VECTOR(dim)` column. +/// An element that is not a number is the wrong kind. +fn extract_vector_floats(arr: &[Value], col_name: &str, dim: u32) -> crate::Result> { arr.iter() - .filter_map(|v| match v { - Value::Float(f) => Some(*f as f32), - Value::Integer(n) => Some(*n as f32), - _ => None, + .map(|v| match v { + Value::Float(f) => Ok(*f as f32), + Value::Integer(n) => Ok(*n as f32), + other => Err(wrong_kind( + col_name, + other, + &format!("a number for VECTOR({dim})"), + )), }) .collect() } -/// Validate dimension count and encode as little-endian bytes. +/// Check the dimension count and encode as little-endian bytes. A wrong +/// count is SQLSTATE `22000` (data_exception), as pgvector gives it. fn validate_and_encode_vector(col_name: &str, dim: u32, floats: &[f32]) -> crate::Result { if floats.len() != dim as usize { - return Err(crate::Error::BadRequest { + return Err(crate::Error::DataException { detail: format!( "column '{col_name}': expected VECTOR({dim}), got {} elements", floats.len() @@ -254,7 +381,8 @@ fn validate_and_encode_vector(col_name: &str, dim: u32, floats: &[f32]) -> crate Ok(Value::Bytes(bytes)) } -/// Convert a typed `Value` to JSON (for pgwire output only). +/// Convert a typed `Value` to JSON (for pgwire output only). A set renders +/// as the JSON array of its members, as the wire rendering gives it. pub fn value_to_json(val: &Value) -> serde_json::Value { match val { Value::Null => serde_json::Value::Null, @@ -273,7 +401,9 @@ pub fn value_to_json(val: &Value) -> serde_json::Value { } Value::Duration(d) => serde_json::Value::String(d.to_string()), Value::Decimal(d) => serde_json::Value::String(d.to_string()), - Value::Array(arr) => serde_json::Value::Array(arr.iter().map(value_to_json).collect()), + Value::Array(arr) | Value::Set(arr) => { + serde_json::Value::Array(arr.iter().map(value_to_json).collect()) + } Value::Object(map) => { let mut obj = serde_json::Map::new(); for (k, v) in map { @@ -282,7 +412,6 @@ pub fn value_to_json(val: &Value) -> serde_json::Value { serde_json::Value::Object(obj) } Value::Geometry(_) - | Value::Set(_) | Value::Regex(_) | Value::Range { .. } | Value::Record { .. } @@ -403,6 +532,315 @@ mod tests { assert!(result.unwrap_err().to_string().contains("not bitemporal")); } + fn dec(text: &str) -> rust_decimal::Decimal { + text.parse().expect("test decimal parses") + } + + fn decimal(precision: i64, scale: i64) -> ColumnType { + ColumnType::Decimal(Some( + nodedb_types::columnar::DecimalTypmod::new(precision, scale) + .expect("test typmod is valid"), + )) + } + + /// Coerce `val` into a nullable column of `col_type` with no width. + fn coerce_typed(val: &Value, col_type: &ColumnType, name: &str) -> crate::Result { + coerce_value(val, &ColumnDef::nullable(name, *col_type)) + } + + /// Coerce `val` into a nullable column declared as `declared`. + fn coerce_declared_as(val: &Value, declared: &str) -> crate::Result { + let col_type: ColumnType = declared.parse().expect("declared type parses"); + coerce_value( + val, + &ColumnDef::nullable("v", col_type).with_declared_width(declared), + ) + } + + /// A value rounds to the declared scale, half away from zero, and carries + /// exactly that many fractional digits. + #[test] + fn decimal_rounds_to_scale_half_away_from_zero() { + for (input, expected) in [ + ("1.005", "1.01"), + ("-1.005", "-1.01"), + ("1.004", "1.00"), + ("12.5", "12.50"), + ("999.994", "999.99"), + ] { + let got = coerce_typed(&Value::Decimal(dec(input)), &decimal(5, 2), "d").unwrap(); + let Value::Decimal(d) = got else { + panic!("{input}: expected a decimal, got {got:?}"); + }; + assert_eq!(d.to_string(), expected, "{input}"); + } + let from_text = coerce_typed(&Value::String("1.005".into()), &decimal(5, 2), "d").unwrap(); + assert_eq!(from_text, Value::Decimal(dec("1.01"))); + let from_int = coerce_typed(&Value::Integer(7), &decimal(5, 2), "d").unwrap(); + assert_eq!(from_int, Value::Decimal(dec("7.00"))); + } + + /// A value whose rounded integer part has more than `precision - scale` + /// digits is refused as numeric value out of range. + #[test] + fn decimal_past_precision_is_numeric_out_of_range() { + for input in ["123456.789", "1000", "999.995", "-1000.00"] { + let err = + coerce_typed(&Value::Decimal(dec(input)), &decimal(5, 2), "d").expect_err(input); + assert!( + matches!(err, crate::Error::NumericValueOutOfRange { .. }), + "{input}: {err:?}" + ); + } + let err = coerce_typed(&Value::Decimal(dec("0.995")), &decimal(2, 2), "d").unwrap_err(); + assert!(matches!(err, crate::Error::NumericValueOutOfRange { .. })); + assert_eq!( + coerce_typed(&Value::Decimal(dec("0.994")), &decimal(2, 2), "d").unwrap(), + Value::Decimal(dec("0.99")) + ); + } + + /// A plain `DECIMAL` keeps every digit it is given. + #[test] + fn unconstrained_decimal_is_not_limited() { + let wide = dec("12345678901234567890.123456789"); + assert_eq!( + coerce_typed(&Value::Decimal(wide), &ColumnType::Decimal(None), "d").unwrap(), + Value::Decimal(wide) + ); + } + + /// A strict or columnar `SMALLINT` column refuses a value past `i16`, in + /// every form the value arrives in, as numeric value out of range. + #[test] + fn smallint_column_refuses_a_value_past_its_width() { + assert_eq!( + coerce_declared_as(&Value::Integer(32767), "SMALLINT").expect("fits"), + Value::Integer(32767) + ); + for value in [ + Value::Integer(40000), + Value::Integer(-40000), + Value::Float(40000.0), + Value::String("40000".into()), + Value::Decimal(dec("40000")), + ] { + let err = coerce_declared_as(&value, "SMALLINT").expect_err("past smallint"); + assert!( + matches!(err, crate::Error::NumericValueOutOfRange { .. }), + "{value:?}: {err:?}" + ); + } + let err = coerce_declared_as(&Value::Integer(i64::from(i32::MAX) + 1), "INTEGER") + .expect_err("past integer"); + assert!( + matches!(err, crate::Error::NumericValueOutOfRange { .. }), + "{err:?}" + ); + assert_eq!( + coerce_declared_as(&Value::Integer(40000), "BIGINT").expect("fits bigint"), + Value::Integer(40000) + ); + } + + /// A `REAL` column refuses a finite value past `f32` and accepts one it + /// rounds. A `DOUBLE PRECISION` column holds the same value. + #[test] + fn real_column_refuses_only_an_overflow() { + let err = coerce_declared_as(&Value::Float(1e39), "REAL").expect_err("past f32"); + assert!( + matches!(err, crate::Error::NumericValueOutOfRange { .. }), + "{err:?}" + ); + assert_eq!( + coerce_declared_as(&Value::Float(1.1), "REAL").expect("rounds"), + Value::Float(1.1) + ); + assert_eq!( + coerce_declared_as(&Value::Float(1e39), "DOUBLE PRECISION").expect("fits f64"), + Value::Float(1e39) + ); + } + + /// The strict encoder applies the width: a row with an out-of-range + /// `SMALLINT` cell does not encode. + #[test] + fn strict_encoder_refuses_a_row_past_a_declared_width() { + let schema = StrictSchema::new(vec![ + ColumnDef::required("id", ColumnType::String).with_primary_key(), + ColumnDef::nullable("v", ColumnType::Int64).with_declared_width("SMALLINT"), + ]) + .expect("valid schema"); + let mut map = std::collections::HashMap::new(); + map.insert("id".into(), Value::String("a".into())); + map.insert("v".into(), Value::Integer(40000)); + let err = super::super::encode::value_to_binary_tuple(&Value::Object(map), &schema, "c") + .expect_err("past smallint"); + assert!( + matches!(err, crate::Error::NumericValueOutOfRange { .. }), + "{err:?}" + ); + } + + /// A float epoch-millisecond count keeps its fraction to the microsecond. + /// NaN, an infinity, and a value past the range overflow the instant. + #[test] + fn float_timestamp_is_exact_or_refused() { + for col_type in [ColumnType::Timestamp, ColumnType::Timestamptz] { + let got = coerce_typed(&Value::Float(1.5), &col_type, "t").unwrap(); + let (Value::NaiveDateTime(dt) | Value::DateTime(dt)) = got else { + panic!("expected an instant, got {got:?}"); + }; + assert_eq!(dt.micros, 1_500); + for bad in [f64::NAN, f64::INFINITY, f64::NEG_INFINITY, 1e300] { + let err = coerce_typed(&Value::Float(bad), &col_type, "t").expect_err("refused"); + assert!( + matches!(err, crate::Error::DatetimeFieldOverflow { .. }), + "{err:?}" + ); + } + } + } + + /// The SQLSTATE `crate::Error` variant a refusal carries. + fn refusal_kind(err: &crate::Error) -> &'static str { + match err { + crate::Error::InvalidTextRepresentation { .. } => "22P02", + crate::Error::NumericValueOutOfRange { .. } => "22003", + crate::Error::DatatypeMismatch { .. } => "42804", + crate::Error::DataException { .. } => "22000", + crate::Error::InvalidDatetimeFormat { .. } => "22007", + crate::Error::DatetimeFieldOverflow { .. } => "22008", + other => panic!("not a value refusal: {other:?}"), + } + } + + /// Each refusal carries the SQLSTATE PostgreSQL gives the same + /// assignment, and its message names the column, the value and the type. + #[test] + fn value_refusals_carry_their_postgres_sqlstate() { + let text = |s: &str| Value::String(s.into()); + let cases: Vec<(Value, ColumnType, &str, &str)> = vec![ + (text("x"), ColumnType::Int64, "22P02", "'x'"), + (text("2.5"), ColumnType::Int64, "22P02", "'2.5'"), + ( + text("99999999999999999999"), + ColumnType::Int64, + "22003", + "99999999999999999999", + ), + (Value::Float(2.5), ColumnType::Int64, "42804", "2.5"), + (Value::Float(f64::NAN), ColumnType::Int64, "22003", "NaN"), + ( + Value::Float(1e19), + ColumnType::Int64, + "22003", + "10000000000000000000", + ), + ( + Value::Decimal(dec("2.5")), + ColumnType::Int64, + "42804", + "2.5", + ), + ( + Value::Decimal(dec("18446744073709551615")), + ColumnType::Int64, + "22003", + "18446744073709551615", + ), + (Value::Bool(true), ColumnType::Int64, "42804", "Bool(true)"), + (text("abc"), ColumnType::Float64, "22P02", "'abc'"), + ( + Value::Bool(true), + ColumnType::Float64, + "42804", + "Bool(true)", + ), + (text("maybe"), ColumnType::Bool, "22P02", "'maybe'"), + ( + Value::Array(vec![Value::Integer(1)]), + ColumnType::Bool, + "42804", + "Array", + ), + (Value::Integer(1), ColumnType::Bytes, "42804", "Integer(1)"), + (text("1.2.3"), ColumnType::Decimal(None), "22P02", "'1.2.3'"), + ( + Value::Float(f64::INFINITY), + ColumnType::Decimal(None), + "22003", + "inf", + ), + ( + text("not a date"), + ColumnType::Timestamp, + "22007", + "'not a date'", + ), + ( + text("not a date"), + ColumnType::Timestamptz, + "22007", + "'not a date'", + ), + ( + Value::Bool(true), + ColumnType::Timestamptz, + "42804", + "Bool(true)", + ), + ( + Value::Integer(i64::MAX), + ColumnType::Timestamp, + "22008", + "9223372036854775807", + ), + ( + text("9223372036854775807"), + ColumnType::Timestamptz, + "22008", + "9223372036854775807", + ), + ( + Value::Float(f64::NAN), + ColumnType::Timestamp, + "22008", + "NaN", + ), + (text("soon"), ColumnType::Duration, "22P02", "'soon'"), + ( + Value::Integer(1), + ColumnType::Geometry, + "42804", + "Integer(1)", + ), + ( + Value::Array(vec![Value::Float(0.5)]), + ColumnType::Vector(2), + "22000", + "1 elements", + ), + ( + Value::Array(vec![Value::Float(0.5), text("x")]), + ColumnType::Vector(2), + "42804", + "String(\"x\")", + ), + ]; + for (value, col_type, sqlstate, shown) in cases { + let err = coerce_typed(&value, &col_type, "c").expect_err("refused"); + assert_eq!( + refusal_kind(&err), + sqlstate, + "{value:?} into {col_type}: {err:?}" + ); + let message = err.to_string(); + assert!(message.contains("'c'"), "names the column: {message}"); + assert!(message.contains(shown), "names the value: {message}"); + } + } + #[test] fn unknown_field_errors() { let schema = test_schema(); diff --git a/nodedb/src/data/executor/strict_format/declared.rs b/nodedb/src/data/executor/strict_format/declared.rs new file mode 100644 index 000000000..8911c1310 --- /dev/null +++ b/nodedb/src/data/executor/strict_format/declared.rs @@ -0,0 +1,259 @@ +// SPDX-License-Identifier: BUSL-1.1 + +//! The write rule for a declared numeric column of a schemaless document or +//! KV collection. +//! +//! A value meets [`coerce_declared_value`], the rule the strict and columnar +//! encoders apply: it is re-typed, then checked against the declared integer +//! or float width. Every write path of the two engines runs it on the row it +//! is about to store, so a literal and a computed value meet the same rule. +//! Fitting is idempotent: a value the planner already fitted fits to itself. +//! +//! A decimal is stored as its text. That is the form the planner writes for +//! a decimal literal, and the read path renders it at the declared scale. +//! `NULL` passes through: nullability is checked elsewhere. + +use std::collections::HashMap; + +use nodedb_physical::physical_plan::DeclaredColumn; +use nodedb_types::value::Value; + +use super::coerce::coerce_declared_value; + +/// `value` as `column` stores it. +/// +/// A value out of the declared range is refused with +/// [`crate::Error::NumericValueOutOfRange`] (SQLSTATE 22003). +pub(crate) fn coerce_declared(value: &Value, column: &DeclaredColumn) -> crate::Result { + if matches!(value, Value::Null) { + return Ok(Value::Null); + } + let typed = coerce_declared_value( + value, + &column.column_type, + column.int_width, + column.float_width, + &column.name, + )?; + Ok(match typed { + Value::Decimal(d) => Value::String(d.to_string()), + other => other, + }) +} + +/// Re-type every declared field of a MessagePack map body. +/// +/// Only the declared fields present in the body are rewritten. Every other +/// field keeps its bytes. A body that is not a map, or holds no declared +/// field, is returned unchanged. +pub(crate) fn coerce_declared_body( + body: Vec, + declared: &[DeclaredColumn], +) -> crate::Result> { + if declared.is_empty() { + return Ok(body); + } + let mut replaced: Vec<(&str, Vec)> = Vec::new(); + for column in declared { + let Some((start, end)) = nodedb_query::msgpack_scan::extract_field(&body, 0, &column.name) + else { + continue; + }; + let Some(field) = body.get(start..end) else { + continue; + }; + let value = + nodedb_types::value_from_msgpack(field).map_err(|e| crate::Error::Serialization { + format: "msgpack".into(), + detail: format!("column '{}': {e}", column.name), + })?; + let coerced = coerce_declared(&value, column)?; + if coerced == value { + continue; + } + let encoded = + nodedb_types::value_to_msgpack(&coerced).map_err(|e| crate::Error::Serialization { + format: "msgpack".into(), + detail: format!("column '{}': {e}", column.name), + })?; + replaced.push((column.name.as_str(), encoded)); + } + if replaced.is_empty() { + return Ok(body); + } + let pairs: Vec<(&str, &[u8])> = replaced + .iter() + .map(|(name, bytes)| (*name, bytes.as_slice())) + .collect(); + Ok(nodedb_query::msgpack_scan::merge_fields(&body, &pairs)) +} + +/// Re-type every declared field of a decoded document, in place. +pub(crate) fn coerce_declared_doc( + doc: &mut serde_json::Value, + declared: &[DeclaredColumn], +) -> crate::Result<()> { + if declared.is_empty() { + return Ok(()); + } + let Some(object) = doc.as_object_mut() else { + return Ok(()); + }; + for column in declared { + let Some(slot) = object.get_mut(&column.name) else { + continue; + }; + let coerced = coerce_declared(&Value::from(slot.clone()), column)?; + *slot = serde_json::Value::from(coerced); + } + Ok(()) +} + +/// Re-type every declared field of a decoded row, in place. +pub(crate) fn coerce_declared_row( + row: &mut HashMap, + declared: &[DeclaredColumn], +) -> crate::Result<()> { + for column in declared { + let Some(slot) = row.get_mut(&column.name) else { + continue; + }; + *slot = coerce_declared(slot, column)?; + } + Ok(()) +} + +#[cfg(test)] +mod tests { + use super::*; + + fn column(declared: &str) -> DeclaredColumn { + DeclaredColumn::from_declared("v", declared).expect("numeric declaration") + } + + fn body(value: &Value) -> Vec { + let mut map = HashMap::new(); + map.insert("id".to_string(), Value::String("a".into())); + map.insert("v".to_string(), value.clone()); + nodedb_types::value_to_msgpack(&Value::Object(map)).expect("encode body") + } + + fn field_v(body: &[u8]) -> Value { + let row = nodedb_types::value_from_msgpack(body).expect("decode body"); + row.get("v").cloned().expect("field v") + } + + #[test] + fn decimal_rounds_to_scale_and_is_stored_as_text() { + let col = column("DECIMAL(5,2)"); + for (input, expected) in [ + (Value::String("1.005".into()), "1.01"), + (Value::Float(2.5), "2.50"), + (Value::Integer(7), "7.00"), + ] { + assert_eq!( + coerce_declared(&input, &col).expect("fits"), + Value::String(expected.into()), + "{input:?}" + ); + } + } + + #[test] + fn decimal_past_precision_is_numeric_out_of_range() { + let err = coerce_declared(&Value::Float(1500.0), &column("DECIMAL(5,2)")) + .expect_err("1500 has four integer digits"); + assert!( + matches!(err, crate::Error::NumericValueOutOfRange { .. }), + "{err:?}" + ); + } + + #[test] + fn integer_width_refuses_a_value_past_smallint() { + let col = column("SMALLINT"); + assert_eq!( + coerce_declared(&Value::Integer(32767), &col).expect("fits"), + Value::Integer(32767) + ); + let err = coerce_declared(&Value::Integer(40000), &col).expect_err("past smallint"); + assert!( + matches!(err, crate::Error::NumericValueOutOfRange { .. }), + "{err:?}" + ); + let err = coerce_declared(&Value::Float(40000.0), &col).expect_err("whole float"); + assert!( + matches!(err, crate::Error::NumericValueOutOfRange { .. }), + "{err:?}" + ); + } + + #[test] + fn real_refuses_only_an_overflow() { + let col = column("REAL"); + assert_eq!( + coerce_declared(&Value::Float(1.1), &col).expect("rounds"), + Value::Float(1.1) + ); + let err = coerce_declared(&Value::Float(1e300), &col).expect_err("overflows f32"); + assert!( + matches!(err, crate::Error::NumericValueOutOfRange { .. }), + "{err:?}" + ); + } + + #[test] + fn null_passes_through() { + assert_eq!( + coerce_declared(&Value::Null, &column("SMALLINT")).expect("null"), + Value::Null + ); + } + + #[test] + fn body_rewrites_only_declared_fields() { + let declared = [column("DECIMAL(5,2)")]; + let rewritten = + coerce_declared_body(body(&Value::String("1.005".into())), &declared).expect("fits"); + assert_eq!(field_v(&rewritten), Value::String("1.01".into())); + + let unchanged = body(&Value::String("1.50".into())); + assert_eq!( + coerce_declared_body(unchanged.clone(), &declared).expect("fits"), + unchanged + ); + + let err = coerce_declared_body(body(&Value::Integer(123456)), &declared) + .expect_err("past precision"); + assert!( + matches!(err, crate::Error::NumericValueOutOfRange { .. }), + "{err:?}" + ); + } + + #[test] + fn body_without_declared_columns_is_untouched() { + let raw = b"not a map".to_vec(); + assert_eq!(coerce_declared_body(raw.clone(), &[]).expect("noop"), raw); + assert_eq!( + coerce_declared_body(raw.clone(), &[column("SMALLINT")]).expect("not a map"), + raw + ); + } + + #[test] + fn doc_rewrites_declared_fields_in_place() { + let mut doc = serde_json::json!({ "id": "a", "v": "1.005", "other": 9.999 }); + coerce_declared_doc(&mut doc, &[column("DECIMAL(5,2)")]).expect("fits"); + assert_eq!(doc["v"], serde_json::json!("1.01")); + assert_eq!(doc["other"], serde_json::json!(9.999)); + + let mut doc = serde_json::json!({ "id": "a", "v": 1500.0 }); + let err = + coerce_declared_doc(&mut doc, &[column("DECIMAL(5,2)")]).expect_err("past precision"); + assert!( + matches!(err, crate::Error::NumericValueOutOfRange { .. }), + "{err:?}" + ); + } +} diff --git a/nodedb/src/data/executor/strict_format/mod.rs b/nodedb/src/data/executor/strict_format/mod.rs index 2b4625d8d..150ab4015 100644 --- a/nodedb/src/data/executor/strict_format/mod.rs +++ b/nodedb/src/data/executor/strict_format/mod.rs @@ -6,9 +6,15 @@ //! JSON is only produced at the read boundary (`binary_tuple_to_json`) for pgwire clients. mod coerce; +mod declared; mod decode; mod encode; +mod render; +pub(crate) use coerce::{coerce_value, float_to_i64}; +pub(crate) use declared::{ + coerce_declared, coerce_declared_body, coerce_declared_doc, coerce_declared_row, +}; pub(crate) use decode::{ binary_tuple_to_json, binary_tuple_to_msgpack, binary_tuple_to_row_value, binary_tuple_to_value, undecodable_strict_row, @@ -17,3 +23,4 @@ pub(super) use encode::{ bytes_to_binary_tuple, bytes_to_binary_tuple_bitemporal, value_to_binary_tuple, value_to_binary_tuple_bitemporal, }; +pub(crate) use render::strict_row_to_msgpack; diff --git a/nodedb/src/data/executor/wal_replay/kv_put.rs b/nodedb/src/data/executor/wal_replay/kv_put.rs index c8df5d994..ef1ed814f 100644 --- a/nodedb/src/data/executor/wal_replay/kv_put.rs +++ b/nodedb/src/data/executor/wal_replay/kv_put.rs @@ -32,8 +32,11 @@ //! verbatim. Recomputing `now_ms + ttl_ms` at replay time would push every //! expiry forward by the crash-to-restart delay. +use std::borrow::Cow; + use crate::data::executor::core_loop::CoreLoop; use crate::data::executor::core_loop::write_index::KeyRepr; +use crate::data::executor::handlers::kv::declared_body::coerce_kv_body; use crate::data::executor::handlers::transaction::undo::UndoEntry; use nodedb_types::Surrogate; @@ -74,6 +77,18 @@ impl CoreLoop { if self.claim_for_validation() { return Some(0); } + // The record holds the body the write supplied. The live apply + // re-typed its declared numeric columns, so replay does too. + let declared = self.declared_columns_of(database_id, tenant_id, &collection); + let coerced = match coerce_kv_body(&value, declared) { + Ok(Cow::Owned(bytes)) => Some(bytes), + Ok(Cow::Borrowed(_)) => None, + Err(e) => { + self.replay_record_unapplied("kv", "put_declared", record_lsn, &e.to_string()); + return Some(0); + } + }; + let value = coerced.unwrap_or(value); if self.recording_redo_undo() { let prior = self.kv_engine @@ -208,6 +223,29 @@ impl CoreLoop { return Some(0); } + // Each entry holds the body the write supplied. The live apply + // re-typed its declared numeric columns, so replay does too. + let mut entries = entries; + for (_, value) in entries.iter_mut() { + let declared = self.declared_columns_of(database_id, tenant_id, &collection); + let coerced = match coerce_kv_body(value, declared) { + Ok(Cow::Owned(bytes)) => Some(bytes), + Ok(Cow::Borrowed(_)) => None, + Err(e) => { + self.replay_record_unapplied( + "kv", + "batch_put_declared", + record_lsn, + &e.to_string(), + ); + return Some(0); + } + }; + if let Some(bytes) = coerced { + *value = bytes; + } + } + let params = crate::engine::kv::KvBatchPutParams { database_id, tenant_id, diff --git a/nodedb/src/data/executor/wal_replay_kv_atomic.rs b/nodedb/src/data/executor/wal_replay_kv_atomic.rs index e8993b9a7..7b3d64699 100644 --- a/nodedb/src/data/executor/wal_replay_kv_atomic.rs +++ b/nodedb/src/data/executor/wal_replay_kv_atomic.rs @@ -22,6 +22,7 @@ use tracing::warn; use super::core_loop::CoreLoop; use crate::data::executor::core_loop::write_index::KeyRepr; +use crate::data::executor::handlers::kv::declared_body::fit_replayed_kv_image; use crate::engine::kv::{AtomicError, AtomicKeyCtx}; impl CoreLoop { @@ -94,6 +95,12 @@ impl CoreLoop { if self.skip_kv_replay_record(tombstones, tenant_id, &collection, record_lsn) { return Some(0); } + // Replay re-applies a write the policy already admitted when it was + // first accepted. The declared column rule is part of the computation. + let declared = self + .declared_columns_of(database_id, tenant_id, &collection) + .to_vec(); + let gate = |image: &[u8]| fit_replayed_kv_image(image, &declared); let result = self.kv_engine.cas( AtomicKeyCtx { database_id, @@ -105,9 +112,7 @@ impl CoreLoop { }, &expected, &new_value, - // Replay re-applies a write the policy already admitted when it - // was first accepted. - &crate::engine::kv::admit_any, + &gate, ); let swapped = match result { Ok(result) => result.success(), @@ -153,6 +158,14 @@ impl CoreLoop { if self.skip_kv_replay_record(tombstones, tenant_id, &collection, record_lsn) { return Some(0); } + // Replay re-applies a write the policy already admitted when it was + // first accepted; re-deciding it here would make recovery depend on + // the policies of whoever happens to be connected. The declared column + // rule is part of the computation, so it runs again. + let declared = self + .declared_columns_of(database_id, tenant_id, &collection) + .to_vec(); + let gate = |image: &[u8]| fit_replayed_kv_image(image, &declared); match self.kv_engine.incr_float( AtomicKeyCtx { database_id, @@ -164,10 +177,7 @@ impl CoreLoop { }, &delta, &shape, - // Replay re-applies a write the policy already admitted when it was - // first accepted; re-deciding it here would make recovery depend on - // the policies of whoever happens to be connected. - &crate::engine::kv::admit_any, + &gate, ) { Ok(_) => { self.note_replay_write_lsn( @@ -209,8 +219,21 @@ impl CoreLoop { ); Some(0) } - // Unreachable by construction: replay hands the engine - // `admit_any`, so there is no predicate here that could refuse an + // The live write refused the same value against the same + // pre-state, so the record converges to the same no-op. + Err(AtomicError::Declared(error)) => { + warn!( + core = self.core_id, + collection = %collection, + key = %String::from_utf8_lossy(&key), + %error, + "WAL kv_incr_float replay: declared column rule refused the value, \ + skipping record" + ); + Some(0) + } + // Unreachable by construction: the replay gate decides no write + // policy, so there is no predicate here that could refuse an // image. Reaching this arm means a redo path acquired a real // write policy, and recovery would then be re-deciding writes that // were already admitted when they were accepted — against @@ -268,6 +291,12 @@ impl CoreLoop { if self.skip_kv_replay_record(tombstones, tenant_id, &collection, record_lsn) { return Some(0); } + // Replay re-applies a write the policy already admitted when it was + // first accepted. The declared column rule is part of the computation. + let declared = self + .declared_columns_of(database_id, tenant_id, &collection) + .to_vec(); + let gate = |image: &[u8]| fit_replayed_kv_image(image, &declared); let result = self.kv_engine.getset( AtomicKeyCtx { database_id, @@ -278,9 +307,7 @@ impl CoreLoop { surrogate: nodedb_types::Surrogate::new(surrogate), }, &new_value, - // Replay re-applies a write the policy already admitted when it - // was first accepted. - &crate::engine::kv::admit_any, + &gate, ); if let Err(error) = result { self.replay_swap_error("getset", &collection, &key, record_lsn, error); @@ -299,10 +326,11 @@ impl CoreLoop { /// Handle a `cas` or `getset` record whose replay computed no value. /// /// The live apply of the same record against the same pre-state failed - /// the same way and wrote nothing, so replay skips it. A refusal by the - /// write gate cannot happen here: replay hands the engine `admit_any`. - /// Reaching it means a redo path re-decides writes that were already - /// admitted, so replay halts and files a forensic report. + /// the same way and wrote nothing, so replay skips it. That covers a + /// declared column rule refusal too. A refusal by the write policy cannot + /// happen here: the replay gate decides no policy. Reaching it means a + /// redo path re-decides writes that were already admitted, so replay + /// halts and files a forensic report. fn replay_swap_error( &mut self, op: &'static str, @@ -314,6 +342,7 @@ impl CoreLoop { let detail = match error { AtomicError::TypeMismatch { detail } | AtomicError::Encode { detail } => detail, AtomicError::Counter(fault) => fault.message().to_string(), + AtomicError::Declared(error) => error.to_string(), AtomicError::Rejected(error) => { self.replay_record_unapplied( "kv", diff --git a/nodedb/src/data/executor/wal_replay_kv_field.rs b/nodedb/src/data/executor/wal_replay_kv_field.rs index cb58c2f08..ef0a6e167 100644 --- a/nodedb/src/data/executor/wal_replay_kv_field.rs +++ b/nodedb/src/data/executor/wal_replay_kv_field.rs @@ -58,7 +58,12 @@ impl CoreLoop { if if_present && current.is_none() { return Some(0); } - let computed = match merge_field_updates(&collection, current.as_deref(), &updates) { + let computed = match merge_field_updates( + &collection, + current.as_deref(), + &updates, + self.declared_columns_of(database_id, tenant_id, &collection), + ) { Ok(c) => c, Err(e) => { warn!( @@ -218,6 +223,7 @@ mod tests { "players", Some(&seed), &[("mana".to_string(), json_field_bytes(serde_json::json!(5)))], + &[], ) .expect("live merge") .new_value; @@ -252,6 +258,7 @@ mod tests { "players", None, &[("hp".to_string(), json_field_bytes(serde_json::json!(100)))], + &[], ) .expect("live merge") .new_value; @@ -298,6 +305,7 @@ mod tests { "players", Some(b"42"), &[("hp".to_string(), json_field_bytes(serde_json::json!(1)))], + &[], ); assert!( matches!( diff --git a/nodedb/src/data/executor/wal_replay_kv_incr.rs b/nodedb/src/data/executor/wal_replay_kv_incr.rs index 7bf5ff4c3..5a53d7453 100644 --- a/nodedb/src/data/executor/wal_replay_kv_incr.rs +++ b/nodedb/src/data/executor/wal_replay_kv_incr.rs @@ -26,6 +26,7 @@ use tracing::warn; use super::core_loop::CoreLoop; use crate::data::executor::core_loop::write_index::KeyRepr; +use crate::data::executor::handlers::kv::declared_body::fit_replayed_kv_image; use crate::engine::kv::{AtomicError, AtomicKeyCtx, IncrStep, Incremented}; /// The decoded `kv_incr` record. @@ -76,8 +77,13 @@ impl CoreLoop { }; // Replay re-applies a write the policy already admitted when it was // first accepted. Re-deciding it here would make recovery depend on - // the policies of whoever happens to be connected. - let admit = &crate::engine::kv::admit_any; + // the policies of whoever happens to be connected. The declared + // column rule is part of the computation, so it runs again. + let declared = self + .declared_columns_of(database_id, tenant_id, &collection) + .to_vec(); + let gate = |image: &[u8]| fit_replayed_kv_image(image, &declared); + let admit: crate::engine::kv::AtomicImageGate<'_> = &gate; let result = match expire_at_ms { Some(expire_at_ms) => self.kv_engine.incr_with_absolute_expiry( ctx, @@ -105,7 +111,7 @@ impl CoreLoop { } /// Shared result handling for a `kv_incr` replay: `Ok` counts as one - /// applied put; `TypeMismatch` / `Counter` / `Encode` are + /// applied put; `TypeMismatch` / `Counter` / `Encode` / `Declared` are /// correctly-converging no-ops (the live dispatch would have failed /// identically), logged and skipped rather than treated as errors. /// @@ -155,8 +161,21 @@ impl CoreLoop { ); 0 } - // Unreachable by construction: replay hands the engine - // `admit_any`, so there is no predicate here that could refuse an + // The live write refused the same value against the same + // pre-state, so the record converges to the same no-op. + Err(AtomicError::Declared(error)) => { + warn!( + core = self.core_id, + collection = %collection, + key = %String::from_utf8_lossy(key), + delta, + %error, + "WAL kv_incr replay: declared column rule refused the value, skipping record" + ); + 0 + } + // Unreachable by construction: the replay gate decides no write + // policy, so there is no predicate here that could refuse an // image. Reaching this arm means a redo path acquired a real // write policy, and recovery would then be re-deciding writes that // were already admitted when they were accepted — against diff --git a/nodedb/src/data/executor/wal_replay_kv_insert_conflict.rs b/nodedb/src/data/executor/wal_replay_kv_insert_conflict.rs index b619c9908..b989b4f0f 100644 --- a/nodedb/src/data/executor/wal_replay_kv_insert_conflict.rs +++ b/nodedb/src/data/executor/wal_replay_kv_insert_conflict.rs @@ -96,9 +96,9 @@ impl CoreLoop { ) } - /// RMW + write-back: absent key installs `value` - /// verbatim (the live handler's insert branch); present key re-runs - /// `merge_kv_conflict_body`, the exact merge the live handler uses. A + /// RMW + write-back through `merge_kv_conflict_body`, the exact post-image + /// the live handler stores: the incoming `value` for an absent key, the + /// merge for a present one. A /// merge failure is logged and the record is skipped rather than /// fabricating a value. fn apply_replayed_insert_on_conflict_update( @@ -122,27 +122,29 @@ impl CoreLoop { .kv_engine .get(database_id, tenant_id, collection, key, now_ms); - let stored_bytes: Vec = match &existing_bytes { - None => value.to_vec(), - Some(existing_raw) => match merge_kv_conflict_body(existing_raw, value, updates) { - Ok(b) => b, - Err(e) => { - // A division/modulo-by-zero here can only come from a - // record logged by a build that did not fail the - // statement at execution time; a shape or decode error - // means the durable bytes no longer hold what the record - // expects. Either way the record is skipped, never - // fabricated, and startup continues. - warn!( - core = self.core_id, - collection = %collection, - key = %String::from_utf8_lossy(key), - ?e, - "WAL kv_insert_on_conflict_update replay: merge failed, skipping record" - ); - return 0; - } - }, + let stored_bytes = match merge_kv_conflict_body( + existing_bytes.as_deref(), + value, + updates, + self.declared_columns_of(database_id, tenant_id, collection), + ) { + Ok(b) => b, + Err(e) => { + // A division/modulo-by-zero here can only come from a + // record logged by a build that did not fail the + // statement at execution time; a shape or decode error + // means the durable bytes no longer hold what the record + // expects. Either way the record is skipped, never + // fabricated, and startup continues. + warn!( + core = self.core_id, + collection = %collection, + key = %String::from_utf8_lossy(key), + ?e, + "WAL kv_insert_on_conflict_update replay: merge failed, skipping record" + ); + return 0; + } }; let params = crate::engine::kv::KvPutParams { diff --git a/nodedb/src/data/executor/wal_replay_kv_predicate.rs b/nodedb/src/data/executor/wal_replay_kv_predicate.rs index fbbf9ed38..4fe215c5d 100644 --- a/nodedb/src/data/executor/wal_replay_kv_predicate.rs +++ b/nodedb/src/data/executor/wal_replay_kv_predicate.rs @@ -56,7 +56,12 @@ impl CoreLoop { let mut written = 0usize; for (key, body) in matched { - let computed = match merge_field_updates(&collection, Some(body.as_slice()), &updates) { + let computed = match merge_field_updates( + &collection, + Some(body.as_slice()), + &updates, + self.declared_columns_of(database_id, tenant_id, &collection), + ) { Ok(c) => c, Err(e) => { warn!( diff --git a/nodedb/src/engine/document/store/config.rs b/nodedb/src/engine/document/store/config.rs index 409633064..26966d74b 100644 --- a/nodedb/src/engine/document/store/config.rs +++ b/nodedb/src/engine/document/store/config.rs @@ -41,6 +41,20 @@ pub struct CollectionConfig { /// MessagePack maps with the same header byte, so byte sniffing silently /// returns tag arrays to the client. pub 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. Empty for every other collection: a strict schema + /// carries its vector columns in `storage_mode`. + pub 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: strict and columnar re-type through their own + /// typed schema. + pub declared_columns: Vec, + /// The declared key column, per `nodedb_types::declared_key`. `None` when + /// rows key by the implicit `id` or `_rowid`. A sparse row's identity + /// renders under this column, else under `id`. + pub declared_key: Option, } impl CollectionConfig { @@ -55,6 +69,9 @@ impl CollectionConfig { conflict_policy: None, timeseries: None, vector_primary: None, + vector_fields: Vec::new(), + declared_columns: Vec::new(), + declared_key: None, } } diff --git a/nodedb/src/engine/kv/counter_refit.rs b/nodedb/src/engine/kv/counter_refit.rs new file mode 100644 index 000000000..b12c1d4fa --- /dev/null +++ b/nodedb/src/engine/kv/counter_refit.rs @@ -0,0 +1,91 @@ +// SPDX-License-Identifier: BUSL-1.1 + +//! The counter value an `INCR_FLOAT` stores once a declared column rule has +//! fitted the image it computed. + +use nodedb_query::msgpack_scan::{KvBodyShape, kv_body_shape}; +use nodedb_types::Value; + +/// The counter value the image `fitted` stores. +/// +/// `value` and `computed` are what `INCR_FLOAT` computed. `fitted` is the +/// image a declared column rule made of `computed`, such as a `DECIMAL(p,s)` +/// column rounding to its scale. A raw body stores the counter as its text. +/// A typed row stores it in the column whose computed cell is `value`, and +/// that column's fitted cell is read as a number. When no cell reads as a +/// number, `value` is the answer: the rule kept it. +pub fn fitted_counter_f64(value: f64, computed: &[u8], fitted: &[u8]) -> f64 { + if kv_body_shape(fitted) == KvBodyShape::Raw { + return std::str::from_utf8(fitted) + .ok() + .and_then(|text| text.trim().parse::().ok()) + .unwrap_or(value); + } + let (Ok(Value::Object(before)), Ok(Value::Object(after))) = ( + nodedb_types::value_from_msgpack(computed), + nodedb_types::value_from_msgpack(fitted), + ) else { + return value; + }; + let mut moved: Vec<&String> = before + .iter() + .filter(|(_, cell)| matches!(cell, Value::Float(f) if *f == value)) + .map(|(name, _)| name) + .collect(); + moved.sort(); + moved + .into_iter() + .filter_map(|name| after.get(name)) + .find_map(cell_f64) + .unwrap_or(value) +} + +/// A stored numeric cell as `f64`. A decimal column stores its value as text. +fn cell_f64(cell: &Value) -> Option { + match cell { + Value::Float(f) => Some(*f), + Value::Integer(i) => Some(*i as f64), + Value::String(text) => text.trim().parse().ok(), + Value::Decimal(d) => d.to_string().parse().ok(), + _ => None, + } +} + +#[cfg(test)] +mod tests { + use std::collections::HashMap; + + use super::*; + + fn row(fields: &[(&str, Value)]) -> Vec { + let map: HashMap = fields + .iter() + .map(|(name, value)| ((*name).to_string(), value.clone())) + .collect(); + nodedb_types::value_to_msgpack(&Value::Object(map)).expect("encode row") + } + + #[test] + fn raw_body_reads_the_fitted_text() { + assert_eq!(fitted_counter_f64(1.005, b"1.005", b"1.01"), 1.01); + } + + #[test] + fn typed_row_reads_the_fitted_counter_cell() { + let computed = row(&[ + ("label", Value::String("gold".into())), + ("price", Value::Float(1.005)), + ]); + let fitted = row(&[ + ("label", Value::String("gold".into())), + ("price", Value::String("1.01".into())), + ]); + assert_eq!(fitted_counter_f64(1.005, &computed, &fitted), 1.01); + } + + #[test] + fn an_unchanged_counter_keeps_its_value() { + let image = row(&[("n", Value::Float(2.5))]); + assert_eq!(fitted_counter_f64(2.5, &image, &image), 2.5); + } +} diff --git a/nodedb/src/engine/kv/engine_atomic.rs b/nodedb/src/engine/kv/engine_atomic.rs index 5b6275d69..92a26881f 100644 --- a/nodedb/src/engine/kv/engine_atomic.rs +++ b/nodedb/src/engine/kv/engine_atomic.rs @@ -9,6 +9,7 @@ use nodedb_physical::kv_atomic::{AtomicComputeError, compute}; use nodedb_physical::physical_plan::KvCounterShape; +use super::counter_refit::fitted_counter_f64; use super::engine::KvEngine; use super::engine_helpers::{expiry_key, table_key}; use super::engine_write::{UnboundKvWrite, require_bound}; @@ -75,8 +76,13 @@ pub enum AtomicError { Counter(crate::bridge::envelope::CounterFault), /// The computed new value failed to re-encode as MessagePack. Encode { detail: String }, - /// The [`AtomicAdmission`] gate refused the computed post-image, so nothing + /// The computed value breaks a declared column rule of the collection, + /// such as a `SMALLINT` range or a `DECIMAL(p,s)` precision, so nothing /// was written. Boxed to keep the error small on the success path. + Declared(Box), + /// The row-level-security write policy refused the computed post-image, + /// so nothing was written. Boxed to keep the error small on the success + /// path. Rejected(Box), /// The write binds its row to `Surrogate::ZERO`, so nothing was written. Unbound(UnboundKvWrite), @@ -101,21 +107,28 @@ impl From for AtomicError { /// A gate consulted with the computed post-image before an atomic commits. /// /// Every atomic computes the value it stores from the stored one: INCR runs -/// arithmetic, and CAS and GETSET swap one column of a typed row. The row a -/// row-level-security write policy has to decide does not exist until that -/// computation has run, and it runs here, inside the engine, in the same pass -/// that persists the result. Passing the decision in keeps the computation in -/// one place. -pub type AtomicAdmission<'a> = &'a dyn Fn(&[u8]) -> crate::Result<()>; - -/// An admission that accepts every image. +/// arithmetic, and CAS and GETSET swap one column of a typed row. The image +/// the collection's declared column rule fits, and the image a row-level +/// security write policy decides, do not exist until that computation has +/// run. It runs here, inside the engine, in the same pass that persists the +/// result. Passing the gate in keeps the computation in one place. /// -/// For replaying a write that was already decided: a WAL redo re-applies a -/// write whose policy verdict was reached when it was first accepted, and -/// re-deciding it against the *current* session's policies would make recovery -/// depend on who happens to be connected. -pub fn admit_any(_image: &[u8]) -> crate::Result<()> { - Ok(()) +/// The gate answers `Ok(None)` to store the computed image as it is, +/// `Ok(Some(image))` to store the image the declared rule fitted, and `Err` +/// to store nothing. +pub type AtomicImageGate<'a> = &'a dyn Fn(&[u8]) -> Result>, AtomicError>; + +/// A gate that stores every computed image as it is. +/// +/// For a collection with no declared column and no write policy, and for +/// tests. +pub fn admit_any(_image: &[u8]) -> Result>, AtomicError> { + Ok(None) +} + +/// The image `gate` decides to store for the computed `image`. +fn gated(gate: AtomicImageGate<'_>, image: Vec) -> Result, AtomicError> { + Ok(gate(&image)?.unwrap_or(image)) } /// Shared key-identity context for a single-key atomic KV operation @@ -151,15 +164,15 @@ impl KvEngine { /// - On i64 overflow: returns `Counter(IntegerOverflow)`. It never wraps. /// - TTL behavior: if `ttl_ms > 0` and key is new, sets TTL. /// If key exists and `ttl_ms > 0`, resets TTL. If `ttl_ms == 0`, preserves. - /// - If `admit` refuses the computed value: returns `Rejected` and writes - /// nothing. + /// - `admit` fits the computed row to the declared columns and decides + /// it. A refusal returns its error and writes nothing. pub fn incr( &mut self, ctx: AtomicKeyCtx<'_>, delta: i64, ttl_ms: u64, shape: &KvCounterShape, - admit: AtomicAdmission<'_>, + admit: AtomicImageGate<'_>, ) -> Result, AtomicError> { self.incr_resolved( ctx, @@ -190,7 +203,7 @@ impl KvEngine { ctx: AtomicKeyCtx<'_>, step: IncrStep<'_>, expire_at_ms: u64, - admit: AtomicAdmission<'_>, + admit: AtomicImageGate<'_>, ) -> Result, AtomicError> { self.incr_resolved(ctx, step, Some(expire_at_ms), admit) } @@ -203,7 +216,7 @@ impl KvEngine { ctx: AtomicKeyCtx<'_>, step: IncrStep<'_>, expire_override: Option, - admit: AtomicAdmission<'_>, + admit: AtomicImageGate<'_>, ) -> Result, AtomicError> { let IncrStep { delta, @@ -215,10 +228,11 @@ impl KvEngine { let table = self.ensure_table(tkey, ctx.tenant_id, ctx.collection); let current = table.get(ctx.key, ctx.now_ms).map(|v| v.to_vec()); - let (value, written) = compute::incr(current.as_deref(), delta, shape)?; + let (value, computed) = compute::incr(current.as_deref(), delta, shape)?; // Decided before `atomic_put`, so a refused image is never durable and - // never reaches the expiry wheel or the secondary indexes. - admit(&written).map_err(|error| AtomicError::Rejected(Box::new(error)))?; + // never reaches the expiry wheel or the secondary indexes. A declared + // rule never changes an integer it accepts, so `value` stands. + let written = gated(admit, computed)?; self.atomic_put( ctx, tkey, @@ -243,23 +257,28 @@ impl KvEngine { /// `Counter(NotAFloat)`. A typed row without a numeric column: returns /// `TypeMismatch`. /// - A NaN or infinite result: returns `Counter(NonFinite)`. - /// - If `admit` refuses the computed value: returns `Rejected` and writes - /// nothing. + /// - `admit` fits the computed row to the declared columns and decides + /// it. A refusal returns its error and writes nothing. The returned + /// value is the one the fitted row stores: a `DECIMAL(p,s)` column + /// rounds it to its scale. pub fn incr_float( &mut self, ctx: AtomicKeyCtx<'_>, delta: &str, shape: &KvCounterShape, - admit: AtomicAdmission<'_>, + admit: AtomicImageGate<'_>, ) -> Result, AtomicError> { require_bound(ctx.collection, ctx.surrogate)?; let tkey = table_key(ctx.database_id, ctx.tenant_id, ctx.collection); let table = self.ensure_table(tkey, ctx.tenant_id, ctx.collection); let current = table.get(ctx.key, ctx.now_ms).map(|v| v.to_vec()); - let (value, written) = compute::incr_float(current.as_deref(), delta, shape)?; + let (value, computed) = compute::incr_float(current.as_deref(), delta, shape)?; // Decided before the value is installed — see `incr_resolved`. - admit(&written).map_err(|error| AtomicError::Rejected(Box::new(error)))?; + let (value, written) = match admit(&computed)? { + None => (value, computed), + Some(fitted) => (fitted_counter_f64(value, &computed, &fitted), fitted), + }; // incr_float always preserves existing TTL (ttl_ms = 0). self.atomic_put(ctx, tkey, &written, 0, current.is_none(), None); @@ -271,14 +290,14 @@ impl KvEngine { /// If current value equals `expected`, sets to `new_value` and returns success. /// If current value differs, returns the actual current value. /// If key doesn't exist and `expected` is empty, creates the key (create-if-not-exists). - /// If `admit` refuses the bytes the swap would store: returns `Rejected` - /// and writes nothing. + /// `admit` fits the bytes the swap would store to the declared columns + /// and decides them. A refusal returns its error and writes nothing. pub fn cas( &mut self, ctx: AtomicKeyCtx<'_>, expected: &[u8], new_value: &[u8], - admit: AtomicAdmission<'_>, + admit: AtomicImageGate<'_>, ) -> Result { require_bound(ctx.collection, ctx.surrogate)?; let tkey = table_key(ctx.database_id, ctx.tenant_id, ctx.collection); @@ -294,7 +313,7 @@ impl KvEngine { }); } // Decided before the value is installed — see `incr_resolved`. - admit(&write_bytes).map_err(|error| AtomicError::Rejected(Box::new(error)))?; + let write_bytes = gated(admit, write_bytes)?; self.atomic_put(ctx, tkey, &write_bytes, 0, current.is_none(), None); Ok(CasResult { written: Some(write_bytes), @@ -306,13 +325,13 @@ impl KvEngine { /// /// If key didn't exist, `old` is `None`. /// Preserves existing TTL. - /// If `admit` refuses the bytes the write would store: returns `Rejected` - /// and writes nothing. + /// `admit` fits the bytes the write would store to the declared columns + /// and decides them. A refusal returns its error and writes nothing. pub fn getset( &mut self, ctx: AtomicKeyCtx<'_>, new_value: &[u8], - admit: AtomicAdmission<'_>, + admit: AtomicImageGate<'_>, ) -> Result { require_bound(ctx.collection, ctx.surrogate)?; let tkey = table_key(ctx.database_id, ctx.tenant_id, ctx.collection); @@ -320,7 +339,7 @@ impl KvEngine { let old = table.get(ctx.key, ctx.now_ms).map(|v| v.to_vec()); let write_bytes = compute::getset(old.as_deref(), new_value)?; // Decided before the value is installed — see `incr_resolved`. - admit(&write_bytes).map_err(|error| AtomicError::Rejected(Box::new(error)))?; + let write_bytes = gated(admit, write_bytes)?; // GetSet preserves existing TTL (ttl_ms = 0). self.atomic_put(ctx, tkey, &write_bytes, 0, old.is_none(), None); @@ -563,11 +582,13 @@ mod tests { .incr(ctx("counters", b"hits"), 7, 0, &RAW, &admit_any) .unwrap(); - let deny = |_: &[u8]| { - Err(crate::Error::RejectedAuthz { - tenant_id: crate::types::TenantId::new(1), - resource: "test".into(), - }) + let deny = |_: &[u8]| -> Result>, AtomicError> { + Err(AtomicError::Rejected(Box::new( + crate::Error::RejectedAuthz { + tenant_id: crate::types::TenantId::new(1), + resource: "test".into(), + }, + ))) }; let result = engine.incr(ctx("counters", b"hits"), 5, 0, &RAW, &deny); assert!(matches!(result, Err(AtomicError::Rejected(_)))); @@ -582,6 +603,60 @@ mod tests { ); } + /// The gate's fitted image is what the engine stores, and the counter + /// value it returns is the fitted one. + #[test] + fn a_fitted_image_is_stored_and_reported() { + let mut engine = make_engine(); + let round = + |_: &[u8]| -> Result>, AtomicError> { Ok(Some(b"1.01".to_vec())) }; + let result = engine + .incr_float(ctx("scores", b"price"), "1.005", &RAW, &round) + .expect("fitted"); + assert_eq!(result.written, b"1.01".to_vec()); + assert_eq!(result.value, 1.01); + assert_eq!( + engine.get(0, 1, "scores", b"price", 1000).as_deref(), + Some(b"1.01".as_slice()) + ); + } + + /// A declared-rule refusal stores nothing for every atomic. + #[test] + fn a_declared_refusal_writes_nothing() { + let mut engine = make_engine(); + engine + .incr(ctx("counters", b"hits"), 7, 0, &RAW, &admit_any) + .expect("seed"); + let refuse = |_: &[u8]| -> Result>, AtomicError> { + Err(AtomicError::Declared(Box::new( + crate::Error::NumericValueOutOfRange { + detail: "test".into(), + }, + ))) + }; + assert!(matches!( + engine.incr(ctx("counters", b"hits"), 1, 0, &RAW, &refuse), + Err(AtomicError::Declared(_)) + )); + assert!(matches!( + engine.incr_float(ctx("counters", b"hits"), "1", &RAW, &refuse), + Err(AtomicError::Declared(_)) + )); + assert!(matches!( + engine.cas(ctx("counters", b"hits"), b"7", b"9", &refuse), + Err(AtomicError::Declared(_)) + )); + assert!(matches!( + engine.getset(ctx("counters", b"hits"), b"9", &refuse), + Err(AtomicError::Declared(_)) + )); + assert_eq!( + engine.get(0, 1, "counters", b"hits", 1000).as_deref(), + Some(b"7".as_slice()) + ); + } + #[test] fn incr_overflow() { let mut engine = make_engine(); diff --git a/nodedb/src/engine/kv/mod.rs b/nodedb/src/engine/kv/mod.rs index e4116d61b..a41f46795 100644 --- a/nodedb/src/engine/kv/mod.rs +++ b/nodedb/src/engine/kv/mod.rs @@ -2,6 +2,7 @@ mod batch_put; mod clock; +mod counter_refit; pub mod engine; pub mod engine_atomic; mod engine_helpers; @@ -22,11 +23,12 @@ pub(crate) mod test_support; pub use batch_put::KvBatchPutParams; pub use clock::current_ms; +pub use counter_refit::fitted_counter_f64; pub use engine::{ KvEngine, KvEntryImage, KvKeyRef, RestoreCompositeIndexParams, RestoreFieldIndexParams, }; pub use engine_atomic::{ - AtomicAdmission, AtomicError, AtomicKeyCtx, CasResult, GetSetResult, IncrStep, Incremented, + AtomicError, AtomicImageGate, AtomicKeyCtx, CasResult, GetSetResult, IncrStep, Incremented, admit_any, }; pub use engine_index::RegisterIndexParams; diff --git a/nodedb/tests/inproc/cases/executor_tests/test_conflict_policy_register.rs b/nodedb/tests/inproc/cases/executor_tests/test_conflict_policy_register.rs index 023259194..2bfab7fab 100644 --- a/nodedb/tests/inproc/cases/executor_tests/test_conflict_policy_register.rs +++ b/nodedb/tests/inproc/cases/executor_tests/test_conflict_policy_register.rs @@ -51,6 +51,9 @@ fn register(ctx: &mut TestCtx, collection: &str, conflict_policy: Option conflict_policy, timeseries: None, vector_primary: None, + vector_fields: Vec::new(), + declared_columns: Vec::new(), + declared_key: None, }), ); assert_eq!(resp.status, Status::Ok, "register document collection"); diff --git a/nodedb/tests/inproc/cases/executor_tests/test_generated_columns.rs b/nodedb/tests/inproc/cases/executor_tests/test_generated_columns.rs index 22493bae2..85e7e1906 100644 --- a/nodedb/tests/inproc/cases/executor_tests/test_generated_columns.rs +++ b/nodedb/tests/inproc/cases/executor_tests/test_generated_columns.rs @@ -44,6 +44,9 @@ fn register_with_generated( conflict_policy: None, timeseries: None, vector_primary: None, + vector_fields: Vec::new(), + declared_columns: Vec::new(), + declared_key: None, }), ); } diff --git a/nodedb/tests/inproc/cases/executor_tests/test_range_scan_bitemporal.rs b/nodedb/tests/inproc/cases/executor_tests/test_range_scan_bitemporal.rs index 11415dc9c..d3ba2419f 100644 --- a/nodedb/tests/inproc/cases/executor_tests/test_range_scan_bitemporal.rs +++ b/nodedb/tests/inproc/cases/executor_tests/test_range_scan_bitemporal.rs @@ -45,6 +45,9 @@ fn register_schemaless_bitemporal(ctx: &mut TestCtx, collection: &str) { conflict_policy: None, timeseries: None, vector_primary: None, + vector_fields: Vec::new(), + declared_columns: Vec::new(), + declared_key: None, }), ); assert_eq!(resp.status, Status::Ok, "register schemaless bitemporal"); @@ -78,6 +81,9 @@ fn register_strict_bitemporal(ctx: &mut TestCtx, collection: &str) { conflict_policy: None, timeseries: None, vector_primary: None, + vector_fields: Vec::new(), + declared_columns: Vec::new(), + declared_key: None, }), ); assert_eq!(resp.status, Status::Ok, "register strict bitemporal"); diff --git a/nodedb/tests/wire/cases/declared_width_computed_writes.rs b/nodedb/tests/wire/cases/declared_width_computed_writes.rs new file mode 100644 index 000000000..003aeb8c2 --- /dev/null +++ b/nodedb/tests/wire/cases/declared_width_computed_writes.rs @@ -0,0 +1,312 @@ +// SPDX-License-Identifier: BUSL-1.1 + +//! A declared numeric column refuses a value past its declared width on every +//! write that stores one, with SQLSTATE 22003, and the row keeps its value. +//! +//! - Strict and columnar schemas carry the declared `SMALLINT` / `REAL` +//! width. An INSERT and a computed UPDATE meet the same rule. +//! - KV operations that compute the value they store (`KV_INCR`, +//! `KV_INCR_FLOAT`, `KV_CAS`, `KV_GETSET`, `TRANSFER`) meet the +//! collection's declared columns before they store it. + +use crate::harness::TestServer; + +const OUT_OF_RANGE: &str = "SQLSTATE 22003"; + +/// The single cell `sql` reads. +async fn cell(srv: &TestServer, sql: &str) -> String { + let rows = srv.query_rows(sql).await.unwrap(); + assert_eq!(rows.len(), 1, "{sql}: {rows:?}"); + rows[0][0].clone() +} + +/// The strict and columnar `CREATE` statements for a collection `name` with +/// an `id` key and a `v` column of `declared` type. +fn typed_engines(name: &str, declared: &str) -> [(String, String); 2] { + [ + ( + format!("{name}_strict"), + format!( + "CREATE COLLECTION {name}_strict (id TEXT PRIMARY KEY, v {declared}) \ + WITH (engine='document_strict')" + ), + ), + ( + format!("{name}_columnar"), + format!( + "CREATE COLLECTION {name}_columnar COLUMNS (id TEXT, v {declared}) \ + WITH (engine='columnar')" + ), + ), + ] +} + +/// Assert `name`'s row `'a'` still holds `1` after a refused UPDATE. +async fn assert_row_kept(srv: &TestServer, name: &str) { + assert_eq!( + cell(srv, &format!("SELECT v FROM {name} WHERE id = 'a'")).await, + "1", + "{name}: the refused UPDATE keeps the row's value" + ); +} + +/// A strict or columnar `SMALLINT` column refuses `40000` on INSERT and on +/// UPDATE, and advertises `int2` on the wire. +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn smallint_refuses_a_value_past_its_width_on_insert_and_update() { + let srv = TestServer::start().await; + for (name, create) in typed_engines("dw_small", "SMALLINT") { + srv.exec(&create).await.unwrap(); + srv.exec(&format!("INSERT INTO {name} (id, v) VALUES ('a', 1)")) + .await + .unwrap(); + + srv.expect_error( + &format!("INSERT INTO {name} (id, v) VALUES ('b', 40000)"), + OUT_OF_RANGE, + ) + .await; + srv.expect_error( + &format!("UPDATE {name} SET v = 40000 WHERE id = 'a'"), + OUT_OF_RANGE, + ) + .await; + assert_row_kept(&srv, &name).await; + assert!( + srv.query_rows(&format!("SELECT v FROM {name} WHERE id = 'b'")) + .await + .unwrap() + .is_empty(), + "{name}: the refused INSERT stores no row" + ); + + let stmt = srv + .client + .prepare(&format!("SELECT v FROM {name} LIMIT 0")) + .await + .expect("prepare describe"); + assert_eq!( + stmt.columns()[0].type_().oid(), + 21, + "{name}: a SMALLINT column advertises int2" + ); + } +} + +/// A computed UPDATE past a strict `SMALLINT` is refused. The planner +/// validates only a literal, so the Data Plane encoder enforces the width. +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn strict_smallint_refuses_a_computed_update_past_its_width() { + let srv = TestServer::start().await; + let [(name, create), _] = typed_engines("dw_cstrict", "SMALLINT"); + srv.exec(&create).await.unwrap(); + srv.exec(&format!("INSERT INTO {name} (id, v) VALUES ('a', 1)")) + .await + .unwrap(); + srv.expect_error( + &format!("UPDATE {name} SET v = v + 39999 WHERE id = 'a'"), + OUT_OF_RANGE, + ) + .await; + assert_row_kept(&srv, &name).await; +} + +/// A computed UPDATE past a columnar `SMALLINT` is refused. The columnar +/// UPDATE handlers fit every post-image to the schema before the first row +/// changes. +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn columnar_smallint_refuses_a_computed_update_past_its_width() { + let srv = TestServer::start().await; + let [_, (name, create)] = typed_engines("dw_ccol", "SMALLINT"); + srv.exec(&create).await.unwrap(); + srv.exec(&format!("INSERT INTO {name} (id, v) VALUES ('a', 1)")) + .await + .unwrap(); + srv.expect_error( + &format!("UPDATE {name} SET v = v + 39999 WHERE id = 'a'"), + OUT_OF_RANGE, + ) + .await; + assert_row_kept(&srv, &name).await; +} + +/// A strict or columnar `REAL` column refuses `1e39`, past `f32`, on INSERT +/// and on UPDATE. +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn real_refuses_a_value_past_f32() { + let srv = TestServer::start().await; + for (name, create) in typed_engines("dw_real", "REAL") { + srv.exec(&create).await.unwrap(); + srv.exec(&format!("INSERT INTO {name} (id, v) VALUES ('a', 1.5)")) + .await + .unwrap(); + + srv.expect_error( + &format!("INSERT INTO {name} (id, v) VALUES ('b', 1e39)"), + OUT_OF_RANGE, + ) + .await; + srv.expect_error( + &format!("UPDATE {name} SET v = 1e39 WHERE id = 'a'"), + OUT_OF_RANGE, + ) + .await; + + assert_eq!( + cell(&srv, &format!("SELECT v FROM {name} WHERE id = 'a'")).await, + "1.5", + "{name}" + ); + } +} + +/// Narrowing a strict column's declared width is refused: existing rows can +/// already hold values past the narrower width. +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn narrowing_a_declared_width_is_refused() { + let srv = TestServer::start().await; + srv.exec( + "CREATE COLLECTION dw_narrow (id TEXT PRIMARY KEY, v INT) \ + WITH (engine='document_strict')", + ) + .await + .unwrap(); + srv.exec("INSERT INTO dw_narrow (id, v) VALUES ('a', 40000)") + .await + .unwrap(); + srv.expect_error( + "ALTER COLLECTION dw_narrow ALTER COLUMN v TYPE SMALLINT", + "SQLSTATE 0A000", + ) + .await; + assert_eq!( + cell(&srv, "SELECT v FROM dw_narrow WHERE id = 'a'").await, + "40000" + ); +} + +/// `KV_INCR` past a declared `SMALLINT` is refused, and the row keeps its +/// value. +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn kv_incr_past_a_declared_smallint_is_refused() { + let srv = TestServer::start().await; + srv.exec( + "CREATE COLLECTION dw_kv_incr (key TEXT PRIMARY KEY, n SMALLINT, label TEXT) \ + WITH (engine='kv')", + ) + .await + .unwrap(); + srv.exec("INSERT INTO dw_kv_incr (key, n, label) VALUES ('a', 32767, 'x')") + .await + .unwrap(); + + srv.expect_error("SELECT KV_INCR('dw_kv_incr', 'a', 1)", OUT_OF_RANGE) + .await; + assert_eq!( + cell(&srv, "SELECT n FROM dw_kv_incr WHERE key = 'a'").await, + "32767" + ); + + srv.query_text("SELECT KV_INCR('dw_kv_incr', 'a', -1)") + .await + .unwrap(); + assert_eq!( + cell(&srv, "SELECT n FROM dw_kv_incr WHERE key = 'a'").await, + "32766" + ); +} + +/// `KV_INCR_FLOAT` on a declared `DECIMAL(5,2)` value rounds to the declared +/// scale, and a result past the declared precision is refused. +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn kv_incr_float_meets_a_declared_decimal() { + let srv = TestServer::start().await; + srv.exec( + "CREATE COLLECTION dw_kv_dec (key TEXT PRIMARY KEY, value DECIMAL(5,2)) \ + WITH (engine='kv')", + ) + .await + .unwrap(); + srv.exec("INSERT INTO dw_kv_dec (key, value) VALUES ('a', 999.99)") + .await + .unwrap(); + srv.exec("INSERT INTO dw_kv_dec (key, value) VALUES ('b', 1.50)") + .await + .unwrap(); + + srv.expect_error("SELECT KV_INCR_FLOAT('dw_kv_dec', 'a', 1)", OUT_OF_RANGE) + .await; + assert_eq!( + cell(&srv, "SELECT value FROM dw_kv_dec WHERE key = 'a'").await, + "999.99" + ); + + srv.query_text("SELECT KV_INCR_FLOAT('dw_kv_dec', 'b', 0.005)") + .await + .unwrap(); + assert_eq!( + cell(&srv, "SELECT value FROM dw_kv_dec WHERE key = 'b'").await, + "1.51" + ); +} + +/// `KV_CAS` and `KV_GETSET` that would store a value past a declared +/// `SMALLINT` are refused, and the row keeps its value. +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn kv_cas_and_getset_past_a_declared_smallint_are_refused() { + let srv = TestServer::start().await; + srv.exec( + "CREATE COLLECTION dw_kv_swap (key TEXT PRIMARY KEY, value SMALLINT) \ + WITH (engine='kv')", + ) + .await + .unwrap(); + srv.exec("INSERT INTO dw_kv_swap (key, value) VALUES ('a', 5)") + .await + .unwrap(); + + srv.expect_error( + "SELECT KV_CAS('dw_kv_swap', 'a', '5', '40000')", + OUT_OF_RANGE, + ) + .await; + srv.expect_error("SELECT KV_GETSET('dw_kv_swap', 'a', '40000')", OUT_OF_RANGE) + .await; + assert_eq!( + cell(&srv, "SELECT value FROM dw_kv_swap WHERE key = 'a'").await, + "5" + ); +} + +/// A `TRANSFER` whose credit takes a declared `SMALLINT` balance past its +/// width is refused, and neither balance moves. +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn kv_transfer_past_a_declared_smallint_is_refused() { + let srv = TestServer::start().await; + srv.exec( + "CREATE COLLECTION dw_kv_acct (key TEXT PRIMARY KEY, balance SMALLINT, note TEXT) \ + WITH (engine='kv')", + ) + .await + .unwrap(); + srv.exec("INSERT INTO dw_kv_acct (key, balance, note) VALUES ('a', 100, 'x')") + .await + .unwrap(); + srv.exec("INSERT INTO dw_kv_acct (key, balance, note) VALUES ('b', 32700, 'y')") + .await + .unwrap(); + + srv.expect_error( + "SELECT TRANSFER('dw_kv_acct', 'a', 'b', 'balance', 100)", + OUT_OF_RANGE, + ) + .await; + assert_eq!( + cell(&srv, "SELECT balance FROM dw_kv_acct WHERE key = 'a'").await, + "100" + ); + assert_eq!( + cell(&srv, "SELECT balance FROM dw_kv_acct WHERE key = 'b'").await, + "32700" + ); +} diff --git a/nodedb/tests/wire/cases/sql_transactions_kv_atomic_overlay.rs b/nodedb/tests/wire/cases/sql_transactions_kv_atomic_overlay.rs index a3aeb388f..a7e0ea19f 100644 --- a/nodedb/tests/wire/cases/sql_transactions_kv_atomic_overlay.rs +++ b/nodedb/tests/wire/cases/sql_transactions_kv_atomic_overlay.rs @@ -184,17 +184,22 @@ async fn incr_on_absent_key_in_tx_creates_it_from_zero() { #[tokio::test(flavor = "multi_thread", worker_threads = 4)] async fn incr_float_in_tx_ryow() { let server = TestServer::start().await; - setup(&server).await; + // A float increment needs a FLOAT value column: an INT column refuses + // the non-whole result. + server + .exec("CREATE COLLECTION cf (key TEXT PRIMARY KEY, n FLOAT) WITH (engine='kv')") + .await + .unwrap(); server.exec("BEGIN").await.unwrap(); let rows = server - .query_text("SELECT KV_INCR_FLOAT('c', 'dmg', 2.5)") + .query_text("SELECT KV_INCR_FLOAT('cf', 'dmg', 2.5)") .await .unwrap(); assert!((json_of(&rows)["value"].as_f64().unwrap() - 2.5).abs() < f64::EPSILON); let rows2 = server - .query_text("SELECT KV_INCR_FLOAT('c', 'dmg', 1.5)") + .query_text("SELECT KV_INCR_FLOAT('cf', 'dmg', 1.5)") .await .unwrap(); assert!( From 2d6751a71068f5f83a8e548555ccb0906e3fe05b Mon Sep 17 00:00:00 2001 From: Farhan Syah Date: Thu, 8 Oct 2026 16:51:29 +0800 Subject: [PATCH 07/24] feat(kv): move DECIMAL balances exactly in TRANSFER A TRANSFER on a field declared DECIMAL takes an exact INT or DECIMAL amount and moves it by decimal arithmetic. Every other field keeps float arithmetic. The amount is typed through the plan, WAL and replication records, and the computed balances meet the collection's declared column rules. --- nodedb-physical/src/physical_plan/kv/mod.rs | 2 + nodedb-physical/src/physical_plan/kv/op.rs | 6 +- .../src/physical_plan/kv/transfer_amount.rs | 141 ++++++ .../src/control/planner/calvin/write_class.rs | 2 +- .../sql_plan_convert/kv_transfer_field.rs | 43 ++ .../control/planner/sql_plan_convert/mod.rs | 1 + .../server/shared/ddl/neutral/transfer.rs | 124 ++++- .../control/server/wal_dispatch_kv/encode.rs | 22 +- .../wal_replication/decode/kv_resolved.rs | 2 +- .../src/control/wal_replication/encode/kv.rs | 2 +- .../src/data/executor/handlers/kv/transfer.rs | 56 ++- .../executor/handlers/kv/transfer_compute.rs | 469 ++++++++++++++++-- .../data/executor/wal_replay_kv_transfer.rs | 64 ++- .../tests/wire/cases/kv_transfer_decimal.rs | 85 ++++ 14 files changed, 926 insertions(+), 93 deletions(-) create mode 100644 nodedb-physical/src/physical_plan/kv/transfer_amount.rs create mode 100644 nodedb/src/control/planner/sql_plan_convert/kv_transfer_field.rs create mode 100644 nodedb/tests/wire/cases/kv_transfer_decimal.rs 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/src/control/planner/calvin/write_class.rs b/nodedb/src/control/planner/calvin/write_class.rs index 8e8072c7b..487a37fb6 100644 --- a/nodedb/src/control/planner/calvin/write_class.rs +++ b/nodedb/src/control/planner/calvin/write_class.rs @@ -544,7 +544,7 @@ mod tests { source_key: b"a".to_vec(), dest_key: b"b".to_vec(), field: "balance".to_owned(), - amount: 10.0, + amount: nodedb_physical::physical_plan::TransferAmount::Float(10.0), debit_surrogate: Surrogate::new(1), credit_surrogate: Surrogate::new(2), rls_write_check: nodedb_types::RlsWriteCheck::pending_injection(), diff --git a/nodedb/src/control/planner/sql_plan_convert/kv_transfer_field.rs b/nodedb/src/control/planner/sql_plan_convert/kv_transfer_field.rs new file mode 100644 index 000000000..4348bd8db --- /dev/null +++ b/nodedb/src/control/planner/sql_plan_convert/kv_transfer_field.rs @@ -0,0 +1,43 @@ +// SPDX-License-Identifier: BUSL-1.1 + +//! The declared type of the field a KV `TRANSFER` moves. + +use std::sync::Arc; + +use nodedb_sql::SqlCatalog; +use nodedb_sql::types_expr::SqlDataType; + +use crate::control::planner::catalog_adapter::OriginCatalog; +use crate::control::planner::plan_error_map::map_plan_error; +use crate::control::state::SharedState; +use crate::types::{DatabaseId, TenantId}; + +/// Whether the catalog declares `field` of `collection` as `DECIMAL`, with +/// or without a typmod. +/// +/// A raw collection, an undeclared field, and a collection the catalog does +/// not hold are not `DECIMAL`. +pub(crate) fn kv_transfer_field_is_decimal( + state: &SharedState, + tenant_id: TenantId, + database_id: DatabaseId, + collection: &str, + field: &str, +) -> crate::Result { + let catalog = OriginCatalog::new( + Arc::clone(&state.credentials), + state.array_catalog.clone(), + tenant_id.as_u64(), + database_id, + Some(Arc::clone(&state.retention_policy_registry)), + ) + .with_sequence_registry(Arc::clone(&state.sequence_registry)); + let info = catalog + .get_collection(database_id, collection) + .map_err(|e| map_plan_error(e.into(), tenant_id))?; + Ok(info.is_some_and(|info| { + info.columns.iter().any(|column| { + column.name == field && matches!(column.data_type, SqlDataType::Decimal(_)) + }) + })) +} diff --git a/nodedb/src/control/planner/sql_plan_convert/mod.rs b/nodedb/src/control/planner/sql_plan_convert/mod.rs index 98041e6d2..51cdb74f5 100644 --- a/nodedb/src/control/planner/sql_plan_convert/mod.rs +++ b/nodedb/src/control/planner/sql_plan_convert/mod.rs @@ -13,6 +13,7 @@ pub mod filter; pub mod filter_scan_side; pub mod group_key_name; pub mod kv_counter_shape; +pub mod kv_transfer_field; pub mod lateral; pub mod output_schema; pub mod output_schema_types; diff --git a/nodedb/src/control/server/shared/ddl/neutral/transfer.rs b/nodedb/src/control/server/shared/ddl/neutral/transfer.rs index 4ad9ea56b..f09f1555f 100644 --- a/nodedb/src/control/server/shared/ddl/neutral/transfer.rs +++ b/nodedb/src/control/server/shared/ddl/neutral/transfer.rs @@ -6,6 +6,8 @@ //! `SELECT TRANSFER(collection, source_key, dest_key, field, amount)` //! — Atomically: source.field -= amount, dest.field += amount. //! — Fails with INSUFFICIENT_BALANCE if source.field < amount. +//! — A `DECIMAL` field takes an exact INT or DECIMAL amount and moves by +//! exact decimal arithmetic. Every other field moves by float arithmetic. //! — Returns: `{ source_key, dest_key, field, amount, source_balance, dest_balance }`. //! //! `SELECT TRANSFER_ITEM(source_collection, dest_collection, item_id, source_owner, dest_owner)` @@ -20,7 +22,8 @@ use crate::control::security::identity::AuthenticatedIdentity; use crate::control::server::shared::session::DmlTxnCtx; use crate::control::state::SharedState; use crate::types::DatabaseId; -use nodedb_physical::physical_plan::{KvOp, PhysicalPlan}; +use nodedb_physical::physical_plan::{KvOp, PhysicalPlan, TransferAmount}; +use rust_decimal::Decimal; use super::super::result::{DdlError, DdlResult}; use super::kv_atomic::{dispatch_and_respond, parse_function_args, unquote}; @@ -46,15 +49,9 @@ pub async fn transfer( let source_key = unquote(&args[1]); let dest_key = unquote(&args[2]); let field = unquote(&args[3]); - let amount_str = args[4].trim().to_string(); - let amount: f64 = amount_str.parse().map_err(|_| { - ddl_err( - "42601", - format!("TRANSFER: amount must be a number, got '{amount_str}'"), - ) - })?; - - if amount <= 0.0 { + let decimal_field = transfer_field_is_decimal(state, identity, &collection, &field)?; + let amount = parse_amount(args[4].trim(), &field, decimal_field)?; + if !amount.is_positive() { return Err(ddl_err("42601", "TRANSFER: amount must be positive")); } @@ -204,6 +201,113 @@ pub async fn transfer_item( // ── Helpers ──────────────────────────────────────────────────────────── +/// Whether the catalog declares `field` of `collection` as `DECIMAL`. +/// +/// The caller's grants are checked first, the pair `dispatch_and_respond` +/// checks: a caller refused the collection learns nothing of its columns. +fn transfer_field_is_decimal( + state: &SharedState, + identity: &AuthenticatedIdentity, + collection: &str, + field: &str, +) -> Result { + let gate = + super::read_gate::CollectionReadGate::for_request(state, identity, DatabaseId::DEFAULT); + gate.authorize(collection)?; + gate.authorize_permission( + collection, + crate::control::security::identity::Permission::Write, + )?; + crate::control::planner::sql_plan_convert::kv_transfer_field::kv_transfer_field_is_decimal( + state, + identity.tenant_id, + DatabaseId::DEFAULT, + collection, + field, + ) + .map_err(|error| DdlError::from_error(&error)) +} + +/// The `TRANSFER` amount literal, typed by the field it moves. +/// +/// A `DECIMAL` field takes an exact INT or DECIMAL literal. A number no +/// `Decimal` holds is out of range for it (`22003`). Every other field takes +/// a finite float. Text that is no number is refused with `42601`. +fn parse_amount(text: &str, field: &str, decimal_field: bool) -> Result { + let not_a_number = || { + ddl_err( + "42601", + format!("TRANSFER: amount must be a number, got '{text}'"), + ) + }; + let float = text + .parse::() + .ok() + .filter(|f| f.is_finite()) + .ok_or_else(not_a_number)?; + if !decimal_field { + return Ok(TransferAmount::Float(float)); + } + let exact = Decimal::from_str_exact(text) + .or_else(|_| Decimal::from_scientific(text)) + .map_err(|_| { + ddl_err( + "22003", + format!("TRANSFER: amount {text} is out of range for DECIMAL field '{field}'"), + ) + })?; + Ok(TransferAmount::Decimal(exact)) +} + fn ddl_err(sqlstate: &str, message: impl Into) -> DdlError { DdlError::new(sqlstate, message) } + +#[cfg(test)] +mod tests { + use super::*; + + fn dec(text: &str) -> Decimal { + text.parse().expect("test decimal parses") + } + + #[test] + fn a_decimal_field_takes_an_exact_amount() { + for (text, expected) in [ + ("30", "30"), + ("0.1", "0.1"), + ("12.345", "12.345"), + ("1.5e2", "150"), + ] { + assert_eq!( + parse_amount(text, "balance", true).expect(text), + TransferAmount::Decimal(dec(expected)), + "{text}" + ); + } + } + + #[test] + fn a_decimal_field_refuses_an_amount_no_decimal_holds() { + let err = parse_amount("1e40", "balance", true).expect_err("past the Decimal range"); + assert_eq!(err.sqlstate, "22003"); + } + + #[test] + fn other_fields_take_a_float_amount() { + assert_eq!( + parse_amount("2.5", "balance", false).expect("float"), + TransferAmount::Float(2.5) + ); + } + + #[test] + fn text_that_is_no_finite_number_is_refused() { + for text in ["abc", "NaN", "inf", ""] { + for decimal_field in [true, false] { + let err = parse_amount(text, "balance", decimal_field).expect_err(text); + assert_eq!(err.sqlstate, "42601", "{text}"); + } + } + } +} diff --git a/nodedb/src/control/server/wal_dispatch_kv/encode.rs b/nodedb/src/control/server/wal_dispatch_kv/encode.rs index 907d5ccf8..be1d77d30 100644 --- a/nodedb/src/control/server/wal_dispatch_kv/encode.rs +++ b/nodedb/src/control/server/wal_dispatch_kv/encode.rs @@ -3,6 +3,7 @@ //! Pure payload encoders for KV WAL records. use nodedb_physical::physical_plan::KvCounterShape; +use nodedb_physical::physical_plan::TransferAmount; use nodedb_physical::physical_plan::UpdateValue; /// Serialize `value` to a MessagePack WAL payload, wrapping any encode error @@ -75,7 +76,7 @@ pub(crate) struct KvTransferFields<'a> { pub source_key: &'a [u8], pub dest_key: &'a [u8], pub field: &'a str, - pub amount: f64, + pub amount: TransferAmount, pub debit_surrogate: u32, pub credit_surrogate: u32, } @@ -403,7 +404,7 @@ pub(crate) fn encode_kv_truncate(collection: &str) -> crate::Result> { #[cfg(test)] mod tests { - use nodedb_physical::physical_plan::{KvCounterShape, UpdateValue}; + use nodedb_physical::physical_plan::{KvCounterShape, TransferAmount, UpdateValue}; use super::{ KvIncrRecord, KvTransferFields, encode_kv_batch_put, encode_kv_cas, encode_kv_expire, @@ -532,7 +533,7 @@ mod tests { source_key: b"alice", dest_key: b"bob", field: "balance", - amount: 30.0, + amount: TransferAmount::Float(30.0), debit_surrogate: 7, credit_surrogate: 8, }) @@ -547,16 +548,23 @@ mod tests { amount, debit_surrogate, credit_surrogate, - ) = zerompk::from_msgpack::<(&str, String, Vec, Vec, String, f64, u32, u32)>( - &entry, - ) + ) = zerompk::from_msgpack::<( + &str, + String, + Vec, + Vec, + String, + TransferAmount, + u32, + u32, + )>(&entry) .unwrap(); assert_eq!(disc, "kv_transfer"); assert_eq!(collection, "accounts"); assert_eq!(source_key, b"alice"); assert_eq!(dest_key, b"bob"); assert_eq!(field, "balance"); - assert_eq!(amount, 30.0); + assert_eq!(amount, TransferAmount::Float(30.0)); assert_eq!(debit_surrogate, 7); assert_eq!(credit_surrogate, 8); } diff --git a/nodedb/src/control/wal_replication/decode/kv_resolved.rs b/nodedb/src/control/wal_replication/decode/kv_resolved.rs index 5c6e36779..9572a0d76 100644 --- a/nodedb/src/control/wal_replication/decode/kv_resolved.rs +++ b/nodedb/src/control/wal_replication/decode/kv_resolved.rs @@ -18,7 +18,7 @@ pub(super) struct TransferFields<'a> { pub(super) source_key: &'a [u8], pub(super) dest_key: &'a [u8], pub(super) field: &'a str, - pub(super) amount: f64, + pub(super) amount: nodedb_physical::physical_plan::TransferAmount, pub(super) debit_surrogate: u32, pub(super) credit_surrogate: u32, } diff --git a/nodedb/src/control/wal_replication/encode/kv.rs b/nodedb/src/control/wal_replication/encode/kv.rs index 5eb44bba5..ab5991c96 100644 --- a/nodedb/src/control/wal_replication/encode/kv.rs +++ b/nodedb/src/control/wal_replication/encode/kv.rs @@ -349,7 +349,7 @@ pub(super) fn transfer( source_key: &[u8], dest_key: &[u8], field: &str, - amount: f64, + amount: nodedb_physical::physical_plan::TransferAmount, debit_surrogate: u32, credit_surrogate: u32, ) -> ReplicatedWrite { diff --git a/nodedb/src/data/executor/handlers/kv/transfer.rs b/nodedb/src/data/executor/handlers/kv/transfer.rs index c2f02c477..4cad0c61a 100644 --- a/nodedb/src/data/executor/handlers/kv/transfer.rs +++ b/nodedb/src/data/executor/handlers/kv/transfer.rs @@ -8,6 +8,7 @@ use tracing::debug; +use super::declared_body::fit_kv_image; use super::transfer_compute::{TransferError, compute_transfer}; use crate::bridge::envelope::{ErrorCode, Response}; use crate::data::executor::core_loop::CoreLoop; @@ -23,7 +24,8 @@ pub(in crate::data::executor) struct TransferParams<'a> { pub source_key: &'a [u8], pub dest_key: &'a [u8], pub field: &'a str, - pub amount: f64, + /// The amount, typed by the field it moves. + pub amount: nodedb_physical::physical_plan::TransferAmount, /// Cross-engine surrogate of the debit (source) row. pub debit_surrogate: nodedb_types::Surrogate, /// Cross-engine surrogate of the credit (dest) row. @@ -74,7 +76,7 @@ impl CoreLoop { credit_surrogate, rls_write_check, } = params; - debug!(core = self.core_id, %collection, %field, amount, "kv transfer"); + debug!(core = self.core_id, %collection, %field, %amount, "kv transfer"); if self.kv_engine.is_over_budget() { return self.response_error(task, ErrorCode::ResourcesExhausted); @@ -100,8 +102,10 @@ impl CoreLoop { } else { Some(dest_bytes.as_slice()) }; - let computed = match compute_transfer(&source_bytes, dest_ref, field, amount) { + let declared = self.declared_columns_of(did, tid, collection); + let computed = match compute_transfer(&source_bytes, dest_ref, field, amount, declared) { Ok(c) => c, + Err(TransferError::Declared(e)) => return self.response_error(task, e), Err(TransferError::TypeMismatch(detail)) => { return self.response_error( task, @@ -123,6 +127,7 @@ impl CoreLoop { }; let new_source = computed.new_source; let new_dest = computed.new_dest; + let moved = computed.amount; let source_balance_after = computed.source_balance_after; let dest_balance_after = computed.dest_balance_after; @@ -205,17 +210,12 @@ impl CoreLoop { "source_key": src_str, "dest_key": dst_str, "field": field, - "amount": amount, - "source_balance": source_balance_after, - "dest_balance": dest_balance_after, + "amount": moved.to_json(), + "source_balance": source_balance_after.to_json(), + "dest_balance": dest_balance_after.to_json(), })) { Ok(payload) => self.response_with_payload(task, payload), - Err(e) => self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ), + Err(e) => self.response_error(task, ErrorCode::from(e)), } } @@ -252,10 +252,21 @@ impl CoreLoop { return self.response_error(task, ErrorCode::NotFound); }; - // The same bytes are two different images to two different policies: - // the row leaving the source, and the row arriving at the destination. - // Both are decided before either half runs, so a move a policy rejects - // cannot delete from the source and then fail to insert at the dest. + // The row arriving at the destination meets the destination's + // declared numeric columns, decided before either half runs. + let fitted = match fit_kv_image( + &item_data, + self.declared_columns_of(did, tid, dest_collection), + ) { + Ok(fitted) => fitted, + Err(e) => return self.response_error(task, e), + }; + let dest_data: &[u8] = fitted.as_deref().unwrap_or(&item_data); + + // The row leaving the source and the row arriving at the destination + // are two images to two different policies. Both are decided before + // either half runs, so a move a policy rejects cannot delete from the + // source and then fail to insert at the dest. if let Err(e) = super::rls::admit_kv_row( source_rls_write_check, &item_data, @@ -267,7 +278,7 @@ impl CoreLoop { } if let Err(e) = super::rls::admit_kv_row( dest_rls_write_check, - &item_data, + dest_data, dest_key, tid, dest_collection, @@ -289,7 +300,7 @@ impl CoreLoop { tenant_id: tid, collection: dest_collection, key: dest_key, - value: &item_data, + value: dest_data, ttl_ms: 0, now_ms, surrogate, @@ -318,7 +329,7 @@ impl CoreLoop { dest_collection, crate::event::WriteOp::Insert, dest_key, - Some(&item_data), + Some(dest_data), None, ); @@ -329,12 +340,7 @@ impl CoreLoop { "dest_collection": dest_collection, })) { Ok(payload) => self.response_with_payload(task, payload), - Err(e) => self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ), + Err(e) => self.response_error(task, ErrorCode::from(e)), } } } diff --git a/nodedb/src/data/executor/handlers/kv/transfer_compute.rs b/nodedb/src/data/executor/handlers/kv/transfer_compute.rs index 11f1b3302..111df60ce 100644 --- a/nodedb/src/data/executor/handlers/kv/transfer_compute.rs +++ b/nodedb/src/data/executor/handlers/kv/transfer_compute.rs @@ -6,11 +6,20 @@ //! value and its COMMIT-time durable replay are always computed by the //! exact same code — mirrors the `nodedb_physical::kv_atomic::compute` / `stage_kv_atomic` //! split for `Incr`/`Cas`/etc. +//! +//! A `DECIMAL` field moves by exact decimal arithmetic. Every other field +//! moves by float arithmetic. use std::collections::HashMap; +use std::fmt; +use nodedb_physical::physical_plan::{DeclaredColumn, TransferAmount}; use nodedb_query::msgpack_scan::{KvBodyShape, kv_body_to_row, row_to_kv_body}; use nodedb_types::Value; +use nodedb_types::columnar::ColumnType; +use rust_decimal::Decimal; + +use crate::data::executor::strict_format::coerce_declared_row; /// Failure modes of [`compute_transfer`], translated to `ErrorCode` at each /// call site (the live handler and the staging handler render slightly @@ -18,7 +27,41 @@ use nodedb_types::Value; #[derive(Debug)] pub(in crate::data::executor) enum TransferError { TypeMismatch(String), - InsufficientBalance { have: f64, need: f64 }, + InsufficientBalance { + have: TransferNumber, + need: TransferNumber, + }, + /// A computed balance breaks a declared column rule of the collection, + /// such as a `SMALLINT` range or a `DECIMAL(p,s)` precision. An exact + /// sum past the `Decimal` range is refused the same way. + Declared(crate::Error), +} + +/// A balance or amount of a transfer, in the arithmetic its field uses. +#[derive(Debug, Clone, Copy, PartialEq)] +pub(in crate::data::executor) enum TransferNumber { + Float(f64), + Decimal(Decimal), +} + +impl TransferNumber { + /// The response form: a float is a JSON number, a decimal its exact + /// text. + pub(in crate::data::executor) fn to_json(self) -> serde_json::Value { + match self { + Self::Float(f) => serde_json::json!(f), + Self::Decimal(d) => serde_json::Value::String(d.to_string()), + } + } +} + +impl fmt::Display for TransferNumber { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + match self { + Self::Float(v) => write!(f, "{v}"), + Self::Decimal(d) => write!(f, "{d}"), + } + } } /// The two updated document bodies and post-transfer balances for a @@ -27,8 +70,12 @@ pub(in crate::data::executor) enum TransferError { pub(in crate::data::executor) struct TransferComputation { pub new_source: Vec, pub new_dest: Vec, - pub source_balance_after: f64, - pub dest_balance_after: f64, + /// The amount moved, in the field's arithmetic. + pub amount: TransferNumber, + /// The stored balances. A decimal balance is the value its declared + /// column fitted. + pub source_balance_after: TransferNumber, + pub dest_balance_after: TransferNumber, } /// Compute the read-validate-write outcome of an atomic fungible transfer. @@ -38,43 +85,167 @@ pub(in crate::data::executor) struct TransferComputation { /// destination that exists but lacks `field` starts from 0. Either side /// holding a bare value (the single-`value` SQL form, RESP `SET`) or a /// non-numeric `field` is a type mismatch, never silently treated as 0. +/// +/// A `DECIMAL` field moves by exact decimal arithmetic: the planner sent a +/// decimal amount, or the collection declares the field `DECIMAL`. The +/// balance check is exact, and the results are stored as decimal text. +/// +/// Both rows then meet the collection's `declared` numeric columns, so a +/// balance past a declared width refuses the whole transfer. A `DECIMAL(p,s)` +/// balance rounds to `s` digits, half away from zero. pub(in crate::data::executor) fn compute_transfer( source_bytes: &[u8], dest_bytes: Option<&[u8]>, field: &str, - amount: f64, + amount: TransferAmount, + declared: &[DeclaredColumn], ) -> Result { let mut source = map_row(source_bytes, "source")?; - let source_balance = numeric_field(&source, field)?.ok_or_else(|| { - TransferError::TypeMismatch(format!("field '{field}' is not numeric or missing")) - })?; + let Moved { + mut dest, + amount, + source_after, + dest_after, + } = if moves_decimal(declared, field, amount) { + move_decimal(&mut source, dest_bytes, field, amount)? + } else { + move_float(&mut source, dest_bytes, field, amount.to_f64())? + }; + coerce_declared_row(&mut source, declared).map_err(TransferError::Declared)?; + coerce_declared_row(&mut dest, declared).map_err(TransferError::Declared)?; + Ok(TransferComputation { + source_balance_after: fitted(&source, field, source_after), + dest_balance_after: fitted(&dest, field, dest_after), + new_source: encode_map(source, "source")?, + new_dest: encode_map(dest, "destination")?, + amount, + }) +} + +/// The destination row and the balances one arithmetic computed. +struct Moved { + dest: HashMap, + amount: TransferNumber, + source_after: TransferNumber, + dest_after: TransferNumber, +} + +/// Whether `field` moves by exact decimal arithmetic. +fn moves_decimal(declared: &[DeclaredColumn], field: &str, amount: TransferAmount) -> bool { + matches!(amount, TransferAmount::Decimal(_)) + || declared + .iter() + .any(|c| c.name == field && matches!(c.column_type, ColumnType::Decimal(_))) +} + +/// Move `amount` by float arithmetic. A whole result stays an integer. +fn move_float( + source: &mut HashMap, + dest_bytes: Option<&[u8]>, + field: &str, + amount: f64, +) -> Result { + let source_balance = float_field(source, field)?.ok_or_else(|| missing(field))?; if source_balance < amount { return Err(TransferError::InsufficientBalance { - have: source_balance, - need: amount, + have: TransferNumber::Float(source_balance), + need: TransferNumber::Float(amount), }); } + let mut dest = dest_row(dest_bytes)?; + let dest_balance = float_field(&dest, field)?.unwrap_or(0.0); - let mut dest = match dest_bytes.filter(|b| !b.is_empty()) { - None => HashMap::with_capacity(1), - Some(bytes) => map_row(bytes, "destination")?, - }; - let dest_balance = numeric_field(&dest, field)?.unwrap_or(0.0); + let source_after = source_balance - amount; + let dest_after = dest_balance + amount; + source.insert(field.to_string(), numeric_value(source_after)); + dest.insert(field.to_string(), numeric_value(dest_after)); + Ok(Moved { + dest, + amount: TransferNumber::Float(amount), + source_after: TransferNumber::Float(source_after), + dest_after: TransferNumber::Float(dest_after), + }) +} - let source_balance_after = source_balance - amount; - let dest_balance_after = dest_balance + amount; - source.insert(field.to_string(), numeric_value(source_balance_after)); - dest.insert(field.to_string(), numeric_value(dest_balance_after)); +/// Move `amount` by exact decimal arithmetic. Both results are stored as +/// decimal text, the form a decimal column stores. +fn move_decimal( + source: &mut HashMap, + dest_bytes: Option<&[u8]>, + field: &str, + amount: TransferAmount, +) -> Result { + let amount = amount.to_decimal().ok_or_else(|| { + out_of_range(format!( + "TRANSFER amount {amount} is out of range for DECIMAL field '{field}'" + )) + })?; + let source_balance = decimal_field(source, field)?.ok_or_else(|| missing(field))?; + if source_balance < amount { + return Err(TransferError::InsufficientBalance { + have: TransferNumber::Decimal(source_balance), + need: TransferNumber::Decimal(amount), + }); + } + let mut dest = dest_row(dest_bytes)?; + let dest_balance = decimal_field(&dest, field)?.unwrap_or(Decimal::ZERO); - Ok(TransferComputation { - new_source: encode_map(source, "source")?, - new_dest: encode_map(dest, "destination")?, - source_balance_after, - dest_balance_after, + let source_after = source_balance.checked_sub(amount).ok_or_else(|| { + out_of_range(format!( + "field '{field}': {source_balance} - {amount} is out of the DECIMAL range" + )) + })?; + let dest_after = dest_balance.checked_add(amount).ok_or_else(|| { + out_of_range(format!( + "field '{field}': {dest_balance} + {amount} is out of the DECIMAL range" + )) + })?; + source.insert(field.to_string(), Value::String(source_after.to_string())); + dest.insert(field.to_string(), Value::String(dest_after.to_string())); + Ok(Moved { + dest, + amount: TransferNumber::Decimal(amount), + source_after: TransferNumber::Decimal(source_after), + dest_after: TransferNumber::Decimal(dest_after), }) } +/// The stored balance of `field` once the declared rule fitted it. A float +/// balance is the computed value. A decimal balance is read back, so the +/// answer shows the rounding its typmod applied. +fn fitted(row: &HashMap, field: &str, computed: TransferNumber) -> TransferNumber { + let TransferNumber::Decimal(_) = computed else { + return computed; + }; + match row.get(field) { + Some(Value::String(text)) => text + .trim() + .parse() + .map_or(computed, TransferNumber::Decimal), + Some(Value::Decimal(d)) => TransferNumber::Decimal(*d), + _ => computed, + } +} + +/// The SQLSTATE `22003` refusal of an exact balance. +fn out_of_range(detail: String) -> TransferError { + TransferError::Declared(crate::Error::NumericValueOutOfRange { detail }) +} + +/// The refusal of a source row that holds no number in `field`. +fn missing(field: &str) -> TransferError { + TransferError::TypeMismatch(format!("field '{field}' is not numeric or missing")) +} + +/// The destination row: empty when the key is absent. +fn dest_row(dest_bytes: Option<&[u8]>) -> Result, TransferError> { + match dest_bytes.filter(|b| !b.is_empty()) { + None => Ok(HashMap::with_capacity(1)), + Some(bytes) => map_row(bytes, "destination"), + } +} + /// Decode a KV body as the typed-column map a transfer operates on. fn map_row(bytes: &[u8], side: &str) -> Result, TransferError> { let (row, shape) = kv_body_to_row(bytes) @@ -95,7 +266,7 @@ fn map_row(bytes: &[u8], side: &str) -> Result, TransferE /// `field` as f64: `Ok(None)` when absent, a type mismatch when present but /// not numeric. -fn numeric_field(map: &HashMap, field: &str) -> Result, TransferError> { +fn float_field(map: &HashMap, field: &str) -> Result, TransferError> { match map.get(field) { None => Ok(None), Some(Value::Float(f)) => Ok(Some(*f)), @@ -107,6 +278,31 @@ fn numeric_field(map: &HashMap, field: &str) -> Result, + field: &str, +) -> Result, TransferError> { + let parsed = match map.get(field) { + None => return Ok(None), + Some(Value::Decimal(d)) => Some(*d), + Some(Value::Integer(i)) => Some(Decimal::from(*i)), + Some(Value::Float(f)) => f.to_string().parse().ok(), + Some(Value::String(text)) => text.trim().parse().ok(), + Some(other) => { + return Err(TransferError::TypeMismatch(format!( + "field '{field}' is {}, not numeric", + other.type_name() + ))); + } + }; + parsed.map(Some).ok_or_else(|| { + TransferError::TypeMismatch(format!("field '{field}' does not hold a DECIMAL value")) + }) +} + /// A whole-number balance stays an integer on disk; anything else is a float. fn numeric_value(v: f64) -> Value { if v.fract() == 0.0 && v >= i64::MIN as f64 && v <= i64::MAX as f64 { @@ -140,20 +336,71 @@ mod tests { nodedb_types::json_to_msgpack(&serde_json::json!({ field: value })).unwrap() } + fn text_doc(field: &str, value: &str) -> Vec { + nodedb_types::json_to_msgpack(&serde_json::json!({ field: value })).unwrap() + } + + fn float(v: f64) -> TransferAmount { + TransferAmount::Float(v) + } + + fn dec(text: &str) -> Decimal { + text.parse().expect("test decimal parses") + } + + fn exact(text: &str) -> TransferAmount { + TransferAmount::Decimal(dec(text)) + } + + fn decimal_10_2() -> [DeclaredColumn; 1] { + [DeclaredColumn::from_declared("balance", "DECIMAL(10,2)").expect("decimal")] + } + + /// The stored cell of `field`. + fn stored(body: &[u8], field: &str) -> Value { + let (row, _) = kv_body_to_row(body).expect("decode body"); + row.get(field).cloned().expect("field present") + } + #[test] fn transfer_moves_balance_between_existing_docs() { let source = doc("balance", 100.0); let dest = doc("balance", 10.0); - let result = compute_transfer(&source, Some(&dest), "balance", 30.0).unwrap(); - assert_eq!(result.source_balance_after, 70.0); - assert_eq!(result.dest_balance_after, 40.0); + let result = compute_transfer(&source, Some(&dest), "balance", float(30.0), &[]).unwrap(); + assert_eq!(result.source_balance_after, TransferNumber::Float(70.0)); + assert_eq!(result.dest_balance_after, TransferNumber::Float(40.0)); + } + + /// A credit that takes a declared `SMALLINT` balance past its width + /// refuses the whole transfer. + #[test] + fn transfer_past_a_declared_width_is_refused() { + let declared = [DeclaredColumn::from_declared("balance", "SMALLINT").expect("smallint")]; + let source = doc("balance", 100.0); + let dest = doc("balance", 32700.0); + let err = compute_transfer(&source, Some(&dest), "balance", float(100.0), &declared) + .expect_err("credit past smallint"); + assert!( + matches!( + err, + TransferError::Declared(crate::Error::NumericValueOutOfRange { .. }) + ), + "{err:?}" + ); + + let fits = compute_transfer(&source, Some(&dest), "balance", float(50.0), &declared) + .expect("credit fits smallint"); + assert_eq!( + extract_numeric_field(&fits.new_dest, "balance"), + Some(32750.0) + ); } #[test] fn transfer_creates_dest_when_absent() { let source = doc("balance", 100.0); - let result = compute_transfer(&source, None, "balance", 30.0).unwrap(); - assert_eq!(result.dest_balance_after, 30.0); + let result = compute_transfer(&source, None, "balance", float(30.0), &[]).unwrap(); + assert_eq!(result.dest_balance_after, TransferNumber::Float(30.0)); assert_eq!( extract_numeric_field(&result.new_dest, "balance"), Some(30.0) @@ -163,7 +410,7 @@ mod tests { #[test] fn transfer_rejects_insufficient_balance() { let source = doc("balance", 10.0); - let err = compute_transfer(&source, None, "balance", 30.0); + let err = compute_transfer(&source, None, "balance", float(30.0), &[]); assert!(matches!( err, Err(TransferError::InsufficientBalance { .. }) @@ -175,7 +422,7 @@ mod tests { // A raw body (the single-`value` form, RESP `SET`) is not a hash: // never treated as a zero balance and re-encoded as a map. let source = doc("balance", 100.0); - let err = compute_transfer(&source, Some(b"5"), "balance", 30.0); + let err = compute_transfer(&source, Some(b"5"), "balance", float(30.0), &[]); assert!( matches!(err, Err(TransferError::TypeMismatch(_))), "{err:?}" @@ -184,7 +431,7 @@ mod tests { #[test] fn transfer_rejects_a_bare_value_source() { - let err = compute_transfer(b"100", None, "balance", 30.0); + let err = compute_transfer(b"100", None, "balance", float(30.0), &[]); assert!( matches!(err, Err(TransferError::TypeMismatch(_))), "{err:?}" @@ -194,8 +441,8 @@ mod tests { #[test] fn transfer_rejects_non_numeric_destination_field() { let source = doc("balance", 100.0); - let dest = nodedb_types::json_to_msgpack(&serde_json::json!({"balance": "abc"})).unwrap(); - let err = compute_transfer(&source, Some(&dest), "balance", 30.0); + let dest = text_doc("balance", "abc"); + let err = compute_transfer(&source, Some(&dest), "balance", float(30.0), &[]); assert!( matches!(err, Err(TransferError::TypeMismatch(_))), "{err:?}" @@ -204,8 +451,158 @@ mod tests { #[test] fn transfer_rejects_non_numeric_field() { - let source = nodedb_types::json_to_msgpack(&serde_json::json!({"balance": "abc"})).unwrap(); - let err = compute_transfer(&source, None, "balance", 30.0); + let source = text_doc("balance", "abc"); + let err = compute_transfer(&source, None, "balance", float(30.0), &[]); assert!(matches!(err, Err(TransferError::TypeMismatch(_)))); } + + /// A `DECIMAL(10,2)` balance moves by the exact amount and is stored as + /// decimal text. + #[test] + fn decimal_transfer_is_exact() { + let source = text_doc("balance", "100.00"); + let dest = text_doc("balance", "5.25"); + let result = compute_transfer( + &source, + Some(&dest), + "balance", + exact("30.10"), + &decimal_10_2(), + ) + .expect("decimal transfer"); + assert_eq!( + stored(&result.new_source, "balance"), + Value::String("69.90".into()) + ); + assert_eq!( + stored(&result.new_dest, "balance"), + Value::String("35.35".into()) + ); + assert_eq!(result.amount, TransferNumber::Decimal(dec("30.10"))); + assert_eq!( + result.source_balance_after, + TransferNumber::Decimal(dec("69.90")) + ); + assert_eq!( + result.dest_balance_after, + TransferNumber::Decimal(dec("35.35")) + ); + } + + /// Each balance rounds to the declared scale, half away from zero. + #[test] + fn decimal_transfer_rounds_each_balance_to_its_scale() { + let source = text_doc("balance", "100.00"); + let dest = text_doc("balance", "5.25"); + let result = compute_transfer( + &source, + Some(&dest), + "balance", + exact("0.005"), + &decimal_10_2(), + ) + .expect("decimal transfer"); + assert_eq!( + stored(&result.new_source, "balance"), + Value::String("100.00".into()) + ); + assert_eq!( + stored(&result.new_dest, "balance"), + Value::String("5.26".into()) + ); + assert_eq!( + result.dest_balance_after, + TransferNumber::Decimal(dec("5.26")) + ); + } + + /// A credit past `DECIMAL(10,2)` is refused with `22003`. + #[test] + fn decimal_credit_past_the_precision_is_refused() { + let source = text_doc("balance", "100.00"); + let dest = text_doc("balance", "99999999.99"); + let err = compute_transfer(&source, Some(&dest), "balance", exact("1"), &decimal_10_2()) + .expect_err("credit past precision"); + assert!( + matches!( + err, + TransferError::Declared(crate::Error::NumericValueOutOfRange { .. }) + ), + "{err:?}" + ); + } + + /// The balance check compares exact decimals. + #[test] + fn decimal_balance_check_is_exact() { + let source = text_doc("balance", "0.10"); + let err = compute_transfer(&source, None, "balance", exact("0.11"), &decimal_10_2()) + .expect_err("0.10 does not cover 0.11"); + assert!( + matches!( + err, + TransferError::InsufficientBalance { + have: TransferNumber::Decimal(_), + need: TransferNumber::Decimal(_), + } + ), + "{err:?}" + ); + let all = compute_transfer(&source, None, "balance", exact("0.10"), &decimal_10_2()) + .expect("0.10 covers 0.10"); + assert_eq!( + stored(&all.new_source, "balance"), + Value::String("0.00".into()) + ); + assert_eq!( + stored(&all.new_dest, "balance"), + Value::String("0.10".into()) + ); + } + + /// An integer amount moves a decimal balance. A float amount converts + /// through its literal text. + #[test] + fn integer_and_float_amounts_move_a_declared_decimal() { + let source = text_doc("balance", "10.50"); + let int = compute_transfer(&source, None, "balance", exact("3"), &decimal_10_2()) + .expect("integer amount"); + assert_eq!( + stored(&int.new_source, "balance"), + Value::String("7.50".into()) + ); + let tenth = compute_transfer(&source, None, "balance", float(0.1), &decimal_10_2()) + .expect("float amount"); + assert_eq!( + stored(&tenth.new_source, "balance"), + Value::String("10.40".into()) + ); + } + + /// A plain `DECIMAL` field is not declared. The planner's decimal amount + /// still moves it exactly. + #[test] + fn decimal_amount_moves_a_plain_decimal_field() { + let source = text_doc("balance", "1.5"); + let result = + compute_transfer(&source, None, "balance", exact("0.25"), &[]).expect("plain decimal"); + assert_eq!( + stored(&result.new_source, "balance"), + Value::String("1.25".into()) + ); + assert_eq!( + stored(&result.new_dest, "balance"), + Value::String("0.25".into()) + ); + } + + #[test] + fn decimal_amount_refuses_non_numeric_text() { + let source = text_doc("balance", "abc"); + let err = compute_transfer(&source, None, "balance", exact("1"), &[]); + assert!( + matches!(err, Err(TransferError::TypeMismatch(_))), + "{err:?}" + ); + } } diff --git a/nodedb/src/data/executor/wal_replay_kv_transfer.rs b/nodedb/src/data/executor/wal_replay_kv_transfer.rs index 2f6b216fc..ec0ae0eac 100644 --- a/nodedb/src/data/executor/wal_replay_kv_transfer.rs +++ b/nodedb/src/data/executor/wal_replay_kv_transfer.rs @@ -9,9 +9,11 @@ //! `handlers/kv/transfer.rs` perform, against whatever state this core's KV //! engine holds at this point in LSN-ordered replay. +use nodedb_physical::physical_plan::TransferAmount; use tracing::warn; use super::core_loop::CoreLoop; +use super::handlers::kv::declared_body::fit_kv_image; use super::handlers::kv::transfer_compute::{TransferError, compute_transfer}; use crate::data::executor::core_loop::write_index::KeyRepr; @@ -25,7 +27,7 @@ pub(super) struct ReplayKvTransferParams<'a> { pub source_key: &'a [u8], pub dest_key: &'a [u8], pub field: &'a str, - pub amount: f64, + pub amount: TransferAmount, pub debit_surrogate: u32, pub credit_surrogate: u32, } @@ -91,8 +93,24 @@ impl CoreLoop { Some(dest_bytes.as_slice()) }; - let computed = match compute_transfer(&source_bytes, dest_ref, p.field, p.amount) { + let declared = self.declared_columns_of(p.database_id, p.tenant_id, p.collection); + let computed = match compute_transfer(&source_bytes, dest_ref, p.field, p.amount, declared) + { Ok(c) => c, + // The live write refused the same balance against the same + // pre-state, so the record converges to the same no-op. + Err(TransferError::Declared(error)) => { + warn!( + core = self.core_id, + collection = p.collection, + source_key = %String::from_utf8_lossy(p.source_key), + dest_key = %String::from_utf8_lossy(p.dest_key), + %error, + "WAL kv_transfer replay: declared column rule refused a balance, \ + skipping record" + ); + return 0; + } Err(TransferError::TypeMismatch(detail)) => { warn!( core = self.core_id, @@ -110,8 +128,8 @@ impl CoreLoop { collection = p.collection, source_key = %String::from_utf8_lossy(p.source_key), dest_key = %String::from_utf8_lossy(p.dest_key), - have, - need, + %have, + %need, "WAL kv_transfer replay: insufficient balance, skipping record" ); return 0; @@ -207,6 +225,25 @@ impl CoreLoop { ); return (0, 0); }; + // The row arriving at the destination meets the destination's + // declared numeric columns. The live move refused a row they refuse, + // so the record converges to the same no-op. + let declared = self.declared_columns_of(p.database_id, p.tenant_id, p.dest_collection); + let dest_data = match fit_kv_image(&item_data, declared) { + Ok(fitted) => fitted.unwrap_or(item_data), + Err(error) => { + warn!( + core = self.core_id, + source_collection = p.source_collection, + dest_collection = p.dest_collection, + item_key = %String::from_utf8_lossy(p.item_key), + %error, + "WAL kv_transfer_item replay: declared column rule refused the row, \ + skipping record" + ); + return (0, 0); + } + }; // The moved row is bound before the source row is deleted. let surrogate = nodedb_types::Surrogate::new(p.surrogate); @@ -231,7 +268,7 @@ impl CoreLoop { tenant_id: p.tenant_id, collection: p.dest_collection, key: p.dest_key, - value: &item_data, + value: &dest_data, ttl_ms: 0, now_ms: p.now_ms, surrogate, @@ -286,9 +323,16 @@ impl CoreLoop { amount, debit_surrogate, credit_surrogate, - ) = zerompk::from_msgpack::<(&str, String, Vec, Vec, String, f64, u32, u32)>( - payload, - ) + ) = zerompk::from_msgpack::<( + &str, + String, + Vec, + Vec, + String, + TransferAmount, + u32, + u32, + )>(payload) .ok()?; if disc != "kv_transfer" { return None; @@ -372,7 +416,7 @@ mod tests { use crate::control::server::wal_dispatch::wal_append_if_write; use crate::types::{DatabaseId, TenantId, VShardId}; use crate::wal::manager::WalManager; - use nodedb_physical::physical_plan::KvOp; + use nodedb_physical::physical_plan::{KvOp, TransferAmount}; use nodedb_types::{QualifiedCollection, RlsWriteCheck, Surrogate}; use nodedb_wal::TombstoneSet; @@ -471,7 +515,7 @@ mod tests { source_key: b"alice".to_vec(), dest_key: b"bob".to_vec(), field: "balance".into(), - amount: 30.0, + amount: TransferAmount::Float(30.0), debit_surrogate: Surrogate::new(1), credit_surrogate: Surrogate::new(2), rls_write_check: RlsWriteCheck::already_decided_elsewhere(), diff --git a/nodedb/tests/wire/cases/kv_transfer_decimal.rs b/nodedb/tests/wire/cases/kv_transfer_decimal.rs new file mode 100644 index 000000000..3e6d69846 --- /dev/null +++ b/nodedb/tests/wire/cases/kv_transfer_decimal.rs @@ -0,0 +1,85 @@ +// SPDX-License-Identifier: BUSL-1.1 + +//! `TRANSFER` on a KV `DECIMAL(p,s)` field moves the balance by exact +//! decimal arithmetic. An INT or DECIMAL amount is accepted. Each balance +//! rounds to the declared scale, half away from zero. A balance past the +//! declared precision is refused with SQLSTATE 22003, and neither row moves. +//! The balance check compares exact decimals. + +use crate::harness::TestServer; + +const OUT_OF_RANGE: &str = "SQLSTATE 22003"; + +/// The `balance` of row `key`. +async fn balance(srv: &TestServer, key: &str) -> String { + let rows = srv + .query_rows(&format!("SELECT balance FROM td_acct WHERE key = '{key}'")) + .await + .unwrap(); + assert_eq!(rows.len(), 1, "{key}: {rows:?}"); + rows[0][0].clone() +} + +async fn transfer(srv: &TestServer, from: &str, to: &str, amount: &str) -> Result<(), String> { + srv.exec(&format!( + "SELECT TRANSFER('td_acct', '{from}', '{to}', 'balance', {amount})" + )) + .await +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn kv_transfer_moves_a_declared_decimal_exactly() { + let srv = TestServer::start().await; + srv.exec( + "CREATE COLLECTION td_acct (key TEXT PRIMARY KEY, balance DECIMAL(10,2)) \ + WITH (engine='kv')", + ) + .await + .unwrap(); + for (key, value) in [("a", "100.00"), ("b", "5.25"), ("c", "99999999.99")] { + srv.exec(&format!( + "INSERT INTO td_acct (key, balance) VALUES ('{key}', {value})" + )) + .await + .unwrap(); + } + + // A DECIMAL amount moves both balances by exactly that amount. + transfer(&srv, "a", "b", "30.10").await.unwrap(); + assert_eq!(balance(&srv, "a").await, "69.90"); + assert_eq!(balance(&srv, "b").await, "35.35"); + + // Each balance rounds to scale 2, half away from zero: + // 69.90 - 0.005 = 69.895 -> 69.90, and 35.35 + 0.005 = 35.355 -> 35.36. + transfer(&srv, "a", "b", "0.005").await.unwrap(); + assert_eq!(balance(&srv, "a").await, "69.90"); + assert_eq!(balance(&srv, "b").await, "35.36"); + + // An INT amount moves a DECIMAL balance. + transfer(&srv, "a", "b", "1").await.unwrap(); + assert_eq!(balance(&srv, "a").await, "68.90"); + assert_eq!(balance(&srv, "b").await, "36.36"); + + // A credit past DECIMAL(10,2) is refused, and neither row moves. + srv.expect_error( + "SELECT TRANSFER('td_acct', 'a', 'c', 'balance', 1)", + OUT_OF_RANGE, + ) + .await; + assert_eq!(balance(&srv, "a").await, "68.90"); + assert_eq!(balance(&srv, "c").await, "99999999.99"); + + // The balance check is exact: 68.90 does not cover 68.91. + srv.expect_error( + "SELECT TRANSFER('td_acct', 'a', 'b', 'balance', 68.91)", + "source has 68.90, need 68.91", + ) + .await; + assert_eq!(balance(&srv, "a").await, "68.90"); + assert_eq!(balance(&srv, "b").await, "36.36"); + + // The whole balance moves when the amount equals it. + transfer(&srv, "a", "b", "68.90").await.unwrap(); + assert_eq!(balance(&srv, "a").await, "0.00"); + assert_eq!(balance(&srv, "b").await, "105.26"); +} From 3c7708d3abf847d70be3f588a3d9c07bbe868da6 Mon Sep 17 00:00:00 2001 From: Farhan Syah Date: Thu, 8 Oct 2026 16:51:46 +0800 Subject: [PATCH 08/24] refactor(sql): give large SqlPlan variants named payload structs Struct-like variants such as Insert, Upsert, Merge, the search plans, the array plans and the vector-primary writes wrap named *Plan structs defined per family under types/plan/variants. Planner, visitor and converter matches follow. The catalog fold and aggregate wrap passes split into module directories. --- nodedb-sql/src/engine_rules/columnar.rs | 8 +- .../src/engine_rules/document_schemaless.rs | 12 +- .../src/engine_rules/document_strict.rs | 12 +- nodedb-sql/src/engine_rules/index_lookup.rs | 4 +- nodedb-sql/src/engine_rules/spatial.rs | 8 +- nodedb-sql/src/engine_rules/timeseries.rs | 12 +- .../planner/aggregate_cp_wrap/expression.rs | 86 ++ .../src/planner/aggregate_cp_wrap/mod.rs | 9 + .../planner/aggregate_cp_wrap/projection.rs | 131 +++ .../wrap.rs} | 213 +---- nodedb-sql/src/planner/array_ddl.rs | 9 +- nodedb-sql/src/planner/array_dml.rs | 9 +- nodedb-sql/src/planner/array_fn/table_fn.rs | 34 +- .../src/planner/bitmap_emit/predicate.rs | 5 +- nodedb-sql/src/planner/catalog_fold/filter.rs | 45 + nodedb-sql/src/planner/catalog_fold/leaf.rs | 162 ++++ nodedb-sql/src/planner/catalog_fold/mod.rs | 9 + .../{catalog_fold.rs => catalog_fold/walk.rs} | 211 +---- .../src/planner/catalog_plan_validate.rs | 82 +- nodedb-sql/src/planner/cte/recursive_scan.rs | 4 +- nodedb-sql/src/planner/cte/recursive_value.rs | 4 +- nodedb-sql/src/planner/dml/insert.rs | 8 +- .../src/planner/dml_helpers/kv_counter.rs | 2 +- .../src/planner/dml_helpers/kv_insert.rs | 4 +- .../planner/dml_helpers/vector_primary_dml.rs | 36 +- .../dml_helpers/vector_primary_insert.rs | 26 +- nodedb-sql/src/planner/index_ddl/create.rs | 23 +- nodedb-sql/src/planner/index_ddl/drop.rs | 13 +- nodedb-sql/src/planner/lateral/plan.rs | 14 +- nodedb-sql/src/planner/select/derived_from.rs | 4 +- .../planner/select/order_by/vector_join.rs | 39 +- nodedb-sql/src/planner/select/post_process.rs | 31 +- nodedb-sql/src/types/mod.rs | 9 + nodedb-sql/src/types/plan/cacheability.rs | 60 +- nodedb-sql/src/types/plan/mod.rs | 10 + nodedb-sql/src/types/plan/variants.rs | 841 ------------------ nodedb-sql/src/types/plan/variants/array.rs | 94 ++ nodedb-sql/src/types/plan/variants/cte.rs | 14 + nodedb-sql/src/types/plan/variants/hybrid.rs | 71 ++ .../src/types/plan/variants/index_ddl.rs | 31 + .../src/types/plan/variants/index_reads.rs | 46 + nodedb-sql/src/types/plan/variants/lateral.rs | 55 ++ nodedb-sql/src/types/plan/variants/merge.rs | 26 + nodedb-sql/src/types/plan/variants/mod.rs | 37 + nodedb-sql/src/types/plan/variants/plan.rs | 471 ++++++++++ .../src/types/plan/variants/recursive.rs | 45 + .../src/types/plan/variants/timeseries.rs | 43 + .../src/types/plan/variants/vector_primary.rs | 81 ++ nodedb-sql/src/types/plan/variants/writes.rs | 89 ++ .../src/visitor/plan_visitor/dispatch.rs | 85 +- .../src/visitor/plan_visitor/dispatch_rest.rs | 74 +- .../sql_suite/cases/limit_offset_bounds.rs | 6 +- .../cases/positional_insert_column_binding.rs | 8 +- .../cases/truncate_engine_routing.rs | 6 +- .../sql_plan_convert/convert_array_arms.rs | 129 --- .../planner/sql_plan_convert/dml/merge.rs | 6 +- .../sql_plan_convert/dml/surrogate_keys.rs | 54 +- .../dml/update_delete/update_from.rs | 6 +- .../planner/sql_plan_convert/lateral.rs | 7 +- .../surrogate_prefetch/bind.rs | 6 +- .../control/server/shared/returning/clause.rs | 23 +- .../executor_tests/test_group_by_alias.rs | 14 +- 62 files changed, 2035 insertions(+), 1681 deletions(-) create mode 100644 nodedb-sql/src/planner/aggregate_cp_wrap/expression.rs create mode 100644 nodedb-sql/src/planner/aggregate_cp_wrap/mod.rs create mode 100644 nodedb-sql/src/planner/aggregate_cp_wrap/projection.rs rename nodedb-sql/src/planner/{aggregate_cp_wrap.rs => aggregate_cp_wrap/wrap.rs} (61%) create mode 100644 nodedb-sql/src/planner/catalog_fold/filter.rs create mode 100644 nodedb-sql/src/planner/catalog_fold/leaf.rs create mode 100644 nodedb-sql/src/planner/catalog_fold/mod.rs rename nodedb-sql/src/planner/{catalog_fold.rs => catalog_fold/walk.rs} (61%) delete mode 100644 nodedb-sql/src/types/plan/variants.rs create mode 100644 nodedb-sql/src/types/plan/variants/array.rs create mode 100644 nodedb-sql/src/types/plan/variants/cte.rs create mode 100644 nodedb-sql/src/types/plan/variants/hybrid.rs create mode 100644 nodedb-sql/src/types/plan/variants/index_ddl.rs create mode 100644 nodedb-sql/src/types/plan/variants/index_reads.rs create mode 100644 nodedb-sql/src/types/plan/variants/lateral.rs create mode 100644 nodedb-sql/src/types/plan/variants/merge.rs create mode 100644 nodedb-sql/src/types/plan/variants/mod.rs create mode 100644 nodedb-sql/src/types/plan/variants/plan.rs create mode 100644 nodedb-sql/src/types/plan/variants/recursive.rs create mode 100644 nodedb-sql/src/types/plan/variants/timeseries.rs create mode 100644 nodedb-sql/src/types/plan/variants/vector_primary.rs create mode 100644 nodedb-sql/src/types/plan/variants/writes.rs delete mode 100644 nodedb/src/control/planner/sql_plan_convert/convert_array_arms.rs diff --git a/nodedb-sql/src/engine_rules/columnar.rs b/nodedb-sql/src/engine_rules/columnar.rs index 2567192f6..a8646bbce 100644 --- a/nodedb-sql/src/engine_rules/columnar.rs +++ b/nodedb-sql/src/engine_rules/columnar.rs @@ -10,7 +10,7 @@ pub struct ColumnarRules; impl EngineRules for ColumnarRules { fn plan_insert(&self, p: InsertParams) -> Result> { - Ok(vec![SqlPlan::Insert { + Ok(vec![SqlPlan::Insert(InsertPlan { collection: p.collection, engine: EngineType::Columnar, route: WriteRoute::ColumnarFamily, @@ -19,7 +19,7 @@ impl EngineRules for ColumnarRules { if_absent: p.if_absent, column_schema: p.column_schema, primary_key: p.primary_key, - }]) + })]) } /// `UPSERT` / `INSERT ... ON CONFLICT (pk) DO UPDATE` on columnar. @@ -27,7 +27,7 @@ impl EngineRules for ColumnarRules { /// bitmap; the new row (or its merged form, when `on_conflict_updates` /// is non-empty) is appended. Sort-key semantics are unchanged. fn plan_upsert(&self, p: UpsertParams) -> Result> { - Ok(vec![SqlPlan::Upsert { + Ok(vec![SqlPlan::Upsert(UpsertPlan { collection: p.collection, engine: EngineType::Columnar, route: WriteRoute::ColumnarFamily, @@ -36,7 +36,7 @@ impl EngineRules for ColumnarRules { on_conflict_updates: p.on_conflict_updates, column_schema: p.column_schema, primary_key: p.primary_key, - }]) + })]) } fn plan_scan(&self, p: ScanParams) -> Result { diff --git a/nodedb-sql/src/engine_rules/document_schemaless.rs b/nodedb-sql/src/engine_rules/document_schemaless.rs index 7f7f1d315..34e2af8c0 100644 --- a/nodedb-sql/src/engine_rules/document_schemaless.rs +++ b/nodedb-sql/src/engine_rules/document_schemaless.rs @@ -10,7 +10,7 @@ pub struct SchemalessRules; impl EngineRules for SchemalessRules { fn plan_insert(&self, p: InsertParams) -> Result> { - Ok(vec![SqlPlan::Insert { + Ok(vec![SqlPlan::Insert(InsertPlan { collection: p.collection, engine: EngineType::DocumentSchemaless, route: WriteRoute::Document, @@ -19,11 +19,11 @@ impl EngineRules for SchemalessRules { if_absent: p.if_absent, column_schema: vec![], primary_key: p.primary_key, - }]) + })]) } fn plan_upsert(&self, p: UpsertParams) -> Result> { - Ok(vec![SqlPlan::Upsert { + Ok(vec![SqlPlan::Upsert(UpsertPlan { collection: p.collection, engine: EngineType::DocumentSchemaless, route: WriteRoute::Document, @@ -32,7 +32,7 @@ impl EngineRules for SchemalessRules { on_conflict_updates: p.on_conflict_updates, column_schema: vec![], primary_key: p.primary_key, - }]) + })]) } fn plan_scan(&self, p: ScanParams) -> Result { @@ -110,7 +110,7 @@ impl EngineRules for SchemalessRules { } fn plan_merge(&self, p: MergeParams) -> Result> { - Ok(vec![SqlPlan::Merge { + Ok(vec![SqlPlan::Merge(MergePlan { target: p.collection, engine: EngineType::DocumentSchemaless, source: p.source, @@ -119,7 +119,7 @@ impl EngineRules for SchemalessRules { source_alias: p.source_alias, clauses: p.clauses, returning: p.returning, - }]) + })]) } fn plan_truncate(&self, p: TruncateParams) -> Result> { diff --git a/nodedb-sql/src/engine_rules/document_strict.rs b/nodedb-sql/src/engine_rules/document_strict.rs index ad447fe07..874baf10c 100644 --- a/nodedb-sql/src/engine_rules/document_strict.rs +++ b/nodedb-sql/src/engine_rules/document_strict.rs @@ -10,7 +10,7 @@ pub struct StrictRules; impl EngineRules for StrictRules { fn plan_insert(&self, p: InsertParams) -> Result> { - Ok(vec![SqlPlan::Insert { + Ok(vec![SqlPlan::Insert(InsertPlan { collection: p.collection, engine: EngineType::DocumentStrict, route: WriteRoute::Document, @@ -19,11 +19,11 @@ impl EngineRules for StrictRules { if_absent: p.if_absent, column_schema: vec![], primary_key: p.primary_key, - }]) + })]) } fn plan_upsert(&self, p: UpsertParams) -> Result> { - Ok(vec![SqlPlan::Upsert { + Ok(vec![SqlPlan::Upsert(UpsertPlan { collection: p.collection, engine: EngineType::DocumentStrict, route: WriteRoute::Document, @@ -32,7 +32,7 @@ impl EngineRules for StrictRules { on_conflict_updates: p.on_conflict_updates, column_schema: vec![], primary_key: p.primary_key, - }]) + })]) } fn plan_scan(&self, p: ScanParams) -> Result { @@ -110,7 +110,7 @@ impl EngineRules for StrictRules { } fn plan_merge(&self, p: MergeParams) -> Result> { - Ok(vec![SqlPlan::Merge { + Ok(vec![SqlPlan::Merge(MergePlan { target: p.collection, engine: EngineType::DocumentStrict, source: p.source, @@ -119,7 +119,7 @@ impl EngineRules for StrictRules { source_alias: p.source_alias, clauses: p.clauses, returning: p.returning, - }]) + })]) } fn plan_truncate(&self, p: TruncateParams) -> Result> { diff --git a/nodedb-sql/src/engine_rules/index_lookup.rs b/nodedb-sql/src/engine_rules/index_lookup.rs index 760193aa3..77c93987d 100644 --- a/nodedb-sql/src/engine_rules/index_lookup.rs +++ b/nodedb-sql/src/engine_rules/index_lookup.rs @@ -90,7 +90,7 @@ pub(crate) fn try_document_index_lookup( value }; - return Some(SqlPlan::DocumentIndexLookup { + return Some(SqlPlan::DocumentIndexLookup(DocumentIndexLookupPlan { collection: params.collection.clone(), alias: params.alias.clone(), engine, @@ -105,7 +105,7 @@ pub(crate) fn try_document_index_lookup( window_functions: params.window_functions.clone(), case_insensitive: idx.case_insensitive, temporal: params.temporal, - }); + })); } None } diff --git a/nodedb-sql/src/engine_rules/spatial.rs b/nodedb-sql/src/engine_rules/spatial.rs index 59bf058fe..0820c7b02 100644 --- a/nodedb-sql/src/engine_rules/spatial.rs +++ b/nodedb-sql/src/engine_rules/spatial.rs @@ -10,7 +10,7 @@ pub struct SpatialRules; impl EngineRules for SpatialRules { fn plan_insert(&self, p: InsertParams) -> Result> { - Ok(vec![SqlPlan::Insert { + Ok(vec![SqlPlan::Insert(InsertPlan { collection: p.collection, engine: EngineType::Spatial, route: WriteRoute::ColumnarFamily, @@ -19,14 +19,14 @@ impl EngineRules for SpatialRules { if_absent: p.if_absent, column_schema: p.column_schema, primary_key: p.primary_key, - }]) + })]) } /// Spatial extends columnar and inherits the same upsert semantics: /// duplicate PK tombstones the prior row; the new row (or merged form /// when `on_conflict_updates` is non-empty) is appended. fn plan_upsert(&self, p: UpsertParams) -> Result> { - Ok(vec![SqlPlan::Upsert { + Ok(vec![SqlPlan::Upsert(UpsertPlan { collection: p.collection, engine: EngineType::Spatial, route: WriteRoute::ColumnarFamily, @@ -35,7 +35,7 @@ impl EngineRules for SpatialRules { on_conflict_updates: p.on_conflict_updates, column_schema: p.column_schema, primary_key: p.primary_key, - }]) + })]) } fn plan_scan(&self, p: ScanParams) -> Result { diff --git a/nodedb-sql/src/engine_rules/timeseries.rs b/nodedb-sql/src/engine_rules/timeseries.rs index f5ea41f9a..e1e40214e 100644 --- a/nodedb-sql/src/engine_rules/timeseries.rs +++ b/nodedb-sql/src/engine_rules/timeseries.rs @@ -15,11 +15,11 @@ pub struct TimeseriesRules; impl EngineRules for TimeseriesRules { fn plan_insert(&self, p: InsertParams) -> Result> { // Timeseries INSERT routes to TimeseriesIngest — append-only semantics. - Ok(vec![SqlPlan::TimeseriesIngest { + Ok(vec![SqlPlan::TimeseriesIngest(TimeseriesIngestPlan { collection: p.collection, rows: p.rows, volatile_defaults: p.volatile_defaults, - }]) + })]) } fn plan_upsert(&self, _p: UpsertParams) -> Result> { @@ -40,7 +40,7 @@ impl EngineRules for TimeseriesRules { } // Timeseries scans use TimeseriesScan for time-range-aware execution. let time_range = default_time_range(); - Ok(SqlPlan::TimeseriesScan { + Ok(SqlPlan::TimeseriesScan(TimeseriesScanPlan { collection: p.collection, time_range, bucket_interval_ms: 0, @@ -53,7 +53,7 @@ impl EngineRules for TimeseriesRules { sort_keys: p.sort_keys, tiered: false, temporal: p.temporal, - }) + })) } fn plan_point_get(&self, _p: PointGetParams) -> Result { @@ -111,7 +111,7 @@ impl EngineRules for TimeseriesRules { ), }); } - Ok(SqlPlan::TimeseriesScan { + Ok(SqlPlan::TimeseriesScan(TimeseriesScanPlan { collection: p.collection, time_range: default_time_range(), bucket_interval_ms: p.bucket_interval_ms.unwrap_or(0), @@ -126,7 +126,7 @@ impl EngineRules for TimeseriesRules { sort_keys: Vec::new(), tiered: p.has_auto_tier, temporal: p.temporal, - }) + })) } fn plan_merge(&self, p: MergeParams) -> Result> { diff --git a/nodedb-sql/src/planner/aggregate_cp_wrap/expression.rs b/nodedb-sql/src/planner/aggregate_cp_wrap/expression.rs new file mode 100644 index 000000000..8105fb41f --- /dev/null +++ b/nodedb-sql/src/planner/aggregate_cp_wrap/expression.rs @@ -0,0 +1,86 @@ +// SPDX-License-Identifier: Apache-2.0 + +//! Group-row column qualification. + +use crate::types::SqlExpr; + +/// `expr` with every column reference stripped of its table qualifier. A +/// finalized group row keys its columns by bare name. +pub(super) fn unqualify_columns(expr: SqlExpr) -> SqlExpr { + match expr { + SqlExpr::Column { name, .. } => SqlExpr::Column { table: None, name }, + SqlExpr::Function { + name, + args, + distinct, + } => SqlExpr::Function { + name, + args: args.into_iter().map(unqualify_columns).collect(), + distinct, + }, + SqlExpr::BinaryOp { left, op, right } => SqlExpr::BinaryOp { + left: Box::new(unqualify_columns(*left)), + op, + right: Box::new(unqualify_columns(*right)), + }, + SqlExpr::UnaryOp { op, expr } => SqlExpr::UnaryOp { + op, + expr: Box::new(unqualify_columns(*expr)), + }, + SqlExpr::Cast { expr, to_type } => SqlExpr::Cast { + expr: Box::new(unqualify_columns(*expr)), + to_type, + }, + SqlExpr::IsNull { expr, negated } => SqlExpr::IsNull { + expr: Box::new(unqualify_columns(*expr)), + negated, + }, + SqlExpr::Case { + operand, + when_then, + else_expr, + } => SqlExpr::Case { + operand: operand.map(|e| Box::new(unqualify_columns(*e))), + when_then: when_then + .into_iter() + .map(|(when, then)| (unqualify_columns(when), unqualify_columns(then))) + .collect(), + else_expr: else_expr.map(|e| Box::new(unqualify_columns(*e))), + }, + SqlExpr::InList { + expr, + list, + negated, + } => SqlExpr::InList { + expr: Box::new(unqualify_columns(*expr)), + list: list.into_iter().map(unqualify_columns).collect(), + negated, + }, + SqlExpr::Between { + expr, + low, + high, + negated, + } => SqlExpr::Between { + expr: Box::new(unqualify_columns(*expr)), + low: Box::new(unqualify_columns(*low)), + high: Box::new(unqualify_columns(*high)), + negated, + }, + SqlExpr::Like { + expr, + pattern, + negated, + case_insensitive, + } => SqlExpr::Like { + expr: Box::new(unqualify_columns(*expr)), + pattern: Box::new(unqualify_columns(*pattern)), + negated, + case_insensitive, + }, + SqlExpr::ArrayLiteral(items) => { + SqlExpr::ArrayLiteral(items.into_iter().map(unqualify_columns).collect()) + } + SqlExpr::Literal(_) | SqlExpr::Subquery(_) | SqlExpr::Wildcard => expr, + } +} diff --git a/nodedb-sql/src/planner/aggregate_cp_wrap/mod.rs b/nodedb-sql/src/planner/aggregate_cp_wrap/mod.rs new file mode 100644 index 000000000..a5372ecfd --- /dev/null +++ b/nodedb-sql/src/planner/aggregate_cp_wrap/mod.rs @@ -0,0 +1,9 @@ +// SPDX-License-Identifier: Apache-2.0 + +//! Control-Plane projection over grouped results. + +mod expression; +mod projection; +mod wrap; + +pub use wrap::wrap_aggregate_cp_items; diff --git a/nodedb-sql/src/planner/aggregate_cp_wrap/projection.rs b/nodedb-sql/src/planner/aggregate_cp_wrap/projection.rs new file mode 100644 index 000000000..043424ee6 --- /dev/null +++ b/nodedb-sql/src/planner/aggregate_cp_wrap/projection.rs @@ -0,0 +1,131 @@ +// SPDX-License-Identifier: Apache-2.0 + +//! Grouped output projection and sequence-accessor checks. + +use super::expression::unqualify_columns; +use crate::aggregate_walk::contains_aggregate; +use crate::error::{Result, SqlError}; +use crate::functions::registry::FunctionRegistry; +use crate::parser::normalize::normalize_ident; +use crate::planner::agg_naming::group_key_row_name; +use crate::planner::aggregate_order::compute_output_order_by_item; +use crate::planner::cp_projection::ast_calls_sequence_accessor; +use crate::resolver::ColumnScope; +use crate::resolver::columns::TableScope; +use crate::resolver::expr::convert_expr; +use crate::types::plan::{first_sequence_accessor, referenced_columns}; +use crate::types::{AggOutputSlot, AggregateExpr, Projection, SqlExpr, SqlPlan}; +use sqlparser::ast; + +/// The projection restating `plan`'s output columns in SELECT-list order, +/// with a [`Projection::CpComputed`] entry at each accessor item's position. +/// `plan` is the `Aggregate` the items were planned into. +pub(super) fn aggregate_cp_projection( + plan: &SqlPlan, + items: &[ast::SelectItem], + functions: &FunctionRegistry, + scope: &TableScope, +) -> Result> { + let (group_by, aggregates) = match plan { + SqlPlan::Aggregate { + group_by, + aggregates, + .. + } => (group_by, aggregates), + other => { + return Err(SqlError::Unsupported { + detail: format!("aggregate wrap over a {} plan", other.variant_name()), + }); + } + }; + let key_names: Vec = group_by + .iter() + .enumerate() + .map(|(index, key)| group_key_row_name(key, index)) + .collect(); + let by_item = compute_output_order_by_item(items, group_by, functions, scope)?; + // The item resolves against the input relations, so an unknown column is + // the usual resolve error, and against the group-key row names, so a + // computed key (`group_0`) is addressable. The reference check below + // then narrows to the keys: the finalized group row carries nothing + // else the Control Plane can read. + let cp_scope = scope + .with_output_names(key_names.iter().cloned()) + .allowing_cp_functions(); + + let mut projection = Vec::with_capacity(items.len()); + for (item, slots) in items.iter().zip(by_item) { + let (expr, alias) = match item { + ast::SelectItem::UnnamedExpr(expr) => (expr, format!("{expr}").to_lowercase()), + ast::SelectItem::ExprWithAlias { expr, alias } => (expr, normalize_ident(alias)), + ast::SelectItem::ExprWithAliases { .. } + | ast::SelectItem::Wildcard(_) + | ast::SelectItem::QualifiedWildcard(..) => continue, + }; + if ast_calls_sequence_accessor(expr) { + let converted = convert_expr(expr, &ColumnScope::Relations(&cp_scope))?; + if let Some(name) = first_sequence_accessor(&converted) { + let name = name.to_string(); + projection.push(cp_item( + converted, name, alias, expr, &key_names, functions, + )?); + continue; + } + } + for slot in slots { + projection.push(Projection::Column(slot_row_name( + slot, &key_names, aggregates, + )?)); + } + } + Ok(projection) +} + +/// The Control-Plane entry for one accessor item over a grouped result. +/// +/// The item holds no aggregate: the Control Plane evaluates it over the +/// finalized group row, which carries aggregate values under their output +/// names but no per-row aggregate state. Every column it references is a +/// group key, unqualified so the reference matches the row's bare key. +fn cp_item( + converted: SqlExpr, + accessor: String, + alias: String, + raw: &ast::Expr, + key_names: &[String], + functions: &FunctionRegistry, +) -> Result { + if contains_aggregate(raw, functions) { + return Err(SqlError::SequencePerRowUnsupported { name: accessor }); + } + for column in referenced_columns(&converted) { + let bare = column.rsplit('.').next().unwrap_or(&column); + if !key_names.iter().any(|key| key.eq_ignore_ascii_case(bare)) { + return Err(SqlError::Unsupported { + detail: format!( + "column '{column}' beside a sequence accessor in a grouped SELECT list \ + must be a GROUP BY key" + ), + }); + } + } + Ok(Projection::CpComputed { + expr: unqualify_columns(converted), + alias, + }) +} + +/// The key a finalized group row carries one output slot under. +fn slot_row_name( + slot: AggOutputSlot, + key_names: &[String], + aggregates: &[AggregateExpr], +) -> Result { + match slot { + AggOutputSlot::GroupKey(index) => key_names.get(index).cloned(), + AggOutputSlot::Aggregate(index) => aggregates.get(index).map(|a| a.alias.clone()), + } + .ok_or_else(|| SqlError::Unsupported { + detail: format!("aggregate output slot {slot:?} names no output column"), + }) +} diff --git a/nodedb-sql/src/planner/aggregate_cp_wrap.rs b/nodedb-sql/src/planner/aggregate_cp_wrap/wrap.rs similarity index 61% rename from nodedb-sql/src/planner/aggregate_cp_wrap.rs rename to nodedb-sql/src/planner/aggregate_cp_wrap/wrap.rs index 51cab4b56..7a71c603e 100644 --- a/nodedb-sql/src/planner/aggregate_cp_wrap.rs +++ b/nodedb-sql/src/planner/aggregate_cp_wrap/wrap.rs @@ -7,23 +7,18 @@ //! item that calls a sequence accessor is neither, so the aggregate cannot //! carry it. The planner wraps the finished aggregate in a `Subquery` whose //! projection restates every output column in SELECT-list order and places a -//! [`Projection::CpComputed`] entry at the item's position. The wrap runs +//! [`crate::types::Projection::CpComputed`] entry at the item's position. The wrap runs //! after ORDER BY and LIMIT are attached, so those stay on the aggregate. +use crate::types::CtePlan; use sqlparser::ast; -use crate::aggregate_walk::contains_aggregate; +use super::projection::aggregate_cp_projection; use crate::error::{Result, SqlError}; use crate::functions::registry::FunctionRegistry; -use crate::parser::normalize::normalize_ident; -use crate::planner::agg_naming::group_key_row_name; -use crate::planner::aggregate_order::compute_output_order_by_item; -use crate::planner::cp_projection::{ast_calls_sequence_accessor, ast_sequence_accessor}; -use crate::resolver::ColumnScope; +use crate::planner::cp_projection::ast_sequence_accessor; use crate::resolver::columns::TableScope; -use crate::resolver::expr::convert_expr; -use crate::types::plan::{first_sequence_accessor, referenced_columns}; -use crate::types::{AggOutputSlot, AggregateExpr, Projection, SqlExpr, SqlPlan}; +use crate::types::SqlPlan; /// Wrap `plan` when `items` hold a sequence accessor the aggregate cannot /// carry. A plan with no such item, or one that is not grouped, returns @@ -52,12 +47,12 @@ pub fn wrap_aggregate_cp_items( return Err(SqlError::SequencePerRowUnsupported { name: accessor }); } match plan { - SqlPlan::Cte { definitions, outer } => Ok(SqlPlan::Cte { + SqlPlan::Cte(CtePlan { definitions, outer }) => Ok(SqlPlan::Cte(CtePlan { definitions, outer: Box::new(wrap_aggregate_cp_items( *outer, items, grouped, functions, scope, )?), - }), + })), SqlPlan::Aggregate { .. } => { let projection = aggregate_cp_projection(&plan, items, functions, scope)?; Ok(SqlPlan::Subquery { @@ -123,200 +118,6 @@ fn item_accessor(item: &ast::SelectItem) -> Option { } } -/// The projection restating `plan`'s output columns in SELECT-list order, -/// with a [`Projection::CpComputed`] entry at each accessor item's position. -/// `plan` is the `Aggregate` the items were planned into. -fn aggregate_cp_projection( - plan: &SqlPlan, - items: &[ast::SelectItem], - functions: &FunctionRegistry, - scope: &TableScope, -) -> Result> { - let (group_by, aggregates) = match plan { - SqlPlan::Aggregate { - group_by, - aggregates, - .. - } => (group_by, aggregates), - other => { - return Err(SqlError::Unsupported { - detail: format!("aggregate wrap over a {} plan", other.variant_name()), - }); - } - }; - let key_names: Vec = group_by - .iter() - .enumerate() - .map(|(index, key)| group_key_row_name(key, index)) - .collect(); - let by_item = compute_output_order_by_item(items, group_by, functions, scope)?; - // The item resolves against the input relations, so an unknown column is - // the usual resolve error, and against the group-key row names, so a - // computed key (`group_0`) is addressable. The reference check below - // then narrows to the keys: the finalized group row carries nothing - // else the Control Plane can read. - let cp_scope = scope - .with_output_names(key_names.iter().cloned()) - .allowing_cp_functions(); - - let mut projection = Vec::with_capacity(items.len()); - for (item, slots) in items.iter().zip(by_item) { - let (expr, alias) = match item { - ast::SelectItem::UnnamedExpr(expr) => (expr, format!("{expr}").to_lowercase()), - ast::SelectItem::ExprWithAlias { expr, alias } => (expr, normalize_ident(alias)), - ast::SelectItem::ExprWithAliases { .. } - | ast::SelectItem::Wildcard(_) - | ast::SelectItem::QualifiedWildcard(..) => continue, - }; - if ast_calls_sequence_accessor(expr) { - let converted = convert_expr(expr, &ColumnScope::Relations(&cp_scope))?; - if let Some(name) = first_sequence_accessor(&converted) { - let name = name.to_string(); - projection.push(cp_item( - converted, name, alias, expr, &key_names, functions, - )?); - continue; - } - } - for slot in slots { - projection.push(Projection::Column(slot_row_name( - slot, &key_names, aggregates, - )?)); - } - } - Ok(projection) -} - -/// The Control-Plane entry for one accessor item over a grouped result. -/// -/// The item holds no aggregate: the Control Plane evaluates it over the -/// finalized group row, which carries aggregate values under their output -/// names but no per-row aggregate state. Every column it references is a -/// group key, unqualified so the reference matches the row's bare key. -fn cp_item( - converted: SqlExpr, - accessor: String, - alias: String, - raw: &ast::Expr, - key_names: &[String], - functions: &FunctionRegistry, -) -> Result { - if contains_aggregate(raw, functions) { - return Err(SqlError::SequencePerRowUnsupported { name: accessor }); - } - for column in referenced_columns(&converted) { - let bare = column.rsplit('.').next().unwrap_or(&column); - if !key_names.iter().any(|key| key.eq_ignore_ascii_case(bare)) { - return Err(SqlError::Unsupported { - detail: format!( - "column '{column}' beside a sequence accessor in a grouped SELECT list \ - must be a GROUP BY key" - ), - }); - } - } - Ok(Projection::CpComputed { - expr: unqualify_columns(converted), - alias, - }) -} - -/// The key a finalized group row carries one output slot under. -fn slot_row_name( - slot: AggOutputSlot, - key_names: &[String], - aggregates: &[AggregateExpr], -) -> Result { - match slot { - AggOutputSlot::GroupKey(index) => key_names.get(index).cloned(), - AggOutputSlot::Aggregate(index) => aggregates.get(index).map(|a| a.alias.clone()), - } - .ok_or_else(|| SqlError::Unsupported { - detail: format!("aggregate output slot {slot:?} names no output column"), - }) -} - -/// `expr` with every column reference stripped of its table qualifier. A -/// finalized group row keys its columns by bare name. -fn unqualify_columns(expr: SqlExpr) -> SqlExpr { - match expr { - SqlExpr::Column { name, .. } => SqlExpr::Column { table: None, name }, - SqlExpr::Function { - name, - args, - distinct, - } => SqlExpr::Function { - name, - args: args.into_iter().map(unqualify_columns).collect(), - distinct, - }, - SqlExpr::BinaryOp { left, op, right } => SqlExpr::BinaryOp { - left: Box::new(unqualify_columns(*left)), - op, - right: Box::new(unqualify_columns(*right)), - }, - SqlExpr::UnaryOp { op, expr } => SqlExpr::UnaryOp { - op, - expr: Box::new(unqualify_columns(*expr)), - }, - SqlExpr::Cast { expr, to_type } => SqlExpr::Cast { - expr: Box::new(unqualify_columns(*expr)), - to_type, - }, - SqlExpr::IsNull { expr, negated } => SqlExpr::IsNull { - expr: Box::new(unqualify_columns(*expr)), - negated, - }, - SqlExpr::Case { - operand, - when_then, - else_expr, - } => SqlExpr::Case { - operand: operand.map(|e| Box::new(unqualify_columns(*e))), - when_then: when_then - .into_iter() - .map(|(when, then)| (unqualify_columns(when), unqualify_columns(then))) - .collect(), - else_expr: else_expr.map(|e| Box::new(unqualify_columns(*e))), - }, - SqlExpr::InList { - expr, - list, - negated, - } => SqlExpr::InList { - expr: Box::new(unqualify_columns(*expr)), - list: list.into_iter().map(unqualify_columns).collect(), - negated, - }, - SqlExpr::Between { - expr, - low, - high, - negated, - } => SqlExpr::Between { - expr: Box::new(unqualify_columns(*expr)), - low: Box::new(unqualify_columns(*low)), - high: Box::new(unqualify_columns(*high)), - negated, - }, - SqlExpr::Like { - expr, - pattern, - negated, - case_insensitive, - } => SqlExpr::Like { - expr: Box::new(unqualify_columns(*expr)), - pattern: Box::new(unqualify_columns(*pattern)), - negated, - case_insensitive, - }, - SqlExpr::ArrayLiteral(items) => { - SqlExpr::ArrayLiteral(items.into_iter().map(unqualify_columns).collect()) - } - SqlExpr::Literal(_) | SqlExpr::Subquery(_) | SqlExpr::Wildcard => expr, - } -} - #[cfg(test)] mod tests { use crate::types::{ diff --git a/nodedb-sql/src/planner/array_ddl.rs b/nodedb-sql/src/planner/array_ddl.rs index 6b7e5020c..ba06f0240 100644 --- a/nodedb-sql/src/planner/array_ddl.rs +++ b/nodedb-sql/src/planner/array_ddl.rs @@ -8,6 +8,7 @@ //! schema-hash + the typed `ArraySchema` are computed in the Origin //! converter where `nodedb-array` is available. +use crate::types::{AlterArrayPlan, CreateArrayPlan}; use std::collections::HashSet; use crate::error::{Result, SqlError}; @@ -102,7 +103,7 @@ pub fn plan_create_array(ast: &CreateArrayAst) -> Result { ), }); } - Ok(SqlPlan::CreateArray { + Ok(SqlPlan::CreateArray(CreateArrayPlan { name: ast.name.clone(), dims: ast.dims.clone(), attrs: ast.attrs.clone(), @@ -112,7 +113,7 @@ pub fn plan_create_array(ast: &CreateArrayAst) -> Result { prefix_bits: ast.prefix_bits, audit_retain_ms: ast.audit_retain_ms, minimum_audit_retain_ms: ast.minimum_audit_retain_ms, - }) + })) } pub fn plan_alter_array(ast: &AlterArrayAst) -> Result { @@ -156,11 +157,11 @@ pub fn plan_alter_array(ast: &AlterArrayAst) -> Result { } } - Ok(SqlPlan::AlterArray { + Ok(SqlPlan::AlterArray(AlterArrayPlan { name: ast.name.clone(), audit_retain_ms, minimum_audit_retain_ms, - }) + })) } pub fn plan_drop_array(ast: &DropArrayAst) -> Result { diff --git a/nodedb-sql/src/planner/array_dml.rs b/nodedb-sql/src/planner/array_dml.rs index fb86c59e7..e812b76fc 100644 --- a/nodedb-sql/src/planner/array_dml.rs +++ b/nodedb-sql/src/planner/array_dml.rs @@ -12,6 +12,7 @@ use crate::catalog::{ArrayCatalogView, SqlCatalog}; use crate::error::{Result, SqlError}; use crate::parser::array_stmt::{DeleteArrayAst, InsertArrayAst}; use crate::types::SqlPlan; +use crate::types::{DeleteArrayPlan, InsertArrayPlan}; use crate::types_array::{ArrayAttrLiteral, ArrayAttrType, ArrayCoordLiteral, ArrayDimType}; pub fn plan_insert_array(ast: &InsertArrayAst, catalog: &dyn SqlCatalog) -> Result> { @@ -29,10 +30,10 @@ pub fn plan_insert_array(ast: &InsertArrayAst, catalog: &dyn SqlCatalog) -> Resu validate_coords(&ast.name, ri, &row.coords, &view)?; validate_attrs(&ast.name, ri, &row.attrs, &view)?; } - Ok(vec![SqlPlan::InsertArray { + Ok(vec![SqlPlan::InsertArray(InsertArrayPlan { name: ast.name.clone(), rows: ast.rows.clone(), - }]) + })]) } pub fn plan_delete_array(ast: &DeleteArrayAst, catalog: &dyn SqlCatalog) -> Result> { @@ -52,10 +53,10 @@ pub fn plan_delete_array(ast: &DeleteArrayAst, catalog: &dyn SqlCatalog) -> Resu for (ri, row) in ast.coords.iter().enumerate() { validate_coords(&ast.name, ri, row, &view)?; } - Ok(vec![SqlPlan::DeleteArray { + Ok(vec![SqlPlan::DeleteArray(DeleteArrayPlan { name: ast.name.clone(), coords: ast.coords.clone(), - }]) + })]) } fn validate_coords( diff --git a/nodedb-sql/src/planner/array_fn/table_fn.rs b/nodedb-sql/src/planner/array_fn/table_fn.rs index e7aa6f669..bd947ca70 100644 --- a/nodedb-sql/src/planner/array_fn/table_fn.rs +++ b/nodedb-sql/src/planner/array_fn/table_fn.rs @@ -3,6 +3,7 @@ //! `SELECT * FROM ARRAY_*(...)` table-valued function planning: //! slice / project / aggregate / elementwise. +use crate::types::{ArrayAggPlan, ArrayElementwisePlan, ArrayProjectPlan, ArraySlicePlan}; use sqlparser::ast; use nodedb_types::Value; @@ -145,13 +146,13 @@ fn plan_slice( 0 }; - Ok(SqlPlan::ArraySlice { + Ok(SqlPlan::ArraySlice(ArraySlicePlan { name, slice: ArraySliceAst { dim_ranges }, attr_projection, limit, temporal, - }) + })) } fn plan_project(args: &[ast::Expr], catalog: &dyn SqlCatalog) -> Result { @@ -182,10 +183,10 @@ fn plan_project(args: &[ast::Expr], catalog: &dyn SqlCatalog) -> Result }); } } - Ok(SqlPlan::ArrayProject { + Ok(SqlPlan::ArrayProject(ArrayProjectPlan { name, attr_projection, - }) + })) } fn plan_agg( @@ -232,13 +233,13 @@ fn plan_agg( None }; - Ok(SqlPlan::ArrayAgg { + Ok(SqlPlan::ArrayAgg(ArrayAggPlan { name, attr, reducer, group_by_dim, temporal, - }) + })) } fn plan_elementwise(args: &[ast::Expr], catalog: &dyn SqlCatalog) -> Result { @@ -279,12 +280,12 @@ fn plan_elementwise(args: &[ast::Expr], catalog: &dyn SqlCatalog) -> Result { + }) => { assert_eq!(name, "g"); assert_eq!(slice.dim_ranges.len(), 2); assert_eq!(attr_projection, vec!["qual".to_string()]); @@ -410,10 +412,10 @@ mod tests { fn project_happy() { let p = plan_one("SELECT * FROM ARRAY_PROJECT('g', ['qual', 'variant'])").unwrap(); match p { - SqlPlan::ArrayProject { + SqlPlan::ArrayProject(ArrayProjectPlan { name, attr_projection, - } => { + }) => { assert_eq!(name, "g"); assert_eq!(attr_projection, vec!["qual".to_string(), "variant".into()]); } @@ -430,13 +432,13 @@ mod tests { fn agg_scalar() { let p = plan_one("SELECT * FROM ARRAY_AGG('g', 'qual', 'sum')").unwrap(); match p { - SqlPlan::ArrayAgg { + SqlPlan::ArrayAgg(ArrayAggPlan { name, attr, reducer, group_by_dim, .. - } => { + }) => { assert_eq!(name, "g"); assert_eq!(attr, "qual"); assert_eq!(reducer, ArrayReducerAst::Sum); @@ -450,11 +452,11 @@ mod tests { fn agg_grouped() { let p = plan_one("SELECT * FROM ARRAY_AGG('g', 'qual', 'mean', 'chrom')").unwrap(); match p { - SqlPlan::ArrayAgg { + SqlPlan::ArrayAgg(ArrayAggPlan { reducer, group_by_dim, .. - } => { + }) => { assert_eq!(reducer, ArrayReducerAst::Mean); assert_eq!(group_by_dim, Some("chrom".into())); } diff --git a/nodedb-sql/src/planner/bitmap_emit/predicate.rs b/nodedb-sql/src/planner/bitmap_emit/predicate.rs index e737f1e20..d7edb67e3 100644 --- a/nodedb-sql/src/planner/bitmap_emit/predicate.rs +++ b/nodedb-sql/src/planner/bitmap_emit/predicate.rs @@ -16,6 +16,7 @@ //! The hint carries the collection name, index field, and the predicate value(s) //! so the converter layer can build an `IndexedFetch` physical sub-plan. +use crate::types::DocumentIndexLookupPlan; use crate::types::{CompareOp, FilterExpr, SqlPlan, SqlValue}; /// Maximum IN-list cardinality for which bitmap emission is attempted. @@ -42,12 +43,12 @@ pub struct BitmapHint { pub fn analyze(plan: &SqlPlan) -> Option { match plan { // Already a single-field equality index lookup — always qualifies. - SqlPlan::DocumentIndexLookup { + SqlPlan::DocumentIndexLookup(DocumentIndexLookupPlan { collection, field, value, .. - } => Some(BitmapHint { + }) => Some(BitmapHint { collection: collection.clone(), field: field.clone(), primary_value: value.clone(), diff --git a/nodedb-sql/src/planner/catalog_fold/filter.rs b/nodedb-sql/src/planner/catalog_fold/filter.rs new file mode 100644 index 000000000..91f002e10 --- /dev/null +++ b/nodedb-sql/src/planner/catalog_fold/filter.rs @@ -0,0 +1,45 @@ +// SPDX-License-Identifier: Apache-2.0 + +//! Catalog folding within filter trees. + +use crate::catalog::SqlCatalog; +use crate::planner::catalog_expr_fold::fold_expr; +use crate::types::{Filter, FilterExpr, SqlExpr}; +use nodedb_types::DatabaseId; + +pub(super) fn fold_filter( + filter: &mut Filter, + catalog: &dyn SqlCatalog, + database_id: DatabaseId, + tenant_id: u64, +) { + fold_filter_expr(&mut filter.expr, catalog, database_id, tenant_id); +} + +fn fold_filter_expr( + expr: &mut FilterExpr, + catalog: &dyn SqlCatalog, + database_id: DatabaseId, + tenant_id: u64, +) { + match expr { + FilterExpr::Expr(sql_expr) => { + let owned = std::mem::replace(sql_expr, SqlExpr::Wildcard); + *sql_expr = fold_expr(owned, catalog, database_id, tenant_id); + } + FilterExpr::And(children) | FilterExpr::Or(children) => { + for child in children { + fold_filter_expr(&mut child.expr, catalog, database_id, tenant_id); + } + } + FilterExpr::Not(child) => { + fold_filter_expr(&mut child.expr, catalog, database_id, tenant_id); + } + // Simple comparison, InList, Between, IsNull, IsNotNull — no sub-expressions to fold. + FilterExpr::Comparison { .. } + | FilterExpr::InList { .. } + | FilterExpr::Between { .. } + | FilterExpr::IsNull { .. } + | FilterExpr::IsNotNull { .. } => {} + } +} diff --git a/nodedb-sql/src/planner/catalog_fold/leaf.rs b/nodedb-sql/src/planner/catalog_fold/leaf.rs new file mode 100644 index 000000000..743657f1e --- /dev/null +++ b/nodedb-sql/src/planner/catalog_fold/leaf.rs @@ -0,0 +1,162 @@ +// SPDX-License-Identifier: Apache-2.0 + +//! Catalog folding within leaf read and write plans. + +use super::filter::fold_filter; +use crate::catalog::SqlCatalog; +use crate::planner::catalog_expr_fold::fold_expr; +use crate::planner::catalog_plan_shapes::{fold_projection, fold_sort_keys, fold_windows}; +use crate::types::{ + DocumentIndexLookupPlan, HybridSearchPlan, HybridSearchTriplePlan, RangeScanPlan, + RecursiveScanPlan, SqlExpr, SqlPlan, TextSearchPlan, VectorPrimaryDeletePlan, + VectorPrimaryUpdatePlan, +}; +use nodedb_types::DatabaseId; + +pub(super) fn fold_leaf( + plan: &mut SqlPlan, + catalog: &dyn SqlCatalog, + database_id: DatabaseId, + tenant_id: u64, +) { + match plan { + SqlPlan::PointGet { projection, .. } + | SqlPlan::RangeScan(RangeScanPlan { projection, .. }) => { + fold_projection(projection, catalog, database_id, tenant_id); + } + SqlPlan::DocumentIndexLookup(DocumentIndexLookupPlan { + filters, + projection, + sort_keys, + window_functions, + .. + }) => { + for filter in filters { + fold_filter(filter, catalog, database_id, tenant_id); + } + fold_projection(projection, catalog, database_id, tenant_id); + fold_sort_keys(sort_keys, catalog, database_id, tenant_id); + fold_windows(window_functions, catalog, database_id, tenant_id); + } + SqlPlan::Delete { filters, .. } + | SqlPlan::VectorPrimaryDelete(VectorPrimaryDeletePlan { filters, .. }) => { + for filter in filters { + fold_filter(filter, catalog, database_id, tenant_id); + } + } + SqlPlan::Update { + assignments, + filters, + .. + } + | SqlPlan::VectorPrimaryUpdate(VectorPrimaryUpdatePlan { + assignments, + filters, + .. + }) => { + for (_, expr) in assignments { + let owned = std::mem::replace(expr, SqlExpr::Wildcard); + *expr = fold_expr(owned, catalog, database_id, tenant_id); + } + for filter in filters { + fold_filter(filter, catalog, database_id, tenant_id); + } + } + SqlPlan::VectorSearch { + filters, + projection, + .. + } => { + for filter in filters { + fold_filter(filter, catalog, database_id, tenant_id); + } + fold_projection(projection, catalog, database_id, tenant_id); + } + SqlPlan::TextSearch(TextSearchPlan { + filters, + projection, + .. + }) + | SqlPlan::HybridSearch(HybridSearchPlan { + filters, + projection, + .. + }) + | SqlPlan::HybridSearchTriple(HybridSearchTriplePlan { + filters, + projection, + .. + }) => { + for filter in filters { + fold_filter(filter, catalog, database_id, tenant_id); + } + fold_projection(projection, catalog, database_id, tenant_id); + } + SqlPlan::SpatialScan { + attribute_filters, + projection, + .. + } => { + for filter in attribute_filters { + fold_filter(filter, catalog, database_id, tenant_id); + } + fold_projection(projection, catalog, database_id, tenant_id); + } + SqlPlan::RecursiveScan(RecursiveScanPlan { + base_filters, + recursive_filters, + projection, + .. + }) => { + for filter in base_filters { + fold_filter(filter, catalog, database_id, tenant_id); + } + for filter in recursive_filters { + fold_filter(filter, catalog, database_id, tenant_id); + } + fold_projection(projection, catalog, database_id, tenant_id); + } + SqlPlan::MultiVectorSearch { projection, .. } + | SqlPlan::SparseSearch { projection, .. } => { + fold_projection(projection, catalog, database_id, tenant_id); + } + // Composite plans: `walk_plan` folds them before reaching a leaf. + SqlPlan::Scan { .. } + | SqlPlan::Union { .. } + | SqlPlan::Intersect { .. } + | SqlPlan::Except { .. } + | SqlPlan::Cte(_) + | SqlPlan::Subquery { .. } + | SqlPlan::Join { .. } + | SqlPlan::UpdateFrom { .. } + | SqlPlan::InsertSelect { .. } + | SqlPlan::Aggregate { .. } + | SqlPlan::LateralTopK(_) + | SqlPlan::LateralLoop(_) + | SqlPlan::Merge(_) => {} + // Plans whose expressions this pass does not fold. + SqlPlan::ConstantResult { .. } + | SqlPlan::Insert(_) + | SqlPlan::KvInsert(_) + | SqlPlan::Upsert(_) + | SqlPlan::Truncate { .. } + | SqlPlan::TimeseriesScan(_) + | SqlPlan::TimeseriesIngest(_) + | SqlPlan::RecursiveValue(_) + | SqlPlan::CreateArray(_) + | SqlPlan::DropArray { .. } + | SqlPlan::AlterArray(_) + | SqlPlan::InsertArray(_) + | SqlPlan::DeleteArray(_) + | SqlPlan::ArraySlice(_) + | SqlPlan::ArrayProject(_) + | SqlPlan::ArrayAgg(_) + | SqlPlan::ArrayElementwise(_) + | SqlPlan::ArrayFlush { .. } + | SqlPlan::ArrayCompact { .. } + | SqlPlan::VectorPrimaryInsert(_) + | SqlPlan::VectorPrimaryTruncate(_) + | SqlPlan::CreateIndex(_) + | SqlPlan::DropIndex(_) => {} + } +} diff --git a/nodedb-sql/src/planner/catalog_fold/mod.rs b/nodedb-sql/src/planner/catalog_fold/mod.rs new file mode 100644 index 000000000..14c522749 --- /dev/null +++ b/nodedb-sql/src/planner/catalog_fold/mod.rs @@ -0,0 +1,9 @@ +// SPDX-License-Identifier: Apache-2.0 + +//! Plan-time catalog expression folding. + +mod filter; +mod leaf; +mod walk; + +pub use walk::fold_catalog_exprs_in_plan; diff --git a/nodedb-sql/src/planner/catalog_fold.rs b/nodedb-sql/src/planner/catalog_fold/walk.rs similarity index 61% rename from nodedb-sql/src/planner/catalog_fold.rs rename to nodedb-sql/src/planner/catalog_fold/walk.rs index 2ba9abb95..dd87de5c5 100644 --- a/nodedb-sql/src/planner/catalog_fold.rs +++ b/nodedb-sql/src/planner/catalog_fold/walk.rs @@ -9,14 +9,19 @@ //! the data-plane evaluator pure (no catalog/session context) while still //! supporting the `'name'::regclass` / `'name'::regtype` PostgreSQL idiom. +use crate::types::{CtePlan, LateralLoopPlan, LateralTopKPlan, MergePlan}; use nodedb_types::DatabaseId; +use super::filter::fold_filter; +use super::leaf::fold_leaf; use crate::catalog::SqlCatalog; -use crate::types::{Filter, FilterExpr, MergePlanAction, SqlExpr, SqlPlan}; +use crate::types::{MergePlanAction, SqlExpr, SqlPlan}; -use super::catalog_expr_fold::fold_expr; -use super::catalog_plan_shapes::{fold_aggregates, fold_projection, fold_sort_keys, fold_windows}; -use super::catalog_plan_validate::validate_catalog_exprs; +use crate::planner::catalog_expr_fold::fold_expr; +use crate::planner::catalog_plan_shapes::{ + fold_aggregates, fold_projection, fold_sort_keys, fold_windows, +}; +use crate::planner::catalog_plan_validate::validate_catalog_exprs; /// Walk every `Filter` in `plan` and fold catalog-dependent cast expressions /// to their constant OID equivalents. @@ -95,13 +100,13 @@ fn walk_plan( all, }, - SqlPlan::Cte { definitions, outer } => SqlPlan::Cte { + SqlPlan::Cte(CtePlan { definitions, outer }) => SqlPlan::Cte(CtePlan { definitions: definitions .into_iter() .map(|(name, plan)| (name, walk_plan(plan, catalog, database_id, tenant_id))) .collect(), outer: Box::new(walk_plan(*outer, catalog, database_id, tenant_id)), - }, + }), // The post-processing tail wraps a body that keeps its own filters. A // wrapper that stopped the walk left every catalog cast inside the body @@ -163,64 +168,6 @@ fn walk_plan( } } - mut plan @ (SqlPlan::PointGet { .. } | SqlPlan::RangeScan { .. }) => { - match &mut plan { - SqlPlan::PointGet { projection, .. } | SqlPlan::RangeScan { projection, .. } => { - fold_projection(projection, catalog, database_id, tenant_id); - } - _ => unreachable!(), - } - plan - } - - mut plan @ (SqlPlan::DocumentIndexLookup { .. } - | SqlPlan::Update { .. } - | SqlPlan::Delete { .. } - | SqlPlan::VectorPrimaryUpdate { .. } - | SqlPlan::VectorPrimaryDelete { .. }) => { - match &mut plan { - SqlPlan::DocumentIndexLookup { - filters, - projection, - sort_keys, - window_functions, - .. - } => { - for filter in filters { - fold_filter(filter, catalog, database_id, tenant_id); - } - fold_projection(projection, catalog, database_id, tenant_id); - fold_sort_keys(sort_keys, catalog, database_id, tenant_id); - fold_windows(window_functions, catalog, database_id, tenant_id); - } - SqlPlan::Delete { filters, .. } | SqlPlan::VectorPrimaryDelete { filters, .. } => { - for filter in filters { - fold_filter(filter, catalog, database_id, tenant_id); - } - } - SqlPlan::Update { - assignments, - filters, - .. - } - | SqlPlan::VectorPrimaryUpdate { - assignments, - filters, - .. - } => { - for (_, expr) in assignments { - let owned = std::mem::replace(expr, SqlExpr::Wildcard); - *expr = fold_expr(owned, catalog, database_id, tenant_id); - } - for filter in filters { - fold_filter(filter, catalog, database_id, tenant_id); - } - } - _ => unreachable!(), - } - plan - } - SqlPlan::UpdateFrom { collection, engine, @@ -295,7 +242,7 @@ fn walk_plan( } } - SqlPlan::LateralTopK { + SqlPlan::LateralTopK(LateralTopKPlan { outer, outer_alias, inner_collection, @@ -306,13 +253,13 @@ fn walk_plan( lateral_alias, mut projection, left_join, - } => { + }) => { for filter in &mut inner_filters { fold_filter(filter, catalog, database_id, tenant_id); } fold_sort_keys(&mut inner_order_by, catalog, database_id, tenant_id); fold_projection(&mut projection, catalog, database_id, tenant_id); - SqlPlan::LateralTopK { + SqlPlan::LateralTopK(LateralTopKPlan { outer: Box::new(walk_plan(*outer, catalog, database_id, tenant_id)), outer_alias, inner_collection, @@ -323,10 +270,10 @@ fn walk_plan( lateral_alias, projection, left_join, - } + }) } - SqlPlan::LateralLoop { + SqlPlan::LateralLoop(LateralLoopPlan { outer, outer_alias, inner, @@ -335,9 +282,9 @@ fn walk_plan( mut projection, outer_row_cap, left_join, - } => { + }) => { fold_projection(&mut projection, catalog, database_id, tenant_id); - SqlPlan::LateralLoop { + SqlPlan::LateralLoop(LateralLoopPlan { outer: Box::new(walk_plan(*outer, catalog, database_id, tenant_id)), outer_alias, inner: Box::new(walk_plan(*inner, catalog, database_id, tenant_id)), @@ -346,10 +293,10 @@ fn walk_plan( projection, outer_row_cap, left_join, - } + }) } - SqlPlan::Merge { + SqlPlan::Merge(MergePlan { target, engine, source, @@ -358,7 +305,7 @@ fn walk_plan( source_alias, mut clauses, returning, - } => { + }) => { for clause in &mut clauses { for filter in &mut clause.extra_predicate { fold_filter(filter, catalog, database_id, tenant_id); @@ -379,7 +326,7 @@ fn walk_plan( MergePlanAction::Delete | MergePlanAction::DoNothing => {} } } - SqlPlan::Merge { + SqlPlan::Merge(MergePlan { target, engine, source: Box::new(walk_plan(*source, catalog, database_id, tenant_id)), @@ -388,117 +335,13 @@ fn walk_plan( source_alias, clauses, returning, - } - } - - mut plan @ (SqlPlan::VectorSearch { .. } - | SqlPlan::TextSearch { .. } - | SqlPlan::SpatialScan { .. } - | SqlPlan::RecursiveScan { .. }) => { - match &mut plan { - SqlPlan::VectorSearch { - filters, - projection, - .. - } => { - for filter in filters { - fold_filter(filter, catalog, database_id, tenant_id); - } - fold_projection(projection, catalog, database_id, tenant_id); - } - SqlPlan::TextSearch { - filters, - projection, - .. - } => { - for filter in filters { - fold_filter(filter, catalog, database_id, tenant_id); - } - fold_projection(projection, catalog, database_id, tenant_id); - } - SqlPlan::SpatialScan { - attribute_filters, - projection, - .. - } => { - for filter in attribute_filters { - fold_filter(filter, catalog, database_id, tenant_id); - } - fold_projection(projection, catalog, database_id, tenant_id); - } - SqlPlan::RecursiveScan { - base_filters, - recursive_filters, - projection, - .. - } => { - for filter in base_filters { - fold_filter(filter, catalog, database_id, tenant_id); - } - for filter in recursive_filters { - fold_filter(filter, catalog, database_id, tenant_id); - } - fold_projection(projection, catalog, database_id, tenant_id); - } - _ => unreachable!(), - } - plan - } - - mut plan @ (SqlPlan::MultiVectorSearch { .. } - | SqlPlan::SparseSearch { .. } - | SqlPlan::HybridSearch { .. } - | SqlPlan::HybridSearchTriple { .. }) => { - match &mut plan { - SqlPlan::MultiVectorSearch { projection, .. } - | SqlPlan::SparseSearch { projection, .. } - | SqlPlan::HybridSearch { projection, .. } - | SqlPlan::HybridSearchTriple { projection, .. } => { - fold_projection(projection, catalog, database_id, tenant_id); - } - _ => unreachable!(), - } - plan + }) } - // Plan variants without expression-bearing filters pass through unchanged. - other => other, - } -} - -fn fold_filter( - filter: &mut Filter, - catalog: &dyn SqlCatalog, - database_id: DatabaseId, - tenant_id: u64, -) { - fold_filter_expr(&mut filter.expr, catalog, database_id, tenant_id); -} - -fn fold_filter_expr( - expr: &mut FilterExpr, - catalog: &dyn SqlCatalog, - database_id: DatabaseId, - tenant_id: u64, -) { - match expr { - FilterExpr::Expr(sql_expr) => { - let owned = std::mem::replace(sql_expr, SqlExpr::Wildcard); - *sql_expr = fold_expr(owned, catalog, database_id, tenant_id); - } - FilterExpr::And(children) | FilterExpr::Or(children) => { - for child in children { - fold_filter_expr(&mut child.expr, catalog, database_id, tenant_id); - } - } - FilterExpr::Not(child) => { - fold_filter_expr(&mut child.expr, catalog, database_id, tenant_id); + // Leaf plans: `fold_leaf` folds the expressions each one carries. + mut other => { + fold_leaf(&mut other, catalog, database_id, tenant_id); + other } - // Simple comparison, InList, Between, IsNull, IsNotNull — no sub-expressions to fold. - FilterExpr::Comparison { .. } - | FilterExpr::InList { .. } - | FilterExpr::Between { .. } - | FilterExpr::IsNull { .. } - | FilterExpr::IsNotNull { .. } => {} } } diff --git a/nodedb-sql/src/planner/catalog_plan_validate.rs b/nodedb-sql/src/planner/catalog_plan_validate.rs index 6230bd833..3667ae690 100644 --- a/nodedb-sql/src/planner/catalog_plan_validate.rs +++ b/nodedb-sql/src/planner/catalog_plan_validate.rs @@ -2,6 +2,11 @@ //! Catalog-dependent expression validation across nested SQL plan shapes. +use crate::types::{ + CtePlan, DocumentIndexLookupPlan, HybridSearchPlan, HybridSearchTriplePlan, LateralLoopPlan, + LateralTopKPlan, MergePlan, RangeScanPlan, RecursiveScanPlan, TextSearchPlan, + VectorPrimaryDeletePlan, VectorPrimaryUpdatePlan, +}; use nodedb_types::DatabaseId; use crate::catalog::SqlCatalog; @@ -26,19 +31,20 @@ pub(super) fn validate_catalog_exprs( window_functions, .. } - | SqlPlan::DocumentIndexLookup { + | SqlPlan::DocumentIndexLookup(DocumentIndexLookupPlan { filters, projection, sort_keys, window_functions, .. - } => { + }) => { validate_filters(filters, catalog, database_id, tenant_id)?; validate_projection(projection, catalog, database_id, tenant_id)?; validate_sort_keys(sort_keys, catalog, database_id, tenant_id)?; validate_windows(window_functions, catalog, database_id, tenant_id)?; } - SqlPlan::PointGet { projection, .. } | SqlPlan::RangeScan { projection, .. } => { + SqlPlan::PointGet { projection, .. } + | SqlPlan::RangeScan(RangeScanPlan { projection, .. }) => { validate_projection(projection, catalog, database_id, tenant_id)?; } SqlPlan::Update { @@ -46,17 +52,18 @@ pub(super) fn validate_catalog_exprs( filters, .. } - | SqlPlan::VectorPrimaryUpdate { + | SqlPlan::VectorPrimaryUpdate(VectorPrimaryUpdatePlan { assignments, filters, .. - } => { + }) => { for (_, expr) in assignments { validate_expr(expr, catalog, database_id, tenant_id)?; } validate_filters(filters, catalog, database_id, tenant_id)?; } - SqlPlan::Delete { filters, .. } | SqlPlan::VectorPrimaryDelete { filters, .. } => { + SqlPlan::Delete { filters, .. } + | SqlPlan::VectorPrimaryDelete(VectorPrimaryDeletePlan { filters, .. }) => { validate_filters(filters, catalog, database_id, tenant_id)? } SqlPlan::UpdateFrom { @@ -99,7 +106,7 @@ pub(super) fn validate_catalog_exprs( validate_catalog_exprs(left, catalog, database_id, tenant_id)?; validate_catalog_exprs(right, catalog, database_id, tenant_id)?; } - SqlPlan::Cte { definitions, outer } => { + SqlPlan::Cte(CtePlan { definitions, outer }) => { for (_, definition) in definitions { validate_catalog_exprs(definition, catalog, database_id, tenant_id)?; } @@ -137,31 +144,31 @@ pub(super) fn validate_catalog_exprs( validate_projection(projection, catalog, database_id, tenant_id)?; validate_filters(filters, catalog, database_id, tenant_id)?; } - SqlPlan::LateralTopK { + SqlPlan::LateralTopK(LateralTopKPlan { outer, inner_filters, inner_order_by, projection, .. - } => { + }) => { validate_catalog_exprs(outer, catalog, database_id, tenant_id)?; validate_filters(inner_filters, catalog, database_id, tenant_id)?; validate_sort_keys(inner_order_by, catalog, database_id, tenant_id)?; validate_projection(projection, catalog, database_id, tenant_id)?; } - SqlPlan::LateralLoop { + SqlPlan::LateralLoop(LateralLoopPlan { outer, inner, projection, .. - } => { + }) => { validate_catalog_exprs(outer, catalog, database_id, tenant_id)?; validate_catalog_exprs(inner, catalog, database_id, tenant_id)?; validate_projection(projection, catalog, database_id, tenant_id)?; } - SqlPlan::Merge { + SqlPlan::Merge(MergePlan { source, clauses, .. - } => { + }) => { validate_catalog_exprs(source, catalog, database_id, tenant_id)?; for clause in clauses { validate_filters(&clause.extra_predicate, catalog, database_id, tenant_id)?; @@ -189,16 +196,27 @@ pub(super) fn validate_catalog_exprs( validate_projection(projection, catalog, database_id, tenant_id)?; } SqlPlan::MultiVectorSearch { projection, .. } - | SqlPlan::SparseSearch { projection, .. } - | SqlPlan::HybridSearch { projection, .. } - | SqlPlan::HybridSearchTriple { projection, .. } => { + | SqlPlan::SparseSearch { projection, .. } => { validate_projection(projection, catalog, database_id, tenant_id)?; } - SqlPlan::TextSearch { + SqlPlan::HybridSearch(HybridSearchPlan { filters, projection, .. - } => { + }) + | SqlPlan::HybridSearchTriple(HybridSearchTriplePlan { + filters, + projection, + .. + }) => { + validate_filters(filters, catalog, database_id, tenant_id)?; + validate_projection(projection, catalog, database_id, tenant_id)?; + } + SqlPlan::TextSearch(TextSearchPlan { + filters, + projection, + .. + }) => { validate_filters(filters, catalog, database_id, tenant_id)?; validate_projection(projection, catalog, database_id, tenant_id)?; } @@ -210,17 +228,39 @@ pub(super) fn validate_catalog_exprs( validate_filters(attribute_filters, catalog, database_id, tenant_id)?; validate_projection(projection, catalog, database_id, tenant_id)?; } - SqlPlan::RecursiveScan { + SqlPlan::RecursiveScan(RecursiveScanPlan { base_filters, recursive_filters, projection, .. - } => { + }) => { validate_filters(base_filters, catalog, database_id, tenant_id)?; validate_filters(recursive_filters, catalog, database_id, tenant_id)?; validate_projection(projection, catalog, database_id, tenant_id)?; } - _ => {} + SqlPlan::ConstantResult { .. } + | SqlPlan::Insert(_) + | SqlPlan::KvInsert(_) + | SqlPlan::Upsert(_) + | SqlPlan::Truncate { .. } + | SqlPlan::TimeseriesScan(_) + | SqlPlan::TimeseriesIngest(_) + | SqlPlan::RecursiveValue(_) + | SqlPlan::CreateArray(_) + | SqlPlan::DropArray { .. } + | SqlPlan::AlterArray(_) + | SqlPlan::InsertArray(_) + | SqlPlan::DeleteArray(_) + | SqlPlan::ArraySlice(_) + | SqlPlan::ArrayProject(_) + | SqlPlan::ArrayAgg(_) + | SqlPlan::ArrayElementwise(_) + | SqlPlan::ArrayFlush { .. } + | SqlPlan::ArrayCompact { .. } + | SqlPlan::VectorPrimaryInsert(_) + | SqlPlan::VectorPrimaryTruncate(_) + | SqlPlan::CreateIndex(_) + | SqlPlan::DropIndex(_) => {} } Ok(()) } diff --git a/nodedb-sql/src/planner/cte/recursive_scan.rs b/nodedb-sql/src/planner/cte/recursive_scan.rs index 76e70ac1f..a0d8fd553 100644 --- a/nodedb-sql/src/planner/cte/recursive_scan.rs +++ b/nodedb-sql/src/planner/cte/recursive_scan.rs @@ -211,7 +211,7 @@ fn plan_recursive_scan_from_parts( _ => Vec::new(), }; - Ok(SqlPlan::RecursiveScan { + Ok(SqlPlan::RecursiveScan(RecursiveScanPlan { collection, base_filters: extract_filters(base), recursive_filters, @@ -220,7 +220,7 @@ fn plan_recursive_scan_from_parts( distinct: *distinct, limit: 10000, projection, - }) + })) } pub(super) fn plan_cte_branch( diff --git a/nodedb-sql/src/planner/cte/recursive_value.rs b/nodedb-sql/src/planner/cte/recursive_value.rs index 9b4fd9bf9..41022656a 100644 --- a/nodedb-sql/src/planner/cte/recursive_value.rs +++ b/nodedb-sql/src/planner/cte/recursive_value.rs @@ -67,7 +67,7 @@ pub(super) fn plan_recursive_value( } } - Ok(SqlPlan::RecursiveValue { + Ok(SqlPlan::RecursiveValue(RecursiveValuePlan { cte_name: cte_name.to_owned(), columns, init_exprs, @@ -75,7 +75,7 @@ pub(super) fn plan_recursive_value( condition, max_depth: DEFAULT_MAX_RECURSION_DEPTH, distinct, - }) + })) } /// One anchor projection item: the expression as SQL text, plus the output diff --git a/nodedb-sql/src/planner/dml/insert.rs b/nodedb-sql/src/planner/dml/insert.rs index 57a54166c..5702178ce 100644 --- a/nodedb-sql/src/planner/dml/insert.rs +++ b/nodedb-sql/src/planner/dml/insert.rs @@ -215,11 +215,11 @@ mod tests { ], )); let plan = plan("INSERT INTO t (id) VALUES ('r1')", &catalog).expect("plans"); - let SqlPlan::Insert { + let SqlPlan::Insert(InsertPlan { rows, volatile_defaults, .. - } = plan + }) = plan else { panic!("expected SqlPlan::Insert, got {plan:?}"); }; @@ -260,11 +260,11 @@ mod tests { ], )); let plan = plan("INSERT INTO t (v) VALUES (1.0)", &catalog).expect("plans"); - let SqlPlan::TimeseriesIngest { + let SqlPlan::TimeseriesIngest(TimeseriesIngestPlan { rows, volatile_defaults, .. - } = plan + }) = plan else { panic!("expected SqlPlan::TimeseriesIngest, got {plan:?}"); }; diff --git a/nodedb-sql/src/planner/dml_helpers/kv_counter.rs b/nodedb-sql/src/planner/dml_helpers/kv_counter.rs index 9498c7be6..078211bfe 100644 --- a/nodedb-sql/src/planner/dml_helpers/kv_counter.rs +++ b/nodedb-sql/src/planner/dml_helpers/kv_counter.rs @@ -103,7 +103,7 @@ pub fn plan_kv_counter_fresh_row( catalog, })?; let cells = match plans.into_iter().next() { - Some(SqlPlan::KvInsert { mut entries, .. }) if entries.len() == 1 => { + Some(SqlPlan::KvInsert(KvInsertPlan { mut entries, .. })) if entries.len() == 1 => { let (_, cells) = entries.remove(0); cells .into_iter() diff --git a/nodedb-sql/src/planner/dml_helpers/kv_insert.rs b/nodedb-sql/src/planner/dml_helpers/kv_insert.rs index 68525971e..e3276e3af 100644 --- a/nodedb-sql/src/planner/dml_helpers/kv_insert.rs +++ b/nodedb-sql/src/planner/dml_helpers/kv_insert.rs @@ -130,14 +130,14 @@ pub(crate) fn build_kv_insert_plan(params: KvInsertParams<'_>) -> Result, target_keys: Vec, ) -> Vec { - vec![SqlPlan::VectorPrimaryDelete { + vec![SqlPlan::VectorPrimaryDelete(VectorPrimaryDeletePlan { collection: collection.to_string(), field: vpc.vector_field.clone(), filters, target_keys, primary_key: info.primary_key.clone(), - }] + })] } /// Build a `SqlPlan::VectorPrimaryTruncate`. @@ -115,11 +117,11 @@ pub(crate) fn build_vector_primary_truncate_plan( vpc: &nodedb_types::VectorPrimaryConfig, restart_identity: bool, ) -> SqlPlan { - SqlPlan::VectorPrimaryTruncate { + SqlPlan::VectorPrimaryTruncate(VectorPrimaryTruncatePlan { collection: collection.to_string(), field: vpc.vector_field.clone(), restart_identity, - } + }) } /// The `f32` components of an assigned vector expression. diff --git a/nodedb-sql/src/planner/dml_helpers/vector_primary_insert.rs b/nodedb-sql/src/planner/dml_helpers/vector_primary_insert.rs index f5df1f448..1cec8cf29 100644 --- a/nodedb-sql/src/planner/dml_helpers/vector_primary_insert.rs +++ b/nodedb-sql/src/planner/dml_helpers/vector_primary_insert.rs @@ -100,16 +100,18 @@ pub(crate) fn build_vector_primary_insert_plan( }); } - Ok(vec![SqlPlan::VectorPrimaryInsert { - collection: collection.to_string(), - field: vpc.vector_field.clone(), - quantization: vpc.quantization, - storage_dtype: vpc.storage_dtype, - payload_indexes: vpc.payload_indexes.clone(), - rows: result_rows, - volatile_defaults, - intent, - on_conflict_updates, - primary_key, - }]) + Ok(vec![SqlPlan::VectorPrimaryInsert( + VectorPrimaryInsertPlan { + collection: collection.to_string(), + field: vpc.vector_field.clone(), + quantization: vpc.quantization, + storage_dtype: vpc.storage_dtype, + payload_indexes: vpc.payload_indexes.clone(), + rows: result_rows, + volatile_defaults, + intent, + on_conflict_updates, + primary_key, + }, + )]) } diff --git a/nodedb-sql/src/planner/index_ddl/create.rs b/nodedb-sql/src/planner/index_ddl/create.rs index f7c206fea..87d5e37c0 100644 --- a/nodedb-sql/src/planner/index_ddl/create.rs +++ b/nodedb-sql/src/planner/index_ddl/create.rs @@ -2,6 +2,7 @@ //! Plan a `CREATE [UNIQUE] INDEX` statement parsed by sqlparser. +use crate::types::CreateIndexPlan; use sqlparser::ast; use crate::SqlPlan; @@ -57,14 +58,14 @@ pub fn plan_create_index(ci: &ast::CreateIndex) -> Result { }); } - Ok(SqlPlan::CreateIndex { + Ok(SqlPlan::CreateIndex(CreateIndexPlan { index_name, collection, field, unique: ci.unique, if_not_exists: ci.if_not_exists, case_insensitive, - }) + })) } /// Strip a `COLLATE ` wrapper from `expr` and return the inner @@ -106,14 +107,14 @@ mod tests { #[test] fn basic_index() { - let SqlPlan::CreateIndex { + let SqlPlan::CreateIndex(CreateIndexPlan { index_name, collection, field, unique, if_not_exists, case_insensitive, - } = plan("CREATE INDEX idx_users_email ON users (email)").unwrap() + }) = plan("CREATE INDEX idx_users_email ON users (email)").unwrap() else { panic!("expected CreateIndex"); }; @@ -127,7 +128,7 @@ mod tests { #[test] fn anonymous_index_name() { - let SqlPlan::CreateIndex { index_name, .. } = + let SqlPlan::CreateIndex(CreateIndexPlan { index_name, .. }) = plan("CREATE INDEX ON users (email)").unwrap() else { panic!("expected CreateIndex"); @@ -137,11 +138,11 @@ mod tests { #[test] fn unique_and_if_not_exists() { - let SqlPlan::CreateIndex { + let SqlPlan::CreateIndex(CreateIndexPlan { unique, if_not_exists, .. - } = plan("CREATE UNIQUE INDEX IF NOT EXISTS u ON users (email)").unwrap() + }) = plan("CREATE UNIQUE INDEX IF NOT EXISTS u ON users (email)").unwrap() else { panic!("expected CreateIndex"); }; @@ -157,9 +158,9 @@ mod tests { "CREATE INDEX i ON users (email COLLATE ci)", "CREATE INDEX i ON users (email COLLATE case_insensitive)", ] { - let SqlPlan::CreateIndex { + let SqlPlan::CreateIndex(CreateIndexPlan { case_insensitive, .. - } = plan(sql).unwrap() + }) = plan(sql).unwrap() else { panic!("expected CreateIndex for {sql}"); }; @@ -169,9 +170,9 @@ mod tests { #[test] fn collate_other_not_case_insensitive() { - let SqlPlan::CreateIndex { + let SqlPlan::CreateIndex(CreateIndexPlan { case_insensitive, .. - } = plan("CREATE INDEX i ON users (email COLLATE \"en_US\")").unwrap() + }) = plan("CREATE INDEX i ON users (email COLLATE \"en_US\")").unwrap() else { panic!("expected CreateIndex"); }; diff --git a/nodedb-sql/src/planner/index_ddl/drop.rs b/nodedb-sql/src/planner/index_ddl/drop.rs index a87ab1231..260a546ed 100644 --- a/nodedb-sql/src/planner/index_ddl/drop.rs +++ b/nodedb-sql/src/planner/index_ddl/drop.rs @@ -2,6 +2,7 @@ //! Plan a `DROP INDEX` statement parsed by sqlparser. +use crate::types::DropIndexPlan; use sqlparser::ast; use crate::SqlPlan; @@ -42,11 +43,11 @@ pub fn plan_drop_index(stmt: &ast::Statement) -> Result { None => None, }; - Ok(SqlPlan::DropIndex { + Ok(SqlPlan::DropIndex(DropIndexPlan { index_name, collection, if_exists: *if_exists, - }) + })) } #[cfg(test)] @@ -61,11 +62,11 @@ mod tests { #[test] fn basic_drop() { - let SqlPlan::DropIndex { + let SqlPlan::DropIndex(DropIndexPlan { index_name, collection, if_exists, - } = plan("DROP INDEX idx_users_email").unwrap() + }) = plan("DROP INDEX idx_users_email").unwrap() else { panic!("expected DropIndex"); }; @@ -76,7 +77,9 @@ mod tests { #[test] fn if_exists_honored() { - let SqlPlan::DropIndex { if_exists, .. } = plan("DROP INDEX IF EXISTS idx").unwrap() else { + let SqlPlan::DropIndex(DropIndexPlan { if_exists, .. }) = + plan("DROP INDEX IF EXISTS idx").unwrap() + else { panic!("expected DropIndex"); }; assert!(if_exists); diff --git a/nodedb-sql/src/planner/lateral/plan.rs b/nodedb-sql/src/planner/lateral/plan.rs index f4ab91a74..d80da0650 100644 --- a/nodedb-sql/src/planner/lateral/plan.rs +++ b/nodedb-sql/src/planner/lateral/plan.rs @@ -91,16 +91,14 @@ pub fn plan_lateral_join(args: LateralJoinArgs<'_>) -> Result { let has_equi = !analysis.equi_keys.is_empty(); let inner_limit = limit_from_query(subquery)?; reject_lateral_offset(subquery)?; - let is_top_k = has_equi && inner_limit.is_some() && analysis.non_equi.is_empty(); - - if is_top_k { + if let Some(inner_limit) = inner_limit.filter(|_| has_equi && analysis.non_equi.is_empty()) { plan_lateral_top_k(LateralTopKPlanArgs { outer_plan, outer_alias, select, subquery, equi_keys: analysis.equi_keys, - inner_limit: inner_limit.expect("checked above"), + inner_limit, lateral_alias, left_join, outer_projection, @@ -170,7 +168,7 @@ pub fn plan_lateral_join(args: LateralJoinArgs<'_>) -> Result { .iter() .map(|c| (c.inner_col.clone(), c.outer_col.clone())) .collect(); - Ok(SqlPlan::LateralLoop { + Ok(SqlPlan::LateralLoop(LateralLoopPlan { outer: Box::new(outer_plan), outer_alias, inner: Box::new(inner_plan), @@ -179,7 +177,7 @@ pub fn plan_lateral_join(args: LateralJoinArgs<'_>) -> Result { projection: outer_projection, outer_row_cap: LATERAL_LOOP_CAP, left_join, - }) + })) } } @@ -327,7 +325,7 @@ fn plan_lateral_top_k(args: LateralTopKPlanArgs<'_>) -> Result { .map(|c| (c.outer_col, c.inner_col)) .collect(); - Ok(SqlPlan::LateralTopK { + Ok(SqlPlan::LateralTopK(LateralTopKPlan { outer: Box::new(outer_plan), outer_alias, inner_collection, @@ -338,5 +336,5 @@ fn plan_lateral_top_k(args: LateralTopKPlanArgs<'_>) -> Result { lateral_alias: lateral_alias.to_string(), projection: outer_projection, left_join, - }) + })) } diff --git a/nodedb-sql/src/planner/select/derived_from.rs b/nodedb-sql/src/planner/select/derived_from.rs index 48a704569..941bc43ef 100644 --- a/nodedb-sql/src/planner/select/derived_from.rs +++ b/nodedb-sql/src/planner/select/derived_from.rs @@ -101,10 +101,10 @@ pub(in crate::planner::select) fn try_plan_derived_from( )?; Ok(Some(PlannedSelect { - plan: SqlPlan::Cte { + plan: SqlPlan::Cte(CtePlan { definitions: vec![(alias_name, inner_plan)], outer: Box::new(outer.plan), - }, + }), scope: outer.scope, })) } diff --git a/nodedb-sql/src/planner/select/order_by/vector_join.rs b/nodedb-sql/src/planner/select/order_by/vector_join.rs index f1a9f3fec..6a2008415 100644 --- a/nodedb-sql/src/planner/select/order_by/vector_join.rs +++ b/nodedb-sql/src/planner/select/order_by/vector_join.rs @@ -7,6 +7,7 @@ //! from the array slice. This module isolates the AST inspection used by //! the trigger detector to recognise the fusion-eligible join shape. +use crate::types::ArraySlicePlan; use crate::types::SqlPlan; /// Result of inspecting a `SqlPlan::Join` for the @@ -27,24 +28,26 @@ pub(super) fn extract_vector_join_target( right: &SqlPlan, ) -> Option { match (left, right) { - (SqlPlan::Scan { collection, .. }, SqlPlan::ArraySlice { name, slice, .. }) => { - Some(VectorJoinTarget { - vector_collection: collection.clone(), - array_prefilter: Some(crate::types::ArrayPrefilter { - array_name: name.clone(), - slice: slice.clone(), - }), - }) - } - (SqlPlan::ArraySlice { name, slice, .. }, SqlPlan::Scan { collection, .. }) => { - Some(VectorJoinTarget { - vector_collection: collection.clone(), - array_prefilter: Some(crate::types::ArrayPrefilter { - array_name: name.clone(), - slice: slice.clone(), - }), - }) - } + ( + SqlPlan::Scan { collection, .. }, + SqlPlan::ArraySlice(ArraySlicePlan { name, slice, .. }), + ) => Some(VectorJoinTarget { + vector_collection: collection.clone(), + array_prefilter: Some(crate::types::ArrayPrefilter { + array_name: name.clone(), + slice: slice.clone(), + }), + }), + ( + SqlPlan::ArraySlice(ArraySlicePlan { name, slice, .. }), + SqlPlan::Scan { collection, .. }, + ) => Some(VectorJoinTarget { + vector_collection: collection.clone(), + array_prefilter: Some(crate::types::ArrayPrefilter { + array_name: name.clone(), + slice: slice.clone(), + }), + }), _ => None, } } diff --git a/nodedb-sql/src/planner/select/post_process.rs b/nodedb-sql/src/planner/select/post_process.rs index 662d440e6..77467482e 100644 --- a/nodedb-sql/src/planner/select/post_process.rs +++ b/nodedb-sql/src/planner/select/post_process.rs @@ -17,6 +17,11 @@ use crate::error::{Result, SqlError}; use crate::planner::qualified_name; +use crate::types::{ + ArrayProjectPlan, ArraySlicePlan, DocumentIndexLookupPlan, HybridSearchPlan, + HybridSearchTriplePlan, LateralLoopPlan, LateralTopKPlan, RangeScanPlan, RecursiveScanPlan, + TextSearchPlan, TimeseriesScanPlan, +}; use crate::types::{Projection, SortKey, SqlExpr, SqlPlan}; /// Wrap `input` in a post-processing node that applies `sort_keys`, `offset` @@ -59,12 +64,12 @@ fn retain_sort_columns(input: &mut SqlPlan, sort_keys: &[SortKey]) -> Result Option<&mut Vec> { SqlPlan::Scan { projection, .. } | SqlPlan::Join { projection, .. } | SqlPlan::PointGet { projection, .. } - | SqlPlan::RangeScan { projection, .. } - | SqlPlan::DocumentIndexLookup { projection, .. } - | SqlPlan::TimeseriesScan { projection, .. } + | SqlPlan::RangeScan(RangeScanPlan { projection, .. }) + | SqlPlan::DocumentIndexLookup(DocumentIndexLookupPlan { projection, .. }) + | SqlPlan::TimeseriesScan(TimeseriesScanPlan { projection, .. }) | SqlPlan::VectorSearch { projection, .. } | SqlPlan::MultiVectorSearch { projection, .. } | SqlPlan::SparseSearch { projection, .. } - | SqlPlan::TextSearch { projection, .. } - | SqlPlan::HybridSearch { projection, .. } - | SqlPlan::HybridSearchTriple { projection, .. } + | SqlPlan::TextSearch(TextSearchPlan { projection, .. }) + | SqlPlan::HybridSearch(HybridSearchPlan { projection, .. }) + | SqlPlan::HybridSearchTriple(HybridSearchTriplePlan { projection, .. }) | SqlPlan::SpatialScan { projection, .. } - | SqlPlan::RecursiveScan { projection, .. } - | SqlPlan::LateralTopK { projection, .. } - | SqlPlan::LateralLoop { projection, .. } + | SqlPlan::RecursiveScan(RecursiveScanPlan { projection, .. }) + | SqlPlan::LateralTopK(LateralTopKPlan { projection, .. }) + | SqlPlan::LateralLoop(LateralLoopPlan { projection, .. }) | SqlPlan::Subquery { projection, .. } => Some(projection), _ => None, } diff --git a/nodedb-sql/src/types/mod.rs b/nodedb-sql/src/types/mod.rs index aa8f578db..22e24a072 100644 --- a/nodedb-sql/src/types/mod.rs +++ b/nodedb-sql/src/types/mod.rs @@ -12,6 +12,15 @@ pub mod query; pub use collection::{CollectionInfo, ColumnInfo, IndexSpec, IndexState}; pub use filter::{CompareOp, Filter, FilterExpr}; +pub use plan::{ + AlterArrayPlan, ArrayAggPlan, ArrayElementwisePlan, ArrayProjectPlan, ArraySlicePlan, + CreateArrayPlan, CreateIndexPlan, CtePlan, DeleteArrayPlan, DocumentIndexLookupPlan, + DropIndexPlan, HybridSearchPlan, HybridSearchTriplePlan, InsertArrayPlan, InsertPlan, + KvInsertPlan, LateralLoopPlan, LateralTopKPlan, MergePlan, RangeScanPlan, RecursiveScanPlan, + RecursiveValuePlan, TextScoreColumn, TextSearchPlan, TextSearchShape, TimeseriesIngestPlan, + TimeseriesScanPlan, UpsertPlan, VectorPrimaryDeletePlan, VectorPrimaryInsertPlan, + VectorPrimaryTruncatePlan, VectorPrimaryUpdatePlan, +}; pub use plan::{ ArrayPrefilter, DistanceMetric, KvInsertIntent, MergeClauseKind, MergePlanAction, MergePlanClause, PlanCacheEligibility, SqlPlan, VectorAnnOptions, VectorPrimaryInsertIntent, diff --git a/nodedb-sql/src/types/plan/cacheability.rs b/nodedb-sql/src/types/plan/cacheability.rs index 5c8717366..077943825 100644 --- a/nodedb-sql/src/types/plan/cacheability.rs +++ b/nodedb-sql/src/types/plan/cacheability.rs @@ -4,6 +4,12 @@ use crate::types::query::EngineType; use super::SqlPlan; use super::expr_scan::projection_is_cp_computed; +use super::variants::{ + CtePlan, DocumentIndexLookupPlan, HybridSearchPlan, HybridSearchTriplePlan, InsertPlan, + KvInsertPlan, LateralLoopPlan, LateralTopKPlan, MergePlan, RangeScanPlan, RecursiveScanPlan, + TextSearchPlan, TimeseriesIngestPlan, TimeseriesScanPlan, UpsertPlan, VectorPrimaryDeletePlan, + VectorPrimaryInsertPlan, VectorPrimaryUpdatePlan, +}; /// Whether a logical plan may be lowered once and reused from the physical-plan cache. #[derive(Debug, Clone, Copy, PartialEq, Eq)] @@ -50,45 +56,45 @@ impl SqlPlan { // must not replay one execution's rows for another. Self::Scan { projection, .. } | Self::PointGet { projection, .. } - | Self::DocumentIndexLookup { projection, .. } - | Self::RangeScan { projection, .. } + | Self::DocumentIndexLookup(DocumentIndexLookupPlan { projection, .. }) + | Self::RangeScan(RangeScanPlan { projection, .. }) | Self::Join { projection, .. } - | Self::TimeseriesScan { projection, .. } + | Self::TimeseriesScan(TimeseriesScanPlan { projection, .. }) | Self::VectorSearch { projection, .. } | Self::MultiVectorSearch { projection, .. } | Self::SparseSearch { projection, .. } - | Self::TextSearch { projection, .. } - | Self::HybridSearch { projection, .. } - | Self::HybridSearchTriple { projection, .. } + | Self::TextSearch(TextSearchPlan { projection, .. }) + | Self::HybridSearch(HybridSearchPlan { projection, .. }) + | Self::HybridSearchTriple(HybridSearchTriplePlan { projection, .. }) | Self::SpatialScan { projection, .. } - | Self::RecursiveScan { projection, .. } + | Self::RecursiveScan(RecursiveScanPlan { projection, .. }) | Self::Subquery { projection, .. } - | Self::LateralTopK { projection, .. } - | Self::LateralLoop { projection, .. } + | Self::LateralTopK(LateralTopKPlan { projection, .. }) + | Self::LateralLoop(LateralLoopPlan { projection, .. }) if projection_is_cp_computed(projection) => { DataDependent } - Self::Insert { + Self::Insert(InsertPlan { volatile_defaults: true, .. - } - | Self::Upsert { + }) + | Self::Upsert(UpsertPlan { volatile_defaults: true, .. - } - | Self::TimeseriesIngest { + }) + | Self::TimeseriesIngest(TimeseriesIngestPlan { volatile_defaults: true, .. - } - | Self::KvInsert { + }) + | Self::KvInsert(KvInsertPlan { volatile_defaults: true, .. - } - | Self::VectorPrimaryInsert { + }) + | Self::VectorPrimaryInsert(VectorPrimaryInsertPlan { volatile_defaults: true, .. - } => DataDependent, + }) => DataDependent, Self::PointGet { engine: EngineType::DocumentSchemaless | EngineType::DocumentStrict, .. @@ -112,8 +118,8 @@ impl SqlPlan { } // A point-key delete or update binds its surrogates while the plan // is lowered, from the catalog state of that moment. - Self::VectorPrimaryDelete { target_keys, .. } - | Self::VectorPrimaryUpdate { target_keys, .. } + Self::VectorPrimaryDelete(VectorPrimaryDeletePlan { target_keys, .. }) + | Self::VectorPrimaryUpdate(VectorPrimaryUpdatePlan { target_keys, .. }) if !target_keys.is_empty() => { DataDependent @@ -121,7 +127,7 @@ impl SqlPlan { Self::InsertSelect { source, .. } | Self::UpdateFrom { source, .. } | Self::Aggregate { input: source, .. } - | Self::Merge { source, .. } => source.cache_eligibility(), + | Self::Merge(MergePlan { source, .. }) => source.cache_eligibility(), Self::Join { left, right, .. } | Self::Intersect { left, right, .. } | Self::Except { left, right, .. } => { @@ -130,14 +136,14 @@ impl SqlPlan { Self::Union { inputs, .. } => inputs.iter().fold(Cacheable, |eligibility, input| { eligibility.combine(input.cache_eligibility()) }), - Self::Cte { definitions, outer } => definitions + Self::Cte(CtePlan { definitions, outer }) => definitions .iter() .fold(outer.cache_eligibility(), |eligibility, (_, plan)| { eligibility.combine(plan.cache_eligibility()) }), Self::Subquery { input, .. } => input.cache_eligibility(), - Self::LateralTopK { outer, .. } => outer.cache_eligibility(), - Self::LateralLoop { outer, inner, .. } => { + Self::LateralTopK(LateralTopKPlan { outer, .. }) => outer.cache_eligibility(), + Self::LateralLoop(LateralLoopPlan { outer, inner, .. }) => { outer.cache_eligibility().combine(inner.cache_eligibility()) } Self::ConstantResult { .. } @@ -318,10 +324,10 @@ mod tests { #[test] fn nested_point_dependency_propagates() { - let plan = SqlPlan::Cte { + let plan = SqlPlan::Cte(CtePlan { definitions: vec![("selected".into(), point_get(EngineType::DocumentStrict))], outer: Box::new(point_get(EngineType::KeyValue)), - }; + }); assert_eq!( plan.cache_eligibility(), PlanCacheEligibility::DataDependent diff --git a/nodedb-sql/src/types/plan/mod.rs b/nodedb-sql/src/types/plan/mod.rs index fc48d02aa..e27478f76 100644 --- a/nodedb-sql/src/types/plan/mod.rs +++ b/nodedb-sql/src/types/plan/mod.rs @@ -19,6 +19,16 @@ pub use expr_scan::{ }; pub use merge_types::{MergeClauseKind, MergePlanAction, MergePlanClause}; pub use row_types::{KvInsertIntent, VectorPrimaryInsertIntent, VectorPrimaryRow, WriteRoute}; +/// Per-family payloads of the larger `SqlPlan` variants. +pub use variants::{ + AlterArrayPlan, ArrayAggPlan, ArrayElementwisePlan, ArrayProjectPlan, ArraySlicePlan, + CreateArrayPlan, CreateIndexPlan, CtePlan, DeleteArrayPlan, DocumentIndexLookupPlan, + DropIndexPlan, HybridSearchPlan, HybridSearchTriplePlan, InsertArrayPlan, InsertPlan, + KvInsertPlan, LateralLoopPlan, LateralTopKPlan, MergePlan, RangeScanPlan, RecursiveScanPlan, + RecursiveValuePlan, TextScoreColumn, TextSearchPlan, TextSearchShape, TimeseriesIngestPlan, + TimeseriesScanPlan, UpsertPlan, VectorPrimaryDeletePlan, VectorPrimaryInsertPlan, + VectorPrimaryTruncatePlan, VectorPrimaryUpdatePlan, +}; pub use variants::{DistanceMetric, SqlPlan}; pub use vector_opts::{ArrayPrefilter, VectorAnnOptions, VectorQuantization}; pub use volatility_scan::expr_is_volatile; diff --git a/nodedb-sql/src/types/plan/variants.rs b/nodedb-sql/src/types/plan/variants.rs deleted file mode 100644 index 63a157a2a..000000000 --- a/nodedb-sql/src/types/plan/variants.rs +++ /dev/null @@ -1,841 +0,0 @@ -// SPDX-License-Identifier: Apache-2.0 - -//! The `SqlPlan` enum — top-level plan produced by the SQL planner. - -use crate::fts_types::FtsQuery; -use crate::temporal::TemporalScope; -use crate::types_array; -use crate::types_expr::{SqlExpr, SqlPayloadAtom, SqlValue}; -pub use nodedb_types::vector_distance::DistanceMetric; - -use crate::types::filter::Filter; -use crate::types::query::{ - AggOutputSlot, AggregateExpr, EngineType, JoinType, Projection, SortKey, SpatialPredicate, - WindowSpec, -}; - -use super::merge_types::MergePlanClause; -use super::row_types::{KvInsertIntent, VectorPrimaryInsertIntent, VectorPrimaryRow, WriteRoute}; -use super::vector_opts::{ArrayPrefilter, VectorAnnOptions}; - -/// The top-level plan produced by the SQL planner. -#[derive(Debug, Clone)] -pub enum SqlPlan { - // ── Constant ── - /// Query with no FROM clause: SELECT 1, SELECT 'hello' AS name, etc. - /// Produces a single row with evaluated constant expressions. - ConstantResult { - columns: Vec, - values: Vec, - /// Whether any projected expression called a `Volatile` function. - /// The values were evaluated while this plan was built, so a cached - /// plan would replay them; a volatile plan is never cached. - volatile: bool, - }, - - // ── Reads ── - Scan { - collection: String, - alias: Option, - engine: EngineType, - filters: Vec, - projection: Vec, - sort_keys: Vec, - limit: Option, - offset: usize, - distinct: bool, - window_functions: Vec, - /// Bitemporal qualifier extracted from `FOR SYSTEM_TIME` / - /// `FOR VALID_TIME`. Default when the scan is current-state. - temporal: TemporalScope, - }, - PointGet { - collection: String, - alias: Option, - engine: EngineType, - key_column: String, - key_value: SqlValue, - /// Resolved SELECT target list, for output-schema derivation. - projection: Vec, - }, - /// Document fetch via a secondary index: equality predicate on an - /// indexed field. The executor performs an index lookup to resolve - /// matching document IDs, reads each document, and applies any - /// remaining filters, projection, sort, and limit. - /// - /// Emitted by `document_schemaless::plan_scan` / - /// `document_strict::plan_scan` when the WHERE clause contains a - /// single equality predicate on a `Ready` indexed field. Any - /// additional predicates fall through as post-filters. - DocumentIndexLookup { - collection: String, - alias: Option, - engine: EngineType, - /// Indexed field path used for the lookup. - field: String, - /// Equality value from the WHERE clause. - value: SqlValue, - /// Remaining filters after extracting the equality used for lookup. - filters: Vec, - projection: Vec, - sort_keys: Vec, - limit: Option, - offset: usize, - distinct: bool, - window_functions: Vec, - /// Whether the chosen index is COLLATE NOCASE — the executor - /// lowercases the lookup value before probing. - case_insensitive: bool, - /// Bitemporal qualifier — mirrors `Scan::temporal`. Document - /// engines must honor it at the Ceiling stage. - temporal: TemporalScope, - }, - RangeScan { - collection: String, - field: String, - lower: Option, - upper: Option, - limit: usize, - /// Resolved SELECT target list, for output-schema derivation. - projection: Vec, - }, - - // ── Writes ── - Insert { - collection: String, - engine: EngineType, - /// The lowering these rows take, chosen by the engine's `EngineRules`. - /// The conversion layer reads it instead of re-deciding from `engine`. - route: WriteRoute, - /// Every declared DEFAULT already materialized and every literal - /// coerced to its declared column type by the planner. - rows: Vec>, - /// Whether a DEFAULT materialized into `rows` was volatile. The - /// planner evaluates declared defaults while building this plan, so a - /// cached plan would replay one execution's value; a volatile plan is - /// never cached. - volatile_defaults: bool, - /// `ON CONFLICT DO NOTHING` semantics: when true, duplicate-PK rows - /// are silently skipped instead of raising `unique_violation`. Plain - /// `INSERT` (no `ON CONFLICT` clause) sets this to `false`. - if_absent: bool, - /// Raw column type strings from the catalog: `(column_name, type_str)`. - /// Forwarded from `InsertParams::column_schema`. Used by columnar - /// converters to reconstruct the exact `ColumnType` for columns whose - /// `SqlDataType` is ambiguous (e.g. JSON and Bytes both map to Bytes). - column_schema: Vec<(String, String)>, - /// Declared `PRIMARY KEY` column name (if any), from - /// `InsertParams::primary_key`. Used by the conversion layer to - /// extract the document id from the correct column instead of - /// guessing at `id`/`document_id`/`key`. - primary_key: Option, - }, - /// KV INSERT: key and value are fundamentally separate. - /// Each entry is `(key, value_columns)`. - KvInsert { - collection: String, - entries: Vec<(SqlValue, Vec<(String, SqlValue)>)>, - /// TTL in seconds (0 = no expiry). Extracted from `ttl` column if present. - ttl_secs: u64, - /// INSERT-vs-UPSERT distinction. `KvOp::Put` is a Redis-SET-style - /// upsert by design; to honor SQL `INSERT` semantics the planner must - /// tell the converter whether a duplicate key should raise (plain - /// `INSERT`, `Insert`), be silently skipped (`ON CONFLICT DO NOTHING`, - /// `InsertIfAbsent`), or overwrite (`UPSERT` / `ON CONFLICT DO - /// UPDATE`, `Put`). - intent: KvInsertIntent, - /// `ON CONFLICT (key) DO UPDATE SET field = expr` assignments, carried - /// through when `intent == Put` via the ON-CONFLICT-DO-UPDATE path. - /// Empty for plain UPSERT (whole-value overwrite) and for INSERT - /// variants. - on_conflict_updates: Vec<(String, SqlExpr)>, - /// Whether any DEFAULT materialized into `entries` came from a - /// `Volatile` expression. The key-value planner evaluates declared - /// defaults while building this plan, so a cached plan would replay - /// one execution's value; a volatile plan is never cached. - volatile_defaults: bool, - }, - /// UPSERT: insert or merge if document exists. - Upsert { - collection: String, - engine: EngineType, - /// The lowering these rows take. Mirrors `Insert::route`. - route: WriteRoute, - /// Defaults materialized and literals coerced, as in `Insert::rows`. - rows: Vec>, - /// Mirrors `Insert::volatile_defaults`. - volatile_defaults: bool, - /// `ON CONFLICT (...) DO UPDATE SET field = expr` assignments. - /// When empty, upsert is a plain merge: new columns overwrite existing. - /// When non-empty, the engine applies these per-row against the - /// *existing* document instead of merging the inserted values. - on_conflict_updates: Vec<(String, SqlExpr)>, - /// Raw column type strings from the catalog: `(column_name, type_str)`. - /// Mirrors `Insert::column_schema` — see that field for rationale. - column_schema: Vec<(String, String)>, - /// Declared `PRIMARY KEY` column name (if any). See `Insert::primary_key`. - primary_key: Option, - }, - InsertSelect { - target: String, - source: Box, - limit: usize, - /// `(target_column, source_expression)`, in target-column order. - /// - /// Empty means passthrough: `INSERT INTO t SELECT * FROM s` copies - /// each source row unchanged. Non-empty means every target column is - /// materialized from the paired expression over the source row. - column_map: Vec<(String, SqlExpr)>, - }, - Update { - collection: String, - engine: EngineType, - assignments: Vec<(String, SqlExpr)>, - filters: Vec, - target_keys: Vec, - returning: bool, - }, - /// `UPDATE target SET col = src.col2 FROM src WHERE target.id = src.id` - /// - /// Two-phase execution: scan `source` with `source_filters`, then for - /// each matched source row that satisfies the join predicates against a - /// target row, apply `assignments` (which may reference source columns - /// via qualified names `src.col`). - /// - /// `join_predicates` are equality pairs `(target_col, source_col)` extracted - /// from the WHERE clause linking the two tables. `target_filters` are - /// remaining WHERE predicates that reference only `target`. - UpdateFrom { - collection: String, - engine: EngineType, - /// The FROM source: a `Scan`, `Join`, or other read plan. - source: Box, - /// Column name used as the target's join key (e.g. `"id"`). - target_join_col: String, - /// Column name used as the source's join key (e.g. `"id"`). - source_join_col: String, - /// SET assignments — RHS may be `SqlExpr::Column { table: Some("src"), .. }`. - assignments: Vec<(String, SqlExpr)>, - /// Filters that apply only to the target collection. - target_filters: Vec, - returning: bool, - }, - Delete { - collection: String, - engine: EngineType, - filters: Vec, - target_keys: Vec, - }, - Truncate { - collection: String, - engine: EngineType, - restart_identity: bool, - }, - - // ── Joins ── - Join { - left: Box, - right: Box, - on: Vec<(String, String)>, - join_type: JoinType, - condition: Option, - /// `None` = no SQL `LIMIT` clause (output bounded downstream by the - /// memory byte budget, never silently truncated); `Some(n)` = explicit - /// `LIMIT n` (output capped at exactly `n`). - limit: Option, - /// Post-join projection: column names to keep (empty = all columns). - projection: Vec, - /// Post-join filters (from WHERE clause). - filters: Vec, - }, - - // ── Aggregation ── - Aggregate { - input: Box, - group_by: Vec, - /// SELECT-list output alias for each GROUP BY key, parallel to - /// `group_by`. `Some(alias)` when the projection aliased the key - /// (`SELECT k AS label ... GROUP BY k` → `Some("label")`); `None` - /// when the projection has no explicit alias for that key, so the - /// output column name falls back to the raw grouped column name. - /// Empty when the plan was built without a projection in scope - /// (treated the same as all-`None`). - group_by_aliases: Vec>, - /// SELECT-list interleaving of `group_by` keys and `aggregates`, in - /// output order. Empty when built without a projection in scope (see - /// `AggOutputSlot`) — output_schema falls back to group-keys-first. - output_order: Vec, - aggregates: Vec, - having: Vec, - limit: usize, - /// When the GROUP BY contains ROLLUP/CUBE/GROUPING SETS, this field holds - /// the expansion. Each inner `Vec` is one grouping set — the indices - /// into `group_by` (the canonical key list) that are *present* (non-NULL) - /// for rows in that set. `None` = plain single-set GROUP BY. - grouping_sets: Option>>, - /// ORDER BY applied to the aggregated rows. Empty = no sort - /// (executor returns groups in hash-map iteration order). - /// Populated by `apply_order_by` when an outer ORDER BY - /// targets a GROUP BY result; the Aggregate executor sorts the - /// finalized group rows before returning. - sort_keys: Vec, - }, - - // ── Timeseries ── - TimeseriesScan { - collection: String, - time_range: (i64, i64), - bucket_interval_ms: i64, - group_by: Vec, - aggregates: Vec, - filters: Vec, - projection: Vec, - gap_fill: String, - limit: usize, - /// ORDER BY applied to the scan result. Empty = the engine's natural - /// order (ascending by the collection's time key). - sort_keys: Vec, - tiered: bool, - /// Bitemporal system-time / valid-time scope. Only non-default - /// on collections created `WITH BITEMPORAL`; `TimeseriesRules::plan_scan` - /// rejects temporal scopes otherwise. - temporal: TemporalScope, - }, - TimeseriesIngest { - collection: String, - /// Defaults materialized and literals coerced, as in `Insert::rows`. - /// A row that omits the `TIME_KEY` column carries its declared - /// default here when one exists; only a row with no time value at - /// all takes the ingest clock. - rows: Vec>, - /// Mirrors `Insert::volatile_defaults`. - volatile_defaults: bool, - }, - - // ── Search (first-class) ── - VectorSearch { - collection: String, - field: String, - query_vector: Vec, - top_k: usize, - ef_search: usize, - /// Distance metric requested by the query operator (`<->`, `<=>`, `<#>`). - /// Overrides the collection-default metric at search time. - metric: DistanceMetric, - filters: Vec, - /// Optional cross-engine prefilter: when set, the ND-array slice - /// runs first and its output cells' surrogates form a bitmap that - /// gates the HNSW candidate set. Set by the planner when an - /// `ORDER BY vector_distance(...) LIMIT k` query is JOINed against - /// `ARRAY_SLICE(...)`. The convert layer lowers this to - /// `VectorOp::Search { inline_prefilter_plan: Some(ArrayOp::SurrogateBitmapScan) }`. - array_prefilter: Option, - /// ANN knobs parsed from the optional third JSON-string argument - /// to `vector_distance(field, query, '{...}')`. - ann_options: VectorAnnOptions, - /// When `true`, the projection contains only the surrogate/PK column - /// and/or `vector_distance(...)` — no payload fields. The Data Plane - /// can skip the document-body fetch entirely for vector-primary - /// collections. Always `false` for non-vector-primary collections - /// (document body is the primary result). - skip_payload_fetch: bool, - /// Predicates against payload-indexed columns on a vector-primary - /// collection. Each atom is `Eq(field, value)`, `In(field, values)`, - /// or `Range(field, ...)`. The convert layer translates SqlValue → - /// nodedb_types::Value and emits them as - /// `VectorOp::Search::payload_filters`. The Data Plane intersects - /// the resulting bitmap with the HNSW candidate set via the - /// per-collection `PayloadIndexSet::pre_filter`. - payload_filters: Vec, - /// Resolved SELECT target list, for output-schema derivation. - projection: Vec, - }, - MultiVectorSearch { - collection: String, - query_vector: Vec, - top_k: usize, - ef_search: usize, - /// Resolved SELECT target list, for output-schema derivation. - projection: Vec, - }, - /// Sparse-vector inverted-index search. - /// - /// Produced by the planner when the leading ORDER BY expression is - /// `sparse_score(field, '{dim: weight, ...}')`. The query literal is - /// parsed into `(dimension, weight)` entries at plan time. The convert - /// layer lowers this to `VectorOp::SparseSearch`, which returns the - /// `top_k` documents with the highest dot-product score — matching the - /// `DESC` (similarity) ordering the SQL author wrote. - SparseSearch { - collection: String, - field: String, - query_entries: Vec<(u32, f32)>, - top_k: usize, - /// Resolved SELECT target list, for output-schema derivation. - projection: Vec, - }, - TextSearch { - collection: String, - /// Structured FTS query. Use `FtsQuery::Plain { text, fuzzy }` for - /// simple keyword search. `FtsQuery::And/Or/Prefix` are supported; - /// `FtsQuery::Phrase` and `FtsQuery::Not` are represented but rejected - /// by the executor with `Unsupported`. - query: FtsQuery, - top_k: usize, - filters: Vec, - /// When set, the SELECT list contains `bm25_score(field, term)` and the - /// caller wants a full-collection scan with the score injected under this - /// alias. The converter emits `TextOp::BM25ScoreScan` instead of - /// `TextOp::Search` so that all documents — including non-matching ones — - /// appear in the response with `null` for the score when they do not - /// contain the term. - score_alias: Option, - /// Resolved SELECT target list, for output-schema derivation. - projection: Vec, - }, - HybridSearch { - collection: String, - query_vector: Vec, - query_text: String, - top_k: usize, - ef_search: usize, - vector_weight: f32, - fuzzy: bool, - /// SELECT-list alias the response should use for the RRF score - /// column. `None` means the executor falls back to the fixed - /// internal field name `rrf_score`. Set by the planner from the - /// SELECT projection's `AS ` for the `rrf_score(...)` call. - score_alias: Option, - /// Resolved SELECT target list, for output-schema derivation. - projection: Vec, - }, - - /// Three-source hybrid search: vector + BM25 text + graph BFS, fused via weighted RRF. - /// - /// Produced when the planner detects `rrf_score(vector_distance(...), - /// bm25_score(...), graph_score(...))` with three source arguments. - HybridSearchTriple { - collection: String, - query_vector: Vec, - query_text: String, - /// Node id used as the BFS seed for the graph leg. - graph_seed_id: String, - /// Maximum BFS depth from the seed node. - graph_depth: usize, - /// Edge label filter for graph BFS. `None` = all edges. - graph_edge_label: Option, - top_k: usize, - ef_search: usize, - fuzzy: bool, - /// Per-source RRF k constants: (vector_k, text_k, graph_k). - rrf_k: (f64, f64, f64), - /// SELECT-list alias for the fused RRF score column. - score_alias: Option, - /// Resolved SELECT target list, for output-schema derivation. - projection: Vec, - }, - SpatialScan { - collection: String, - field: String, - predicate: SpatialPredicate, - query_geometry: nodedb_types::geometry::Geometry, - distance_meters: f64, - attribute_filters: Vec, - limit: usize, - projection: Vec, - }, - - // ── Composite ── - Union { - inputs: Vec, - distinct: bool, - }, - Intersect { - left: Box, - right: Box, - all: bool, - }, - Except { - left: Box, - right: Box, - all: bool, - }, - RecursiveScan { - collection: String, - base_filters: Vec, - recursive_filters: Vec, - /// Equi-join link for tree-traversal recursion: - /// `(collection_field, working_table_field)`. - /// e.g. `("parent_id", "id")` means each iteration finds rows - /// where `collection.parent_id` matches a `working_table.id`. - join_link: Option<(String, String)>, - max_iterations: usize, - distinct: bool, - limit: usize, - /// Resolved SELECT target list, for output-schema derivation. - projection: Vec, - }, - - /// Value-generating recursive CTE (`WITH RECURSIVE name(cols) AS (anchor UNION [ALL] step)`). - /// - /// Unlike `RecursiveScan`, this variant carries no collection reference — the anchor - /// row is produced entirely from literal expressions and each iteration applies the - /// step expressions to the previous row. The executor evaluates this iteratively - /// in the Data Plane without touching storage. - /// - /// All expressions are stored as raw SQL text so they can be serialised across the - /// SPSC bridge without requiring `SqlExpr` to implement `Serialize`. The executor - /// parses them at execution time via the same lightweight expression evaluator used - /// by the procedural executor. - RecursiveValue { - /// CTE name (used in error messages). - cte_name: String, - /// Column names declared on the CTE (e.g. `(n)` in `c(n) AS ...`). - columns: Vec, - /// Anchor SELECT expressions as raw SQL text (one per column). - init_exprs: Vec, - /// Recursive step SELECT expressions as raw SQL text (one per column). - /// May reference column names from `columns`. - step_exprs: Vec, - /// Optional WHERE condition as raw SQL text applied to each new row - /// to decide whether to continue. `None` → run until fixed point. - condition: Option, - /// Maximum iterations before a `RecursionDepthExceeded` error is raised. - max_depth: usize, - /// `false` → UNION ALL (keep duplicates); `true` → UNION (deduplicate). - distinct: bool, - }, - - /// Non-recursive CTE: execute each definition, then the outer query. - Cte { - /// CTE definitions: `(name, subquery_plan)`. - definitions: Vec<(String, SqlPlan)>, - /// The outer query that references CTE names. - outer: Box, - }, - - /// Relational post-processing over a subquery/derived-table body whose leaf - /// plan cannot absorb the outer query's constraints. - /// - /// Produced by CTE / derived-table inlining when the referenced body is not - /// a plain `Scan` (which carries its own filters/sort/limit/offset/distinct) - /// and the outer reference adds constraints the body has no slot for — an - /// `ORDER BY`, `OFFSET`, `DISTINCT`, or a `LIMIT` that must apply *after* a - /// reorder. Without this node those constraints were silently dropped, - /// returning the body's unordered/unoffset/undeduplicated rows. - /// - /// The Data Plane materializes `input`'s rows, then applies, in this order: - /// filter → offset → sort → distinct → project → limit — matching - /// `QueryOp::ProviderScan` semantics (which this lowers onto). Constraints - /// the body CAN absorb (e.g. an unordered `LIMIT` folded into a vector - /// search's `top_k`, or a `WHERE` pushed into the engine as a post-filter) - /// are applied at the leaf during inlining and are NOT repeated here. - Subquery { - /// The subquery body whose rows are post-processed. - input: Box, - /// Predicates applied to the materialized rows (outer `WHERE` that the - /// body leaf did not absorb). - filters: Vec, - /// Outer projection (target list). Empty = inherit the body's columns. - projection: Vec, - /// Window functions evaluated over the post-processed rows. Empty = none. - window_functions: Vec, - /// Outer `ORDER BY` keys applied over the materialized rows. - sort_keys: Vec, - /// Outer `OFFSET` (0 = none). - offset: usize, - /// Outer `DISTINCT`. - distinct: bool, - /// Outer `LIMIT` (`None` = unbounded). - limit: Option, - }, - - // ── Array (ND sparse) ───────────────────────────────────── - /// `CREATE ARRAY DIMS (...) ATTRS (...) TILE_EXTENTS (...)`. - /// AST is engine-agnostic — the Origin converter builds the typed - /// `nodedb_array::ArraySchema` and persists the catalog row. - CreateArray { - name: String, - dims: Vec, - attrs: Vec, - tile_extents: Vec, - cell_order: types_array::ArrayCellOrderAst, - tile_order: types_array::ArrayTileOrderAst, - /// Hilbert-prefix bits for vShard routing (1–16, default 8). - prefix_bits: u8, - /// Audit-retention horizon in milliseconds. `None` = non-bitemporal. - audit_retain_ms: Option, - /// Compliance floor for `audit_retain_ms`. `None` = no floor. - minimum_audit_retain_ms: Option, - }, - /// `DROP ARRAY [IF EXISTS] ` — pure Control-Plane catalog - /// mutation. Per-core array store cleanup happens lazily. - DropArray { name: String, if_exists: bool }, - /// `ALTER ARRAY SET (audit_retain_ms = N, ...)`. - /// - /// Double-`Option` semantics for each diff field: - /// - `None` = key was absent from SET clause → field unchanged. - /// - `Some(None)` = key present with value `NULL` → field set to NULL. - /// - `Some(Some(v))` = key present with integer value → field set to v. - AlterArray { - name: String, - /// New value for `audit_retain_ms`. `Some(None)` unregisters - /// the array from the bitemporal retention registry. - audit_retain_ms: Option>, - /// New value for `minimum_audit_retain_ms`. Cannot be NULL. - minimum_audit_retain_ms: Option, - }, - /// `INSERT INTO ARRAY COORDS (...) VALUES (...) [, ...]`. - InsertArray { - name: String, - rows: Vec, - }, - /// `DELETE FROM ARRAY WHERE COORDS IN ((...), (...))`. - DeleteArray { - name: String, - coords: Vec>, - }, - /// `SELECT * FROM ARRAY_SLICE(name, {dim:[lo,hi],..}, [attrs], limit)`. - ArraySlice { - name: String, - slice: types_array::ArraySliceAst, - /// Attribute names. Empty = all attrs. - attr_projection: Vec, - /// 0 = unlimited. - limit: u32, - /// Bitemporal qualifier. When both axes are `None` / `Any`, the Data - /// Plane returns the live (current) state — the default fast path. - /// Populated from `AS OF SYSTEM TIME` / `AS OF VALID TIME` clauses. - temporal: TemporalScope, - }, - /// `SELECT * FROM ARRAY_PROJECT(name, [attrs])`. - ArrayProject { - name: String, - /// Attribute names. Must be non-empty. - attr_projection: Vec, - }, - /// `SELECT * FROM ARRAY_AGG(name, attr, reducer [, group_by_dim])`. - ArrayAgg { - name: String, - attr: String, - reducer: types_array::ArrayReducerAst, - /// `None` = scalar fold; `Some(name)` = group by that dim. - group_by_dim: Option, - /// Bitemporal qualifier. When both axes are `None` / `Any`, the Data - /// Plane aggregates against the live (current) state — the default - /// fast path. Populated from `AS OF SYSTEM TIME` / `AS OF VALID TIME`. - temporal: TemporalScope, - }, - /// `SELECT * FROM ARRAY_ELEMENTWISE(left, right, op, attr)`. - ArrayElementwise { - left: String, - right: String, - op: types_array::ArrayBinaryOpAst, - attr: String, - }, - /// `SELECT ARRAY_FLUSH(name)` — returns one row `{result: BOOL}`. - ArrayFlush { name: String }, - /// `SELECT ARRAY_COMPACT(name)` — returns one row `{result: BOOL}`. - ArrayCompact { name: String }, - - // ── MERGE ────────────────────────────────────────────────────────── - /// `MERGE INTO target USING source ON ... WHEN ... THEN ...` - /// - /// Supported only for `document_schemaless` and `document_strict` engines. - /// The Data Plane handler evaluates WHEN arms in declaration order and - /// applies the first matching action to each joined or unmatched row. - Merge { - target: String, - engine: EngineType, - /// Source plan (Scan, DocumentIndexLookup, or Join of a sub-select). - source: Box, - /// Column in the target used for the equi-join (from ON clause). - target_join_col: String, - /// Column in the source used for the equi-join (from ON clause). - source_join_col: String, - /// Alias used to qualify source columns in expressions (e.g. `src.col`). - source_alias: String, - /// WHEN arms in declaration order. - clauses: Vec, - returning: bool, - }, - - // ── Lateral joins ─────────────────────────────────────────────────── - /// LATERAL subquery that is equi-correlated and has ORDER BY + LIMIT k. - /// - /// Emitted when the inner subquery has an equi-key correlation to the outer - /// table plus an `ORDER BY ... LIMIT k` clause. The Data Plane scans the - /// inner collection once per outer row applying the equi-filter, sorts by - /// `inner_order_by`, and retains at most `inner_limit` rows. - /// - /// `correlation_keys` is `(outer_col, inner_col)` — the equi-join pairs - /// that correlate inner to outer. - LateralTopK { - /// Plan producing outer rows. - outer: Box, - /// Alias used to qualify outer table columns (e.g. `"u"`). - outer_alias: Option, - /// Inner collection to scan. - inner_collection: String, - /// Pre-filter applied to inner rows (non-correlated filters). - inner_filters: Vec, - /// Sort keys for the inner per-outer-row result. - inner_order_by: Vec, - /// Maximum number of inner rows per outer row. - inner_limit: usize, - /// Equi-join pairs `(outer_col, inner_col)`. - correlation_keys: Vec<(String, String)>, - /// Alias under which inner rows are presented. - lateral_alias: String, - /// Post-lateral projection. - projection: Vec, - /// LEFT join semantics: preserve outer rows even when inner is empty. - left_join: bool, - }, - - /// General LATERAL subquery — per-outer-row correlated nested loop. - /// - /// Emitted for LATERAL subqueries that cannot be rewritten as equi-join - /// hash joins or `LateralTopK`. The Control Plane drives execution: it - /// materialises outer rows, then for each row substitutes the correlation - /// values as additional filters on the inner plan and re-dispatches it. - /// - /// Bounded by `outer_row_cap`; queries that exceed the cap receive a typed - /// `SqlError::Unsupported` before any data is returned. - LateralLoop { - /// Plan producing outer rows. - outer: Box, - /// Alias used to qualify outer table columns. - outer_alias: Option, - /// Inner subquery plan (correlation predicates are injected at runtime). - inner: Box, - /// Correlated predicates extracted from the inner WHERE that reference - /// outer columns. Each entry is `(inner_field, outer_field)`. - correlation_predicates: Vec<(String, String)>, - /// Alias under which inner rows are presented. - lateral_alias: String, - /// Post-lateral projection. - projection: Vec, - /// Maximum outer rows allowed. Queries exceeding this return an error. - outer_row_cap: usize, - /// LEFT join semantics: preserve outer rows even when inner is empty. - left_join: bool, - }, - - // ── Vector-primary ────────────────────────────────────────────────── - /// INSERT / UPSERT into a vector-primary collection. - /// - /// Emitted by the planner instead of the generic `Insert` / `Upsert` - /// variants when the target collection has `primary = - /// PrimaryEngine::Vector`. Each row lowers to one of - /// `VectorOp::DirectInsert` / `DirectInsertIfAbsent` / `DirectUpsert` - /// per `intent`, bypassing full-document MessagePack encoding. - VectorPrimaryInsert { - collection: String, - /// Vector column name (matches `VectorPrimaryConfig::vector_field`). - /// Plumbed to the direct write op so the Data Plane keys its HNSW - /// index by `(tid, collection, field)` — the same key the SELECT - /// path uses. - field: String, - /// Collection-level quantization. Applied via `set_quantization` on - /// the first DirectUpsert so subsequent seals trigger codec-dispatch - /// rebuilds against the configured codec. - quantization: nodedb_types::VectorQuantization, - /// Native storage dtype for vector values (F32 / F16 / BF16). - storage_dtype: nodedb_types::VectorStorageDtype, - /// Payload field names that get equality bitmap indexes. Registered - /// via `payload.add_index` on the first DirectUpsert. - payload_indexes: Vec<(String, nodedb_types::PayloadIndexKind)>, - rows: Vec, - /// Whether any row carries a value materialized from a volatile - /// DEFAULT such as `nextval`. The value was evaluated while the plan - /// was built, so caching the lowered tasks would replay one - /// execution's value into every later one. - volatile_defaults: bool, - /// What an existing primary key means for each row. Mirrors - /// `KvInsert::intent`. - intent: VectorPrimaryInsertIntent, - /// `ON CONFLICT (pk) DO UPDATE SET field = expr` assignments, carried - /// when `intent == Upsert`. Empty means whole-row replace. - on_conflict_updates: Vec<(String, SqlExpr)>, - /// Resolved primary-key column name. See `Insert::primary_key`. - primary_key: Option, - }, - /// DELETE on a vector-primary collection. - /// - /// `target_keys` carries the primary keys when the WHERE clause is a - /// pure primary-key equality (or IN / OR of equalities). Otherwise - /// `filters` is evaluated against every sidecar row on the Data Plane. - VectorPrimaryDelete { - collection: String, - /// Vector column name; keys the HNSW index the rows live in. - field: String, - filters: Vec, - target_keys: Vec, - /// Resolved primary-key column name. See `Insert::primary_key`. - primary_key: Option, - }, - /// TRUNCATE on a vector-primary collection. - /// - /// Removes every row from the HNSW index and its payload sidecar. - VectorPrimaryTruncate { - collection: String, - /// Vector column name; keys the HNSW index the rows live in. - field: String, - restart_identity: bool, - }, - /// UPDATE on a vector-primary collection. - /// - /// `new_vector` is the literal the statement assigns to the vector - /// column, when it assigns one. Every other assignment stays in - /// `assignments` and patches the payload sidecar. - VectorPrimaryUpdate { - collection: String, - /// Vector column name; keys the HNSW index the rows live in. - field: String, - quantization: nodedb_types::VectorQuantization, - storage_dtype: nodedb_types::VectorStorageDtype, - payload_indexes: Vec<(String, nodedb_types::PayloadIndexKind)>, - new_vector: Option>, - assignments: Vec<(String, SqlExpr)>, - filters: Vec, - target_keys: Vec, - returning: bool, - /// Resolved primary-key column name. See `Insert::primary_key`. - primary_key: Option, - }, - - // ── Index DDL ─────────────────────────────────────────────────────── - /// `CREATE [UNIQUE] INDEX [IF NOT EXISTS] name ON collection (field)` - /// - /// Registers a secondary index on the named field of a document or KV - /// collection. The executor backtracks existing rows into the index so - /// it is immediately consistent. - CreateIndex { - /// Name of the index. `None` requests an auto-generated name. - index_name: Option, - /// Target collection. - collection: String, - /// Indexed field path. - field: String, - /// Whether the index enforces uniqueness. - unique: bool, - /// `IF NOT EXISTS` — succeed silently if the index already exists. - if_not_exists: bool, - /// Case-insensitive string collation (`COLLATE NOCASE`). - case_insensitive: bool, - }, - - /// `DROP INDEX [IF EXISTS] name [ON collection]` - /// - /// Removes a secondary index from a collection. All index entries are - /// erased and the index metadata is unregistered. - DropIndex { - /// Name of the index. - index_name: String, - /// Target collection (may be inferred from the index catalog). - collection: Option, - /// `IF EXISTS` — succeed silently if the index does not exist. - if_exists: bool, - }, -} diff --git a/nodedb-sql/src/types/plan/variants/array.rs b/nodedb-sql/src/types/plan/variants/array.rs new file mode 100644 index 000000000..5864c1080 --- /dev/null +++ b/nodedb-sql/src/types/plan/variants/array.rs @@ -0,0 +1,94 @@ +// SPDX-License-Identifier: Apache-2.0 + +//! ND sparse array plan payloads. + +use crate::temporal::TemporalScope; +use crate::types_array; + +/// Payload of [`SqlPlan::CreateArray`](crate::types::SqlPlan::CreateArray). +#[derive(Debug, Clone)] +pub struct CreateArrayPlan { + pub name: String, + pub dims: Vec, + pub attrs: Vec, + pub tile_extents: Vec, + pub cell_order: types_array::ArrayCellOrderAst, + pub tile_order: types_array::ArrayTileOrderAst, + /// Hilbert-prefix bits for vShard routing (1–16, default 8). + pub prefix_bits: u8, + /// Audit-retention horizon in milliseconds. `None` = non-bitemporal. + pub audit_retain_ms: Option, + /// Compliance floor for `audit_retain_ms`. `None` = no floor. + pub minimum_audit_retain_ms: Option, +} + +/// Payload of [`SqlPlan::AlterArray`](crate::types::SqlPlan::AlterArray). +#[derive(Debug, Clone)] +pub struct AlterArrayPlan { + pub name: String, + /// New value for `audit_retain_ms`. `Some(None)` unregisters + /// the array from the bitemporal retention registry. + pub audit_retain_ms: Option>, + /// New value for `minimum_audit_retain_ms`. Cannot be NULL. + pub minimum_audit_retain_ms: Option, +} + +/// Payload of [`SqlPlan::InsertArray`](crate::types::SqlPlan::InsertArray). +#[derive(Debug, Clone)] +pub struct InsertArrayPlan { + pub name: String, + pub rows: Vec, +} + +/// Payload of [`SqlPlan::DeleteArray`](crate::types::SqlPlan::DeleteArray). +#[derive(Debug, Clone)] +pub struct DeleteArrayPlan { + pub name: String, + pub coords: Vec>, +} + +/// Payload of [`SqlPlan::ArraySlice`](crate::types::SqlPlan::ArraySlice). +#[derive(Debug, Clone)] +pub struct ArraySlicePlan { + pub name: String, + pub slice: types_array::ArraySliceAst, + /// Attribute names. Empty = all attrs. + pub attr_projection: Vec, + /// 0 = unlimited. + pub limit: u32, + /// Bitemporal qualifier. When both axes are `None` / `Any`, the Data + /// Plane returns the live (current) state — the default fast path. + /// Populated from `AS OF SYSTEM TIME` / `AS OF VALID TIME` clauses. + pub temporal: TemporalScope, +} + +/// Payload of [`SqlPlan::ArrayProject`](crate::types::SqlPlan::ArrayProject). +#[derive(Debug, Clone)] +pub struct ArrayProjectPlan { + pub name: String, + /// Attribute names. Must be non-empty. + pub attr_projection: Vec, +} + +/// Payload of [`SqlPlan::ArrayAgg`](crate::types::SqlPlan::ArrayAgg). +#[derive(Debug, Clone)] +pub struct ArrayAggPlan { + pub name: String, + pub attr: String, + pub reducer: types_array::ArrayReducerAst, + /// `None` = scalar fold; `Some(name)` = group by that dim. + pub group_by_dim: Option, + /// Bitemporal qualifier. When both axes are `None` / `Any`, the Data + /// Plane aggregates against the live (current) state — the default + /// fast path. Populated from `AS OF SYSTEM TIME` / `AS OF VALID TIME`. + pub temporal: TemporalScope, +} + +/// Payload of [`SqlPlan::ArrayElementwise`](crate::types::SqlPlan::ArrayElementwise). +#[derive(Debug, Clone)] +pub struct ArrayElementwisePlan { + pub left: String, + pub right: String, + pub op: types_array::ArrayBinaryOpAst, + pub attr: String, +} diff --git a/nodedb-sql/src/types/plan/variants/cte.rs b/nodedb-sql/src/types/plan/variants/cte.rs new file mode 100644 index 000000000..0be8acbcb --- /dev/null +++ b/nodedb-sql/src/types/plan/variants/cte.rs @@ -0,0 +1,14 @@ +// SPDX-License-Identifier: Apache-2.0 + +//! Non-recursive CTE plan payload. + +use super::plan::SqlPlan; + +/// Payload of [`SqlPlan::Cte`](crate::types::SqlPlan::Cte). +#[derive(Debug, Clone)] +pub struct CtePlan { + /// CTE definitions: `(name, subquery_plan)`. + pub definitions: Vec<(String, SqlPlan)>, + /// The outer query that references CTE names. + pub outer: Box, +} diff --git a/nodedb-sql/src/types/plan/variants/hybrid.rs b/nodedb-sql/src/types/plan/variants/hybrid.rs new file mode 100644 index 000000000..b678ade16 --- /dev/null +++ b/nodedb-sql/src/types/plan/variants/hybrid.rs @@ -0,0 +1,71 @@ +// SPDX-License-Identifier: Apache-2.0 + +//! Hybrid search plan payloads: vector + text, and vector + text + graph. + +use nodedb_types::text_search::QueryMode; + +use crate::types::filter::Filter; +use crate::types::query::Projection; + +/// Payload of [`SqlPlan::HybridSearch`](crate::types::SqlPlan::HybridSearch). +#[derive(Debug, Clone)] +pub struct HybridSearchPlan { + pub collection: String, + /// Vector column of the `vector_distance(column, q)` leg. + pub vector_field: String, + pub query_vector: Vec, + /// Column of the `bm25_score(column, q)` leg. `None` for `*`: the + /// whole-document index. + pub text_field: Option, + pub query_text: String, + /// Residual WHERE predicates. They restrict both legs before fusion. + pub filters: Vec, + pub top_k: usize, + pub ef_search: usize, + pub vector_weight: f32, + /// `mode => 'and' | 'or'` of the `bm25_score` leg. + pub mode: QueryMode, + /// `fuzzy => true | false` of the `bm25_score` leg. + pub fuzzy: bool, + /// SELECT-list alias the response should use for the RRF score + /// column. `None` means the executor falls back to the fixed + /// internal field name `rrf_score`. Set by the planner from the + /// SELECT projection's `AS ` for the `rrf_score(...)` call. + pub score_alias: Option, + /// Resolved SELECT target list, for output-schema derivation. + pub projection: Vec, +} + +/// Payload of [`SqlPlan::HybridSearchTriple`](crate::types::SqlPlan::HybridSearchTriple). +#[derive(Debug, Clone)] +pub struct HybridSearchTriplePlan { + pub collection: String, + /// Vector column of the `vector_distance(column, q)` leg. + pub vector_field: String, + pub query_vector: Vec, + /// Column of the `bm25_score(column, q)` leg. `None` for `*`: the + /// whole-document index. + pub text_field: Option, + pub query_text: String, + /// Residual WHERE predicates. They restrict the vector and text legs + /// before fusion. + pub filters: Vec, + /// Node id used as the BFS seed for the graph leg. + pub graph_seed_id: String, + /// Maximum BFS depth from the seed node. + pub graph_depth: usize, + /// Edge label filter for graph BFS. `None` = all edges. + pub graph_edge_label: Option, + pub top_k: usize, + pub ef_search: usize, + /// `mode => 'and' | 'or'` of the `bm25_score` leg. + pub mode: QueryMode, + /// `fuzzy => true | false` of the `bm25_score` leg. + pub fuzzy: bool, + /// Per-source RRF k constants: (vector_k, text_k, graph_k). + pub rrf_k: (f64, f64, f64), + /// SELECT-list alias for the fused RRF score column. + pub score_alias: Option, + /// Resolved SELECT target list, for output-schema derivation. + pub projection: Vec, +} diff --git a/nodedb-sql/src/types/plan/variants/index_ddl.rs b/nodedb-sql/src/types/plan/variants/index_ddl.rs new file mode 100644 index 000000000..33e5e312f --- /dev/null +++ b/nodedb-sql/src/types/plan/variants/index_ddl.rs @@ -0,0 +1,31 @@ +// SPDX-License-Identifier: Apache-2.0 + +//! Index DDL plan payloads. + +/// Payload of [`SqlPlan::CreateIndex`](crate::types::SqlPlan::CreateIndex). +#[derive(Debug, Clone)] +pub struct CreateIndexPlan { + /// Name of the index. `None` requests an auto-generated name. + pub index_name: Option, + /// Target collection. + pub collection: String, + /// Indexed field path. + pub field: String, + /// Whether the index enforces uniqueness. + pub unique: bool, + /// `IF NOT EXISTS` — succeed silently if the index already exists. + pub if_not_exists: bool, + /// Case-insensitive string collation (`COLLATE NOCASE`). + pub case_insensitive: bool, +} + +/// Payload of [`SqlPlan::DropIndex`](crate::types::SqlPlan::DropIndex). +#[derive(Debug, Clone)] +pub struct DropIndexPlan { + /// Name of the index. + pub index_name: String, + /// Target collection (may be inferred from the index catalog). + pub collection: Option, + /// `IF EXISTS` — succeed silently if the index does not exist. + pub if_exists: bool, +} diff --git a/nodedb-sql/src/types/plan/variants/index_reads.rs b/nodedb-sql/src/types/plan/variants/index_reads.rs new file mode 100644 index 000000000..0d911ea46 --- /dev/null +++ b/nodedb-sql/src/types/plan/variants/index_reads.rs @@ -0,0 +1,46 @@ +// SPDX-License-Identifier: Apache-2.0 + +//! Secondary-index read plan payloads: index lookup and range scan. + +use crate::temporal::TemporalScope; +use crate::types::filter::Filter; +use crate::types::query::{EngineType, Projection, SortKey, WindowSpec}; +use crate::types_expr::SqlValue; + +/// Payload of [`SqlPlan::DocumentIndexLookup`](crate::types::SqlPlan::DocumentIndexLookup). +#[derive(Debug, Clone)] +pub struct DocumentIndexLookupPlan { + pub collection: String, + pub alias: Option, + pub engine: EngineType, + /// Indexed field path used for the lookup. + pub field: String, + /// Equality value from the WHERE clause. + pub value: SqlValue, + /// Remaining filters after extracting the equality used for lookup. + pub filters: Vec, + pub projection: Vec, + pub sort_keys: Vec, + pub limit: Option, + pub offset: usize, + pub distinct: bool, + pub window_functions: Vec, + /// Whether the chosen index is COLLATE NOCASE — the executor + /// lowercases the lookup value before probing. + pub case_insensitive: bool, + /// Bitemporal qualifier — mirrors `Scan::temporal`. Document + /// engines must honor it at the Ceiling stage. + pub temporal: TemporalScope, +} + +/// Payload of [`SqlPlan::RangeScan`](crate::types::SqlPlan::RangeScan). +#[derive(Debug, Clone)] +pub struct RangeScanPlan { + pub collection: String, + pub field: String, + pub lower: Option, + pub upper: Option, + pub limit: usize, + /// Resolved SELECT target list, for output-schema derivation. + pub projection: Vec, +} diff --git a/nodedb-sql/src/types/plan/variants/lateral.rs b/nodedb-sql/src/types/plan/variants/lateral.rs new file mode 100644 index 000000000..7dfc8d9dc --- /dev/null +++ b/nodedb-sql/src/types/plan/variants/lateral.rs @@ -0,0 +1,55 @@ +// SPDX-License-Identifier: Apache-2.0 + +//! LATERAL join plan payloads. + +use crate::types::filter::Filter; +use crate::types::query::{Projection, SortKey}; + +use super::plan::SqlPlan; + +/// Payload of [`SqlPlan::LateralTopK`](crate::types::SqlPlan::LateralTopK). +#[derive(Debug, Clone)] +pub struct LateralTopKPlan { + /// Plan producing outer rows. + pub outer: Box, + /// Alias used to qualify outer table columns (e.g. `"u"`). + pub outer_alias: Option, + /// Inner collection to scan. + pub inner_collection: String, + /// Pre-filter applied to inner rows (non-correlated filters). + pub inner_filters: Vec, + /// Sort keys for the inner per-outer-row result. + pub inner_order_by: Vec, + /// Maximum number of inner rows per outer row. + pub inner_limit: usize, + /// Equi-join pairs `(outer_col, inner_col)`. + pub correlation_keys: Vec<(String, String)>, + /// Alias under which inner rows are presented. + pub lateral_alias: String, + /// Post-lateral projection. + pub projection: Vec, + /// LEFT join semantics: preserve outer rows even when inner is empty. + pub left_join: bool, +} + +/// Payload of [`SqlPlan::LateralLoop`](crate::types::SqlPlan::LateralLoop). +#[derive(Debug, Clone)] +pub struct LateralLoopPlan { + /// Plan producing outer rows. + pub outer: Box, + /// Alias used to qualify outer table columns. + pub outer_alias: Option, + /// Inner subquery plan (correlation predicates are injected at runtime). + pub inner: Box, + /// Correlated predicates extracted from the inner WHERE that reference + /// outer columns. Each entry is `(inner_field, outer_field)`. + pub correlation_predicates: Vec<(String, String)>, + /// Alias under which inner rows are presented. + pub lateral_alias: String, + /// Post-lateral projection. + pub projection: Vec, + /// Maximum outer rows allowed. Queries exceeding this return an error. + pub outer_row_cap: usize, + /// LEFT join semantics: preserve outer rows even when inner is empty. + pub left_join: bool, +} diff --git a/nodedb-sql/src/types/plan/variants/merge.rs b/nodedb-sql/src/types/plan/variants/merge.rs new file mode 100644 index 000000000..44e321ef1 --- /dev/null +++ b/nodedb-sql/src/types/plan/variants/merge.rs @@ -0,0 +1,26 @@ +// SPDX-License-Identifier: Apache-2.0 + +//! `MERGE` plan payload. + +use crate::types::query::EngineType; + +use super::super::merge_types::MergePlanClause; +use super::plan::SqlPlan; + +/// Payload of [`SqlPlan::Merge`](crate::types::SqlPlan::Merge). +#[derive(Debug, Clone)] +pub struct MergePlan { + pub target: String, + pub engine: EngineType, + /// Source plan (Scan, DocumentIndexLookup, or Join of a sub-select). + pub source: Box, + /// Column in the target used for the equi-join (from ON clause). + pub target_join_col: String, + /// Column in the source used for the equi-join (from ON clause). + pub source_join_col: String, + /// Alias used to qualify source columns in expressions (e.g. `src.col`). + pub source_alias: String, + /// WHEN arms in declaration order. + pub clauses: Vec, + pub returning: bool, +} diff --git a/nodedb-sql/src/types/plan/variants/mod.rs b/nodedb-sql/src/types/plan/variants/mod.rs new file mode 100644 index 000000000..d7240b902 --- /dev/null +++ b/nodedb-sql/src/types/plan/variants/mod.rs @@ -0,0 +1,37 @@ +// SPDX-License-Identifier: Apache-2.0 + +//! `SqlPlan` and its per-family payload structs. + +mod array; +mod cte; +mod hybrid; +mod index_ddl; +mod index_reads; +mod lateral; +mod merge; +mod plan; +mod recursive; +mod text; +mod timeseries; +mod vector_primary; +mod writes; + +pub use array::{ + AlterArrayPlan, ArrayAggPlan, ArrayElementwisePlan, ArrayProjectPlan, ArraySlicePlan, + CreateArrayPlan, DeleteArrayPlan, InsertArrayPlan, +}; +pub use cte::CtePlan; +pub use hybrid::{HybridSearchPlan, HybridSearchTriplePlan}; +pub use index_ddl::{CreateIndexPlan, DropIndexPlan}; +pub use index_reads::{DocumentIndexLookupPlan, RangeScanPlan}; +pub use lateral::{LateralLoopPlan, LateralTopKPlan}; +pub use merge::MergePlan; +pub use plan::{DistanceMetric, SqlPlan}; +pub use recursive::{RecursiveScanPlan, RecursiveValuePlan}; +pub use text::{TextScoreColumn, TextSearchPlan, TextSearchShape}; +pub use timeseries::{TimeseriesIngestPlan, TimeseriesScanPlan}; +pub use vector_primary::{ + VectorPrimaryDeletePlan, VectorPrimaryInsertPlan, VectorPrimaryTruncatePlan, + VectorPrimaryUpdatePlan, +}; +pub use writes::{InsertPlan, KvInsertPlan, UpsertPlan}; diff --git a/nodedb-sql/src/types/plan/variants/plan.rs b/nodedb-sql/src/types/plan/variants/plan.rs new file mode 100644 index 000000000..a4e831565 --- /dev/null +++ b/nodedb-sql/src/types/plan/variants/plan.rs @@ -0,0 +1,471 @@ +// SPDX-License-Identifier: Apache-2.0 + +//! The `SqlPlan` enum — top-level plan produced by the SQL planner. Larger +//! payloads live in per-family structs beside it. + +use crate::temporal::TemporalScope; +use crate::types_expr::{SqlExpr, SqlPayloadAtom, SqlValue}; +pub use nodedb_types::vector_distance::DistanceMetric; + +use crate::types::filter::Filter; +use crate::types::query::{ + AggOutputSlot, AggregateExpr, EngineType, JoinType, Projection, SortKey, SpatialPredicate, + WindowSpec, +}; + +use super::super::vector_opts::{ArrayPrefilter, VectorAnnOptions}; + +use super::array::{ + AlterArrayPlan, ArrayAggPlan, ArrayElementwisePlan, ArrayProjectPlan, ArraySlicePlan, + CreateArrayPlan, DeleteArrayPlan, InsertArrayPlan, +}; +use super::cte::CtePlan; +use super::hybrid::{HybridSearchPlan, HybridSearchTriplePlan}; +use super::index_ddl::{CreateIndexPlan, DropIndexPlan}; +use super::index_reads::{DocumentIndexLookupPlan, RangeScanPlan}; +use super::lateral::{LateralLoopPlan, LateralTopKPlan}; +use super::merge::MergePlan; +use super::recursive::{RecursiveScanPlan, RecursiveValuePlan}; +use super::text::TextSearchPlan; +use super::timeseries::{TimeseriesIngestPlan, TimeseriesScanPlan}; +use super::vector_primary::{ + VectorPrimaryDeletePlan, VectorPrimaryInsertPlan, VectorPrimaryTruncatePlan, + VectorPrimaryUpdatePlan, +}; +use super::writes::{InsertPlan, KvInsertPlan, UpsertPlan}; + +/// The top-level plan produced by the SQL planner. +#[derive(Debug, Clone)] +pub enum SqlPlan { + // ── Constant ── + /// Query with no FROM clause: SELECT 1, SELECT 'hello' AS name, etc. + /// Produces a single row with evaluated constant expressions. + ConstantResult { + columns: Vec, + values: Vec, + /// Whether any projected expression called a `Volatile` function. + /// The values were evaluated while this plan was built, so a cached + /// plan would replay them; a volatile plan is never cached. + volatile: bool, + }, + + // ── Reads ── + Scan { + collection: String, + alias: Option, + engine: EngineType, + filters: Vec, + projection: Vec, + sort_keys: Vec, + limit: Option, + offset: usize, + distinct: bool, + window_functions: Vec, + /// Bitemporal qualifier extracted from `FOR SYSTEM_TIME` / + /// `FOR VALID_TIME`. Default when the scan is current-state. + temporal: TemporalScope, + }, + PointGet { + collection: String, + alias: Option, + engine: EngineType, + key_column: String, + key_value: SqlValue, + /// Resolved SELECT target list, for output-schema derivation. + projection: Vec, + }, + /// Document fetch via a secondary index: equality predicate on an + /// indexed field. The executor performs an index lookup to resolve + /// matching document IDs, reads each document, and applies any + /// remaining filters, projection, sort, and limit. + /// + /// Emitted by `document_schemaless::plan_scan` / + /// `document_strict::plan_scan` when the WHERE clause contains a + /// single equality predicate on a `Ready` indexed field. Any + /// additional predicates fall through as post-filters. + DocumentIndexLookup(DocumentIndexLookupPlan), + RangeScan(RangeScanPlan), + + // ── Writes ── + Insert(InsertPlan), + /// KV INSERT: key and value are fundamentally separate. + /// Each entry is `(key, value_columns)`. + KvInsert(KvInsertPlan), + /// UPSERT: insert or merge if document exists. + Upsert(UpsertPlan), + InsertSelect { + target: String, + source: Box, + limit: usize, + /// `(target_column, source_expression)`, in target-column order. + /// + /// Empty means passthrough: `INSERT INTO t SELECT * FROM s` copies + /// each source row unchanged. Non-empty means every target column is + /// materialized from the paired expression over the source row. + column_map: Vec<(String, SqlExpr)>, + }, + Update { + collection: String, + engine: EngineType, + assignments: Vec<(String, SqlExpr)>, + filters: Vec, + target_keys: Vec, + returning: bool, + }, + /// `UPDATE target SET col = src.col2 FROM src WHERE target.id = src.id` + /// + /// Two-phase execution: scan `source` with `source_filters`, then for + /// each matched source row that satisfies the join predicates against a + /// target row, apply `assignments` (which may reference source columns + /// via qualified names `src.col`). + /// + /// `join_predicates` are equality pairs `(target_col, source_col)` extracted + /// from the WHERE clause linking the two tables. `target_filters` are + /// remaining WHERE predicates that reference only `target`. + UpdateFrom { + collection: String, + engine: EngineType, + /// The FROM source: a `Scan`, `Join`, or other read plan. + source: Box, + /// Column name used as the target's join key (e.g. `"id"`). + target_join_col: String, + /// Column name used as the source's join key (e.g. `"id"`). + source_join_col: String, + /// SET assignments — RHS may be `SqlExpr::Column { table: Some("src"), .. }`. + assignments: Vec<(String, SqlExpr)>, + /// Filters that apply only to the target collection. + target_filters: Vec, + returning: bool, + }, + Delete { + collection: String, + engine: EngineType, + filters: Vec, + target_keys: Vec, + }, + Truncate { + collection: String, + engine: EngineType, + restart_identity: bool, + }, + + // ── Joins ── + Join { + left: Box, + right: Box, + on: Vec<(String, String)>, + join_type: JoinType, + condition: Option, + /// `None` = no SQL `LIMIT` clause (output bounded downstream by the + /// memory byte budget, never silently truncated); `Some(n)` = explicit + /// `LIMIT n` (output capped at exactly `n`). + limit: Option, + /// Post-join projection: column names to keep (empty = all columns). + projection: Vec, + /// Post-join filters (from WHERE clause). + filters: Vec, + }, + + // ── Aggregation ── + Aggregate { + input: Box, + group_by: Vec, + /// SELECT-list output alias for each GROUP BY key, parallel to + /// `group_by`. `Some(alias)` when the projection aliased the key + /// (`SELECT k AS label ... GROUP BY k` → `Some("label")`); `None` + /// when the projection has no explicit alias for that key, so the + /// output column name falls back to the raw grouped column name. + /// Empty when the plan was built without a projection in scope + /// (treated the same as all-`None`). + group_by_aliases: Vec>, + /// SELECT-list interleaving of `group_by` keys and `aggregates`, in + /// output order. Empty when built without a projection in scope (see + /// `AggOutputSlot`) — output_schema falls back to group-keys-first. + output_order: Vec, + aggregates: Vec, + having: Vec, + limit: usize, + /// When the GROUP BY contains ROLLUP/CUBE/GROUPING SETS, this field holds + /// the expansion. Each inner `Vec` is one grouping set — the indices + /// into `group_by` (the canonical key list) that are *present* (non-NULL) + /// for rows in that set. `None` = plain single-set GROUP BY. + grouping_sets: Option>>, + /// ORDER BY applied to the aggregated rows. Empty = no sort + /// (executor returns groups in hash-map iteration order). + /// Populated by `apply_order_by` when an outer ORDER BY + /// targets a GROUP BY result; the Aggregate executor sorts the + /// finalized group rows before returning. + sort_keys: Vec, + }, + + // ── Timeseries ── + TimeseriesScan(TimeseriesScanPlan), + TimeseriesIngest(TimeseriesIngestPlan), + + // ── Search (first-class) ── + VectorSearch { + collection: String, + field: String, + query_vector: Vec, + top_k: usize, + ef_search: usize, + /// Distance metric requested by the query operator (`<->`, `<=>`, `<#>`). + /// Overrides the collection-default metric at search time. + metric: DistanceMetric, + filters: Vec, + /// Optional cross-engine prefilter: when set, the ND-array slice + /// runs first and its output cells' surrogates form a bitmap that + /// gates the HNSW candidate set. Set by the planner when an + /// `ORDER BY vector_distance(...) LIMIT k` query is JOINed against + /// `ARRAY_SLICE(...)`. The convert layer lowers this to + /// `VectorOp::Search { inline_prefilter_plan: Some(ArrayOp::SurrogateBitmapScan) }`. + array_prefilter: Option, + /// ANN knobs parsed from the optional third JSON-string argument + /// to `vector_distance(field, query, '{...}')`. + ann_options: VectorAnnOptions, + /// When `true`, the projection contains only the surrogate/PK column + /// and/or `vector_distance(...)` — no payload fields. The Data Plane + /// can skip the document-body fetch entirely for vector-primary + /// collections. Always `false` for non-vector-primary collections + /// (document body is the primary result). + skip_payload_fetch: bool, + /// Predicates against payload-indexed columns on a vector-primary + /// collection. Each atom is `Eq(field, value)`, `In(field, values)`, + /// or `Range(field, ...)`. The convert layer translates SqlValue → + /// nodedb_types::Value and emits them as + /// `VectorOp::Search::payload_filters`. The Data Plane intersects + /// the resulting bitmap with the HNSW candidate set via the + /// per-collection `PayloadIndexSet::pre_filter`. + payload_filters: Vec, + /// Primary keys a top-level `WHERE pk = v` / `pk IN (...)` conjunct + /// names. The search ranks only those rows: the convert layer lowers + /// the keys to the candidate bitmap the index search honors. `None`: + /// no key restriction. `Some(empty)`: no row is a candidate. + pk_prefilter: Option>, + /// Resolved SELECT target list, for output-schema derivation. + projection: Vec, + }, + MultiVectorSearch { + collection: String, + query_vector: Vec, + top_k: usize, + ef_search: usize, + /// Resolved SELECT target list, for output-schema derivation. + projection: Vec, + }, + /// Sparse-vector inverted-index search. + /// + /// Produced by the planner when the leading ORDER BY expression is + /// `sparse_score(field, '{dim: weight, ...}')`. The query literal is + /// parsed into `(dimension, weight)` entries at plan time. The convert + /// layer lowers this to `VectorOp::SparseSearch`, which returns the + /// `top_k` documents with the highest dot-product score — matching the + /// `DESC` (similarity) ordering the SQL author wrote. + SparseSearch { + collection: String, + field: String, + query_entries: Vec<(u32, f32)>, + top_k: usize, + /// Resolved SELECT target list, for output-schema derivation. + projection: Vec, + }, + /// Full-text search: `WHERE text_match(...)` matches, or a + /// `bm25_score(...)` score scan over every row the filters admit. + TextSearch(TextSearchPlan), + HybridSearch(HybridSearchPlan), + + /// Three-source hybrid search: vector + BM25 text + graph BFS, fused via weighted RRF. + /// + /// Produced when the planner detects `rrf_score(vector_distance(...), + /// bm25_score(...), graph_score(...))` with three source arguments. + HybridSearchTriple(HybridSearchTriplePlan), + SpatialScan { + collection: String, + field: String, + predicate: SpatialPredicate, + query_geometry: nodedb_types::geometry::Geometry, + distance_meters: f64, + attribute_filters: Vec, + limit: usize, + projection: Vec, + }, + + // ── Composite ── + Union { + inputs: Vec, + distinct: bool, + }, + Intersect { + left: Box, + right: Box, + all: bool, + }, + Except { + left: Box, + right: Box, + all: bool, + }, + RecursiveScan(RecursiveScanPlan), + + /// Value-generating recursive CTE (`WITH RECURSIVE name(cols) AS (anchor UNION [ALL] step)`). + /// + /// Unlike `RecursiveScan`, this variant carries no collection reference — the anchor + /// row is produced entirely from literal expressions and each iteration applies the + /// step expressions to the previous row. The executor evaluates this iteratively + /// in the Data Plane without touching storage. + /// + /// All expressions are stored as raw SQL text so they can be serialised across the + /// SPSC bridge without requiring `SqlExpr` to implement `Serialize`. The executor + /// parses them at execution time via the same lightweight expression evaluator used + /// by the procedural executor. + RecursiveValue(RecursiveValuePlan), + + /// Non-recursive CTE: execute each definition, then the outer query. + Cte(CtePlan), + + /// Relational post-processing over a subquery/derived-table body whose leaf + /// plan cannot absorb the outer query's constraints. + /// + /// Produced by CTE / derived-table inlining when the referenced body is not + /// a plain `Scan` (which carries its own filters/sort/limit/offset/distinct) + /// and the outer reference adds constraints the body has no slot for — an + /// `ORDER BY`, `OFFSET`, `DISTINCT`, or a `LIMIT` that must apply *after* a + /// reorder. Without this node those constraints were silently dropped, + /// returning the body's unordered/unoffset/undeduplicated rows. + /// + /// The Data Plane materializes `input`'s rows, then applies, in this order: + /// filter → offset → sort → distinct → project → limit — matching + /// `QueryOp::ProviderScan` semantics (which this lowers onto). Constraints + /// the body CAN absorb (e.g. an unordered `LIMIT` folded into a vector + /// search's `top_k`, or a `WHERE` pushed into the engine as a post-filter) + /// are applied at the leaf during inlining and are NOT repeated here. + Subquery { + /// The subquery body whose rows are post-processed. + input: Box, + /// Predicates applied to the materialized rows (outer `WHERE` that the + /// body leaf did not absorb). + filters: Vec, + /// Outer projection (target list). Empty = inherit the body's columns. + projection: Vec, + /// Window functions evaluated over the post-processed rows. Empty = none. + window_functions: Vec, + /// Outer `ORDER BY` keys applied over the materialized rows. + sort_keys: Vec, + /// Outer `OFFSET` (0 = none). + offset: usize, + /// Outer `DISTINCT`. + distinct: bool, + /// Outer `LIMIT` (`None` = unbounded). + limit: Option, + }, + + // ── Array (ND sparse) ───────────────────────────────────── + /// `CREATE ARRAY DIMS (...) ATTRS (...) TILE_EXTENTS (...)`. + /// AST is engine-agnostic — the Origin converter builds the typed + /// `nodedb_array::ArraySchema` and persists the catalog row. + CreateArray(CreateArrayPlan), + /// `DROP ARRAY [IF EXISTS] ` — pure Control-Plane catalog + /// mutation. Per-core array store cleanup happens lazily. + DropArray { + name: String, + if_exists: bool, + }, + /// `ALTER ARRAY SET (audit_retain_ms = N, ...)`. + /// + /// Double-`Option` semantics for each diff field: + /// - `None` = key was absent from SET clause → field unchanged. + /// - `Some(None)` = key present with value `NULL` → field set to NULL. + /// - `Some(Some(v))` = key present with integer value → field set to v. + AlterArray(AlterArrayPlan), + /// `INSERT INTO ARRAY COORDS (...) VALUES (...) [, ...]`. + InsertArray(InsertArrayPlan), + /// `DELETE FROM ARRAY WHERE COORDS IN ((...), (...))`. + DeleteArray(DeleteArrayPlan), + /// `SELECT * FROM ARRAY_SLICE(name, {dim:[lo,hi],..}, [attrs], limit)`. + ArraySlice(ArraySlicePlan), + /// `SELECT * FROM ARRAY_PROJECT(name, [attrs])`. + ArrayProject(ArrayProjectPlan), + /// `SELECT * FROM ARRAY_AGG(name, attr, reducer [, group_by_dim])`. + ArrayAgg(ArrayAggPlan), + /// `SELECT * FROM ARRAY_ELEMENTWISE(left, right, op, attr)`. + ArrayElementwise(ArrayElementwisePlan), + /// `SELECT ARRAY_FLUSH(name)` — returns one row `{result: BOOL}`. + ArrayFlush { + name: String, + }, + /// `SELECT ARRAY_COMPACT(name)` — returns one row `{result: BOOL}`. + ArrayCompact { + name: String, + }, + + // ── MERGE ────────────────────────────────────────────────────────── + /// `MERGE INTO target USING source ON ... WHEN ... THEN ...` + /// + /// Supported only for `document_schemaless` and `document_strict` engines. + /// The Data Plane handler evaluates WHEN arms in declaration order and + /// applies the first matching action to each joined or unmatched row. + Merge(MergePlan), + + // ── Lateral joins ─────────────────────────────────────────────────── + /// LATERAL subquery that is equi-correlated and has ORDER BY + LIMIT k. + /// + /// Emitted when the inner subquery has an equi-key correlation to the outer + /// table plus an `ORDER BY ... LIMIT k` clause. The Data Plane scans the + /// inner collection once per outer row applying the equi-filter, sorts by + /// `inner_order_by`, and retains at most `inner_limit` rows. + /// + /// `correlation_keys` is `(outer_col, inner_col)` — the equi-join pairs + /// that correlate inner to outer. + LateralTopK(LateralTopKPlan), + + /// General LATERAL subquery — per-outer-row correlated nested loop. + /// + /// Emitted for LATERAL subqueries that cannot be rewritten as equi-join + /// hash joins or `LateralTopK`. The Control Plane drives execution: it + /// materialises outer rows, then for each row substitutes the correlation + /// values as additional filters on the inner plan and re-dispatches it. + /// + /// Bounded by `outer_row_cap`; queries that exceed the cap receive a typed + /// `SqlError::Unsupported` before any data is returned. + LateralLoop(LateralLoopPlan), + + // ── Vector-primary ────────────────────────────────────────────────── + /// INSERT / UPSERT into a vector-primary collection. + /// + /// Emitted by the planner instead of the generic `Insert` / `Upsert` + /// variants when the target collection has `primary = + /// PrimaryEngine::Vector`. Each row lowers to one of + /// `VectorOp::DirectInsert` / `DirectInsertIfAbsent` / `DirectUpsert` + /// per `intent`, bypassing full-document MessagePack encoding. + VectorPrimaryInsert(VectorPrimaryInsertPlan), + /// DELETE on a vector-primary collection. + /// + /// `target_keys` carries the primary keys when the WHERE clause is a + /// pure primary-key equality (or IN / OR of equalities). Otherwise + /// `filters` is evaluated against every sidecar row on the Data Plane. + VectorPrimaryDelete(VectorPrimaryDeletePlan), + /// TRUNCATE on a vector-primary collection. + /// + /// Removes every row from the HNSW index and its payload sidecar. + VectorPrimaryTruncate(VectorPrimaryTruncatePlan), + /// UPDATE on a vector-primary collection. + /// + /// `new_vector` is the literal the statement assigns to the vector + /// column, when it assigns one. Every other assignment stays in + /// `assignments` and patches the payload sidecar. + VectorPrimaryUpdate(VectorPrimaryUpdatePlan), + + // ── Index DDL ─────────────────────────────────────────────────────── + /// `CREATE [UNIQUE] INDEX [IF NOT EXISTS] name ON collection (field)` + /// + /// Registers a secondary index on the named field of a document or KV + /// collection. The executor backtracks existing rows into the index so + /// it is immediately consistent. + CreateIndex(CreateIndexPlan), + + /// `DROP INDEX [IF EXISTS] name [ON collection]` + /// + /// Removes a secondary index from a collection. All index entries are + /// erased and the index metadata is unregistered. + DropIndex(DropIndexPlan), +} diff --git a/nodedb-sql/src/types/plan/variants/recursive.rs b/nodedb-sql/src/types/plan/variants/recursive.rs new file mode 100644 index 000000000..0b16943ed --- /dev/null +++ b/nodedb-sql/src/types/plan/variants/recursive.rs @@ -0,0 +1,45 @@ +// SPDX-License-Identifier: Apache-2.0 + +//! Recursive CTE plan payloads. + +use crate::types::filter::Filter; +use crate::types::query::Projection; + +/// Payload of [`SqlPlan::RecursiveScan`](crate::types::SqlPlan::RecursiveScan). +#[derive(Debug, Clone)] +pub struct RecursiveScanPlan { + pub collection: String, + pub base_filters: Vec, + pub recursive_filters: Vec, + /// Equi-join link for tree-traversal recursion: + /// `(collection_field, working_table_field)`. + /// e.g. `("parent_id", "id")` means each iteration finds rows + /// where `collection.parent_id` matches a `working_table.id`. + pub join_link: Option<(String, String)>, + pub max_iterations: usize, + pub distinct: bool, + pub limit: usize, + /// Resolved SELECT target list, for output-schema derivation. + pub projection: Vec, +} + +/// Payload of [`SqlPlan::RecursiveValue`](crate::types::SqlPlan::RecursiveValue). +#[derive(Debug, Clone)] +pub struct RecursiveValuePlan { + /// CTE name (used in error messages). + pub cte_name: String, + /// Column names declared on the CTE (e.g. `(n)` in `c(n) AS ...`). + pub columns: Vec, + /// Anchor SELECT expressions as raw SQL text (one per column). + pub init_exprs: Vec, + /// Recursive step SELECT expressions as raw SQL text (one per column). + /// May reference column names from `columns`. + pub step_exprs: Vec, + /// Optional WHERE condition as raw SQL text applied to each new row + /// to decide whether to continue. `None` → run until fixed point. + pub condition: Option, + /// Maximum iterations before a `RecursionDepthExceeded` error is raised. + pub max_depth: usize, + /// `false` → UNION ALL (keep duplicates); `true` → UNION (deduplicate). + pub distinct: bool, +} diff --git a/nodedb-sql/src/types/plan/variants/timeseries.rs b/nodedb-sql/src/types/plan/variants/timeseries.rs new file mode 100644 index 000000000..17b9c4316 --- /dev/null +++ b/nodedb-sql/src/types/plan/variants/timeseries.rs @@ -0,0 +1,43 @@ +// SPDX-License-Identifier: Apache-2.0 + +//! Timeseries plan payloads: scan and ingest. + +use crate::temporal::TemporalScope; +use crate::types::filter::Filter; +use crate::types::query::{AggregateExpr, Projection, SortKey}; +use crate::types_expr::SqlValue; + +/// Payload of [`SqlPlan::TimeseriesScan`](crate::types::SqlPlan::TimeseriesScan). +#[derive(Debug, Clone)] +pub struct TimeseriesScanPlan { + pub collection: String, + pub time_range: (i64, i64), + pub bucket_interval_ms: i64, + pub group_by: Vec, + pub aggregates: Vec, + pub filters: Vec, + pub projection: Vec, + pub gap_fill: String, + pub limit: usize, + /// ORDER BY applied to the scan result. Empty = the engine's natural + /// order (ascending by the collection's time key). + pub sort_keys: Vec, + pub tiered: bool, + /// Bitemporal system-time / valid-time scope. Only non-default + /// on collections created `WITH BITEMPORAL`; `TimeseriesRules::plan_scan` + /// rejects temporal scopes otherwise. + pub temporal: TemporalScope, +} + +/// Payload of [`SqlPlan::TimeseriesIngest`](crate::types::SqlPlan::TimeseriesIngest). +#[derive(Debug, Clone)] +pub struct TimeseriesIngestPlan { + pub collection: String, + /// Defaults materialized and literals coerced, as in `Insert::rows`. + /// A row that omits the `TIME_KEY` column carries its declared + /// default here when one exists; only a row with no time value at + /// all takes the ingest clock. + pub rows: Vec>, + /// Mirrors `Insert::volatile_defaults`. + pub volatile_defaults: bool, +} diff --git a/nodedb-sql/src/types/plan/variants/vector_primary.rs b/nodedb-sql/src/types/plan/variants/vector_primary.rs new file mode 100644 index 000000000..c632a0249 --- /dev/null +++ b/nodedb-sql/src/types/plan/variants/vector_primary.rs @@ -0,0 +1,81 @@ +// SPDX-License-Identifier: Apache-2.0 + +//! Vector-primary collection write plan payloads. + +use crate::types::filter::Filter; +use crate::types_expr::{SqlExpr, SqlValue}; + +use super::super::row_types::{VectorPrimaryInsertIntent, VectorPrimaryRow}; + +/// Payload of [`SqlPlan::VectorPrimaryInsert`](crate::types::SqlPlan::VectorPrimaryInsert). +#[derive(Debug, Clone)] +pub struct VectorPrimaryInsertPlan { + pub collection: String, + /// Vector column name (matches `VectorPrimaryConfig::vector_field`). + /// Plumbed to the direct write op so the Data Plane keys its HNSW + /// index by `(tid, collection, field)` — the same key the SELECT + /// path uses. + pub field: String, + /// Collection-level quantization. Applied via `set_quantization` on + /// the first DirectUpsert so subsequent seals trigger codec-dispatch + /// rebuilds against the configured codec. + pub quantization: nodedb_types::VectorQuantization, + /// Native storage dtype for vector values (F32 / F16 / BF16). + pub storage_dtype: nodedb_types::VectorStorageDtype, + /// Payload field names that get equality bitmap indexes. Registered + /// via `payload.add_index` on the first DirectUpsert. + pub payload_indexes: Vec<(String, nodedb_types::PayloadIndexKind)>, + pub rows: Vec, + /// Whether any row carries a value materialized from a volatile + /// DEFAULT such as `nextval`. The value was evaluated while the plan + /// was built, so caching the lowered tasks would replay one + /// execution's value into every later one. + pub volatile_defaults: bool, + /// What an existing primary key means for each row. Mirrors + /// `KvInsert::intent`. + pub intent: VectorPrimaryInsertIntent, + /// `ON CONFLICT (pk) DO UPDATE SET field = expr` assignments, carried + /// when `intent == Upsert`. Empty means whole-row replace. + pub on_conflict_updates: Vec<(String, SqlExpr)>, + /// Resolved primary-key column name. See `Insert::primary_key`. + pub primary_key: Option, +} + +/// Payload of [`SqlPlan::VectorPrimaryDelete`](crate::types::SqlPlan::VectorPrimaryDelete). +#[derive(Debug, Clone)] +pub struct VectorPrimaryDeletePlan { + pub collection: String, + /// Vector column name; keys the HNSW index the rows live in. + pub field: String, + pub filters: Vec, + pub target_keys: Vec, + /// Resolved primary-key column name. See `Insert::primary_key`. + pub primary_key: Option, +} + +/// Payload of [`SqlPlan::VectorPrimaryTruncate`](crate::types::SqlPlan::VectorPrimaryTruncate). +#[derive(Debug, Clone)] +pub struct VectorPrimaryTruncatePlan { + pub collection: String, + /// Vector column name; keys the HNSW index the rows live in. + pub field: String, + pub restart_identity: bool, +} + +/// Payload of [`SqlPlan::VectorPrimaryUpdate`](crate::types::SqlPlan::VectorPrimaryUpdate). +#[derive(Debug, Clone)] +pub struct VectorPrimaryUpdatePlan { + pub collection: String, + /// Vector column name; keys the HNSW index the rows live in. + pub field: String, + pub quantization: nodedb_types::VectorQuantization, + pub storage_dtype: nodedb_types::VectorStorageDtype, + pub payload_indexes: Vec<(String, nodedb_types::PayloadIndexKind)>, + pub new_vector: Option>, + pub assignments: Vec<(String, SqlExpr)>, + pub filters: Vec, + pub target_keys: Vec, + pub returning: bool, + /// Resolved primary-key column name. See `Insert::primary_key`. + pub primary_key: Option, +} diff --git a/nodedb-sql/src/types/plan/variants/writes.rs b/nodedb-sql/src/types/plan/variants/writes.rs new file mode 100644 index 000000000..8f9be9af3 --- /dev/null +++ b/nodedb-sql/src/types/plan/variants/writes.rs @@ -0,0 +1,89 @@ +// SPDX-License-Identifier: Apache-2.0 + +//! Write-plan payloads: `INSERT`, key-value `INSERT`, `UPSERT`. + +use crate::types::query::EngineType; +use crate::types_expr::{SqlExpr, SqlValue}; + +use super::super::row_types::{KvInsertIntent, WriteRoute}; + +/// Payload of [`SqlPlan::Insert`](crate::types::SqlPlan::Insert). +#[derive(Debug, Clone)] +pub struct InsertPlan { + pub collection: String, + pub engine: EngineType, + /// The lowering these rows take, chosen by the engine's `EngineRules`. + /// The conversion layer reads it instead of re-deciding from `engine`. + pub route: WriteRoute, + /// Every declared DEFAULT already materialized and every literal + /// coerced to its declared column type by the planner. + pub rows: Vec>, + /// Whether a DEFAULT materialized into `rows` was volatile. The + /// planner evaluates declared defaults while building this plan, so a + /// cached plan would replay one execution's value; a volatile plan is + /// never cached. + pub volatile_defaults: bool, + /// `ON CONFLICT DO NOTHING` semantics: when true, duplicate-PK rows + /// are silently skipped instead of raising `unique_violation`. Plain + /// `INSERT` (no `ON CONFLICT` clause) sets this to `false`. + pub if_absent: bool, + /// Raw column type strings from the catalog: `(column_name, type_str)`. + /// Forwarded from `InsertParams::column_schema`. Used by columnar + /// converters to reconstruct the exact `ColumnType` for columns whose + /// `SqlDataType` is ambiguous (e.g. JSON and Bytes both map to Bytes). + pub column_schema: Vec<(String, String)>, + /// Declared `PRIMARY KEY` column name (if any), from + /// `InsertParams::primary_key`. Used by the conversion layer to + /// extract the document id from the correct column instead of + /// guessing at `id`/`document_id`/`key`. + pub primary_key: Option, +} + +/// Payload of [`SqlPlan::KvInsert`](crate::types::SqlPlan::KvInsert). +#[derive(Debug, Clone)] +pub struct KvInsertPlan { + pub collection: String, + pub entries: Vec<(SqlValue, Vec<(String, SqlValue)>)>, + /// TTL in seconds (0 = no expiry). Extracted from `ttl` column if present. + pub ttl_secs: u64, + /// INSERT-vs-UPSERT distinction. `KvOp::Put` is a Redis-SET-style + /// upsert by design; to honor SQL `INSERT` semantics the planner must + /// tell the converter whether a duplicate key should raise (plain + /// `INSERT`, `Insert`), be silently skipped (`ON CONFLICT DO NOTHING`, + /// `InsertIfAbsent`), or overwrite (`UPSERT` / `ON CONFLICT DO + /// UPDATE`, `Put`). + pub intent: KvInsertIntent, + /// `ON CONFLICT (key) DO UPDATE SET field = expr` assignments, carried + /// through when `intent == Put` via the ON-CONFLICT-DO-UPDATE path. + /// Empty for plain UPSERT (whole-value overwrite) and for INSERT + /// variants. + pub on_conflict_updates: Vec<(String, SqlExpr)>, + /// Whether any DEFAULT materialized into `entries` came from a + /// `Volatile` expression. The key-value planner evaluates declared + /// defaults while building this plan, so a cached plan would replay + /// one execution's value; a volatile plan is never cached. + pub volatile_defaults: bool, +} + +/// Payload of [`SqlPlan::Upsert`](crate::types::SqlPlan::Upsert). +#[derive(Debug, Clone)] +pub struct UpsertPlan { + pub collection: String, + pub engine: EngineType, + /// The lowering these rows take. Mirrors `Insert::route`. + pub route: WriteRoute, + /// Defaults materialized and literals coerced, as in `Insert::rows`. + pub rows: Vec>, + /// Mirrors `Insert::volatile_defaults`. + pub volatile_defaults: bool, + /// `ON CONFLICT (...) DO UPDATE SET field = expr` assignments. + /// When empty, upsert is a plain merge: new columns overwrite existing. + /// When non-empty, the engine applies these per-row against the + /// *existing* document instead of merging the inserted values. + pub on_conflict_updates: Vec<(String, SqlExpr)>, + /// Raw column type strings from the catalog: `(column_name, type_str)`. + /// Mirrors `Insert::column_schema` — see that field for rationale. + pub column_schema: Vec<(String, String)>, + /// Declared `PRIMARY KEY` column name (if any). See `Insert::primary_key`. + pub primary_key: Option, +} diff --git a/nodedb-sql/src/visitor/plan_visitor/dispatch.rs b/nodedb-sql/src/visitor/plan_visitor/dispatch.rs index b4fa81e56..eea3bc407 100644 --- a/nodedb-sql/src/visitor/plan_visitor/dispatch.rs +++ b/nodedb-sql/src/visitor/plan_visitor/dispatch.rs @@ -8,12 +8,17 @@ use super::args::{ AggregateVisitArgs, DocumentIndexLookupVisitArgs, HybridSearchTripleVisitArgs, HybridSearchVisitArgs, InsertVisitArgs, JoinVisitArgs, RecursiveScanVisitArgs, - RecursiveValueVisitArgs, ScanVisitArgs, SpatialScanVisitArgs, TimeseriesScanVisitArgs, - UpdateFromVisitArgs, UpsertVisitArgs, VectorSearchVisitArgs, + RecursiveValueVisitArgs, ScanVisitArgs, SpatialScanVisitArgs, TextSearchVisitArgs, + TimeseriesScanVisitArgs, UpdateFromVisitArgs, UpsertVisitArgs, VectorSearchVisitArgs, }; use super::dispatch_rest::dispatch_rest; use super::trait_def::PlanVisitor; use crate::types::SqlPlan; +use crate::types::{ + DocumentIndexLookupPlan, HybridSearchPlan, HybridSearchTriplePlan, InsertPlan, KvInsertPlan, + RangeScanPlan, RecursiveScanPlan, RecursiveValuePlan, TextSearchPlan, TimeseriesIngestPlan, + TimeseriesScanPlan, UpsertPlan, +}; pub fn dispatch(visitor: &mut V, plan: &SqlPlan) -> Result { match plan { @@ -53,7 +58,7 @@ pub fn dispatch(visitor: &mut V, plan: &SqlPlan) -> Result visitor.point_get(collection, alias.as_deref(), *engine, key_column, key_value), - SqlPlan::DocumentIndexLookup { + SqlPlan::DocumentIndexLookup(DocumentIndexLookupPlan { collection, alias, engine, @@ -68,7 +73,7 @@ pub fn dispatch(visitor: &mut V, plan: &SqlPlan) -> Result visitor.document_index_lookup(DocumentIndexLookupVisitArgs { + }) => visitor.document_index_lookup(DocumentIndexLookupVisitArgs { collection, alias: alias.as_deref(), engine: *engine, @@ -84,15 +89,15 @@ pub fn dispatch(visitor: &mut V, plan: &SqlPlan) -> Result visitor.range_scan(collection, field, lower.as_ref(), upper.as_ref(), *limit), - SqlPlan::Insert { + }) => visitor.range_scan(collection, field, lower.as_ref(), upper.as_ref(), *limit), + SqlPlan::Insert(InsertPlan { collection, engine, route, @@ -101,7 +106,7 @@ pub fn dispatch(visitor: &mut V, plan: &SqlPlan) -> Result visitor.insert(InsertVisitArgs { + }) => visitor.insert(InsertVisitArgs { collection, engine: *engine, route: *route, @@ -110,15 +115,15 @@ pub fn dispatch(visitor: &mut V, plan: &SqlPlan) -> Result visitor.kv_insert(collection, entries, *ttl_secs, *intent, on_conflict_updates), - SqlPlan::Upsert { + }) => visitor.kv_insert(collection, entries, *ttl_secs, *intent, on_conflict_updates), + SqlPlan::Upsert(UpsertPlan { collection, engine, route, @@ -127,7 +132,7 @@ pub fn dispatch(visitor: &mut V, plan: &SqlPlan) -> Result visitor.upsert(UpsertVisitArgs { + }) => visitor.upsert(UpsertVisitArgs { collection, engine: *engine, route: *route, @@ -225,7 +230,7 @@ pub fn dispatch(visitor: &mut V, plan: &SqlPlan) -> Result(visitor: &mut V, plan: &SqlPlan) -> Result visitor.timeseries_scan(TimeseriesScanVisitArgs { + }) => visitor.timeseries_scan(TimeseriesScanVisitArgs { collection, time_range: *time_range, bucket_interval_ms: *bucket_interval_ms, @@ -252,11 +257,11 @@ pub fn dispatch(visitor: &mut V, plan: &SqlPlan) -> Result visitor.timeseries_ingest(collection, rows), + }) => visitor.timeseries_ingest(collection, rows), SqlPlan::VectorSearch { collection, field, @@ -269,6 +274,7 @@ pub fn dispatch(visitor: &mut V, plan: &SqlPlan) -> Result visitor.vector_search(VectorSearchVisitArgs { collection, @@ -282,6 +288,7 @@ pub fn dispatch(visitor: &mut V, plan: &SqlPlan) -> Result(visitor: &mut V, plan: &SqlPlan) -> Result visitor.sparse_search(collection, field, query_entries, *top_k), - SqlPlan::TextSearch { + SqlPlan::TextSearch(TextSearchPlan { collection, - query, - top_k, + shape, filters, - score_alias, + scores, projection: _, - } => visitor.text_search(collection, query, *top_k, filters, score_alias.as_deref()), - SqlPlan::HybridSearch { + }) => visitor.text_search(TextSearchVisitArgs { + collection, + shape, + filters, + scores, + }), + SqlPlan::HybridSearch(HybridSearchPlan { collection, + vector_field, query_vector, + text_field, query_text, + filters, top_k, ef_search, vector_weight, + mode, fuzzy, score_alias, projection: _, - } => visitor.hybrid_search(HybridSearchVisitArgs { + }) => visitor.hybrid_search(HybridSearchVisitArgs { collection, + vector_field, query_vector, + text_field: text_field.as_deref(), query_text, + filters, top_k: *top_k, ef_search: *ef_search, vector_weight: *vector_weight, + mode: *mode, fuzzy: *fuzzy, score_alias: score_alias.as_deref(), }), - SqlPlan::HybridSearchTriple { + SqlPlan::HybridSearchTriple(HybridSearchTriplePlan { collection, + vector_field, query_vector, + text_field, query_text, + filters, graph_seed_id, graph_depth, graph_edge_label, top_k, ef_search, + mode, fuzzy, rrf_k, score_alias, projection: _, - } => visitor.hybrid_search_triple(HybridSearchTripleVisitArgs { + }) => visitor.hybrid_search_triple(HybridSearchTripleVisitArgs { collection, + vector_field, query_vector, + text_field: text_field.as_deref(), query_text, + filters, graph_seed_id, graph_depth: *graph_depth, graph_edge_label: graph_edge_label.as_deref(), top_k: *top_k, ef_search: *ef_search, + mode: *mode, fuzzy: *fuzzy, rrf_k: *rrf_k, score_alias: score_alias.as_deref(), @@ -370,7 +397,7 @@ pub fn dispatch(visitor: &mut V, plan: &SqlPlan) -> Result(visitor: &mut V, plan: &SqlPlan) -> Result visitor.recursive_scan(RecursiveScanVisitArgs { + }) => visitor.recursive_scan(RecursiveScanVisitArgs { collection, base_filters, recursive_filters, @@ -388,7 +415,7 @@ pub fn dispatch(visitor: &mut V, plan: &SqlPlan) -> Result(visitor: &mut V, plan: &SqlPlan) -> Result visitor.recursive_value(RecursiveValueVisitArgs { + }) => visitor.recursive_value(RecursiveValueVisitArgs { cte_name, columns, init_exprs, diff --git a/nodedb-sql/src/visitor/plan_visitor/dispatch_rest.rs b/nodedb-sql/src/visitor/plan_visitor/dispatch_rest.rs index e27b08462..1f049bfcc 100644 --- a/nodedb-sql/src/visitor/plan_visitor/dispatch_rest.rs +++ b/nodedb-sql/src/visitor/plan_visitor/dispatch_rest.rs @@ -13,6 +13,12 @@ use super::args::{ }; use super::trait_def::PlanVisitor; use crate::types::SqlPlan; +use crate::types::{ + AlterArrayPlan, ArrayAggPlan, ArrayElementwisePlan, ArrayProjectPlan, ArraySlicePlan, + CreateArrayPlan, CreateIndexPlan, CtePlan, DeleteArrayPlan, DropIndexPlan, InsertArrayPlan, + LateralLoopPlan, LateralTopKPlan, MergePlan, VectorPrimaryDeletePlan, VectorPrimaryInsertPlan, + VectorPrimaryTruncatePlan, VectorPrimaryUpdatePlan, +}; pub(super) fn dispatch_rest( visitor: &mut V, @@ -22,7 +28,7 @@ pub(super) fn dispatch_rest( SqlPlan::Union { inputs, distinct } => visitor.union(inputs, *distinct), SqlPlan::Intersect { left, right, all } => visitor.intersect(left, right, *all), SqlPlan::Except { left, right, all } => visitor.except(left, right, *all), - SqlPlan::Cte { definitions, outer } => visitor.cte(definitions, outer), + SqlPlan::Cte(CtePlan { definitions, outer }) => visitor.cte(definitions, outer), SqlPlan::Subquery { input, filters, @@ -42,7 +48,7 @@ pub(super) fn dispatch_rest( distinct: *distinct, limit: *limit, }), - SqlPlan::CreateArray { + SqlPlan::CreateArray(CreateArrayPlan { name, dims, attrs, @@ -52,7 +58,7 @@ pub(super) fn dispatch_rest( prefix_bits, audit_retain_ms, minimum_audit_retain_ms, - } => visitor.create_array(CreateArrayVisitArgs { + }) => visitor.create_array(CreateArrayVisitArgs { name, dims, attrs, @@ -64,40 +70,42 @@ pub(super) fn dispatch_rest( minimum_audit_retain_ms: *minimum_audit_retain_ms, }), SqlPlan::DropArray { name, if_exists } => visitor.drop_array(name, *if_exists), - SqlPlan::AlterArray { + SqlPlan::AlterArray(AlterArrayPlan { name, audit_retain_ms, minimum_audit_retain_ms, - } => visitor.alter_array(name, *audit_retain_ms, *minimum_audit_retain_ms), - SqlPlan::InsertArray { name, rows } => visitor.insert_array(name, rows), - SqlPlan::DeleteArray { name, coords } => visitor.delete_array(name, coords), - SqlPlan::ArraySlice { + }) => visitor.alter_array(name, *audit_retain_ms, *minimum_audit_retain_ms), + SqlPlan::InsertArray(InsertArrayPlan { name, rows }) => visitor.insert_array(name, rows), + SqlPlan::DeleteArray(DeleteArrayPlan { name, coords }) => { + visitor.delete_array(name, coords) + } + SqlPlan::ArraySlice(ArraySlicePlan { name, slice, attr_projection, limit, temporal, - } => visitor.array_slice(name, slice, attr_projection, *limit, temporal), - SqlPlan::ArrayProject { + }) => visitor.array_slice(name, slice, attr_projection, *limit, temporal), + SqlPlan::ArrayProject(ArrayProjectPlan { name, attr_projection, - } => visitor.array_project(name, attr_projection), - SqlPlan::ArrayAgg { + }) => visitor.array_project(name, attr_projection), + SqlPlan::ArrayAgg(ArrayAggPlan { name, attr, reducer, group_by_dim, temporal, - } => visitor.array_agg(name, attr, reducer, group_by_dim.as_deref(), temporal), - SqlPlan::ArrayElementwise { + }) => visitor.array_agg(name, attr, reducer, group_by_dim.as_deref(), temporal), + SqlPlan::ArrayElementwise(ArrayElementwisePlan { left, right, op, attr, - } => visitor.array_elementwise(left, right, *op, attr), + }) => visitor.array_elementwise(left, right, *op, attr), SqlPlan::ArrayFlush { name } => visitor.array_flush(name), SqlPlan::ArrayCompact { name } => visitor.array_compact(name), - SqlPlan::Merge { + SqlPlan::Merge(MergePlan { target, engine, source, @@ -106,7 +114,7 @@ pub(super) fn dispatch_rest( source_alias, clauses, returning, - } => visitor.merge(MergeVisitArgs { + }) => visitor.merge(MergeVisitArgs { target, engine: *engine, source, @@ -116,7 +124,7 @@ pub(super) fn dispatch_rest( clauses, returning: *returning, }), - SqlPlan::LateralTopK { + SqlPlan::LateralTopK(LateralTopKPlan { outer, outer_alias, inner_collection, @@ -127,7 +135,7 @@ pub(super) fn dispatch_rest( lateral_alias, projection, left_join, - } => visitor.lateral_top_k(LateralTopKVisitArgs { + }) => visitor.lateral_top_k(LateralTopKVisitArgs { outer, outer_alias: outer_alias.as_deref(), inner_collection, @@ -139,7 +147,7 @@ pub(super) fn dispatch_rest( projection, left_join: *left_join, }), - SqlPlan::LateralLoop { + SqlPlan::LateralLoop(LateralLoopPlan { outer, outer_alias, inner, @@ -148,7 +156,7 @@ pub(super) fn dispatch_rest( projection, outer_row_cap, left_join, - } => visitor.lateral_loop(LateralLoopVisitArgs { + }) => visitor.lateral_loop(LateralLoopVisitArgs { outer, outer_alias: outer_alias.as_deref(), inner, @@ -158,7 +166,7 @@ pub(super) fn dispatch_rest( outer_row_cap: *outer_row_cap, left_join: *left_join, }), - SqlPlan::VectorPrimaryInsert { + SqlPlan::VectorPrimaryInsert(VectorPrimaryInsertPlan { collection, field, quantization, @@ -170,7 +178,7 @@ pub(super) fn dispatch_rest( intent, on_conflict_updates, primary_key, - } => visitor.vector_primary_insert(VectorPrimaryInsertVisitArgs { + }) => visitor.vector_primary_insert(VectorPrimaryInsertVisitArgs { collection, field, quantization: *quantization, @@ -181,25 +189,25 @@ pub(super) fn dispatch_rest( on_conflict_updates, primary_key: primary_key.as_deref(), }), - SqlPlan::VectorPrimaryDelete { + SqlPlan::VectorPrimaryDelete(VectorPrimaryDeletePlan { collection, field, filters, target_keys, primary_key, - } => visitor.vector_primary_delete(VectorPrimaryDeleteVisitArgs { + }) => visitor.vector_primary_delete(VectorPrimaryDeleteVisitArgs { collection, field, filters, target_keys, primary_key: primary_key.as_deref(), }), - SqlPlan::VectorPrimaryTruncate { + SqlPlan::VectorPrimaryTruncate(VectorPrimaryTruncatePlan { collection, field, restart_identity, - } => visitor.vector_primary_truncate(collection, field, *restart_identity), - SqlPlan::VectorPrimaryUpdate { + }) => visitor.vector_primary_truncate(collection, field, *restart_identity), + SqlPlan::VectorPrimaryUpdate(VectorPrimaryUpdatePlan { collection, field, quantization, @@ -211,7 +219,7 @@ pub(super) fn dispatch_rest( target_keys, returning, primary_key, - } => visitor.vector_primary_update(VectorPrimaryUpdateVisitArgs { + }) => visitor.vector_primary_update(VectorPrimaryUpdateVisitArgs { collection, field, quantization: *quantization, @@ -224,14 +232,14 @@ pub(super) fn dispatch_rest( returning: *returning, primary_key: primary_key.as_deref(), }), - SqlPlan::CreateIndex { + SqlPlan::CreateIndex(CreateIndexPlan { index_name, collection, field, unique, if_not_exists, case_insensitive, - } => visitor.create_index( + }) => visitor.create_index( index_name.as_deref(), collection, field, @@ -239,11 +247,11 @@ pub(super) fn dispatch_rest( *if_not_exists, *case_insensitive, ), - SqlPlan::DropIndex { + SqlPlan::DropIndex(DropIndexPlan { index_name, collection, if_exists, - } => visitor.drop_index(index_name, collection.as_deref(), *if_exists), + }) => visitor.drop_index(index_name, collection.as_deref(), *if_exists), _ => unreachable!( "dispatch() already handles every remaining SqlPlan variant before forwarding here" ), diff --git a/nodedb-sql/tests/sql_suite/cases/limit_offset_bounds.rs b/nodedb-sql/tests/sql_suite/cases/limit_offset_bounds.rs index a6a8ff99e..d21d6001c 100644 --- a/nodedb-sql/tests/sql_suite/cases/limit_offset_bounds.rs +++ b/nodedb-sql/tests/sql_suite/cases/limit_offset_bounds.rs @@ -20,7 +20,7 @@ //! `LIMIT NULL` to derive the row description (see //! `fn prepared_limit_placeholder_null_plans_cleanly` below). -use nodedb_sql::types::{CollectionInfo, EngineType, SqlPlan}; +use nodedb_sql::types::{CollectionInfo, CtePlan, DocumentIndexLookupPlan, EngineType, SqlPlan}; use nodedb_sql::{SqlCatalog, SqlCatalogError, plan_sql}; use nodedb_types::DatabaseId; @@ -79,9 +79,9 @@ impl SqlCatalog for Catalog { fn bound_of(plan: &SqlPlan) -> String { match plan { SqlPlan::Scan { limit, offset, .. } - | SqlPlan::DocumentIndexLookup { limit, offset, .. } + | SqlPlan::DocumentIndexLookup(DocumentIndexLookupPlan { limit, offset, .. }) | SqlPlan::Subquery { limit, offset, .. } => format!("limit={limit:?} offset={offset}"), - SqlPlan::Cte { outer, .. } => format!("cte outer -> {}", bound_of(outer)), + SqlPlan::Cte(CtePlan { outer, .. }) => format!("cte outer -> {}", bound_of(outer)), other => format!("{other:?}"), } } diff --git a/nodedb-sql/tests/sql_suite/cases/positional_insert_column_binding.rs b/nodedb-sql/tests/sql_suite/cases/positional_insert_column_binding.rs index 30568f9a1..cc7f4e2b4 100644 --- a/nodedb-sql/tests/sql_suite/cases/positional_insert_column_binding.rs +++ b/nodedb-sql/tests/sql_suite/cases/positional_insert_column_binding.rs @@ -14,7 +14,9 @@ //! All tests use `plan_sql()` with a minimal catalog (mirrors //! `point_get_operand_order.rs` / `schema_qualified_rejection.rs`). -use nodedb_sql::types::{CollectionInfo, ColumnInfo, EngineType, SqlDataType}; +use nodedb_sql::types::{ + CollectionInfo, ColumnInfo, EngineType, InsertPlan, SqlDataType, UpsertPlan, +}; use nodedb_sql::{SqlCatalog, SqlCatalogError, SqlError, SqlPlan, SqlValue, plan_sql}; use nodedb_types::DatabaseId; @@ -118,8 +120,8 @@ fn plan_one(sql: &str) -> SqlPlan { fn insert_rows(plan: SqlPlan) -> Vec> { match plan { - SqlPlan::Insert { rows, .. } => rows, - SqlPlan::Upsert { rows, .. } => rows, + SqlPlan::Insert(InsertPlan { rows, .. }) => rows, + SqlPlan::Upsert(UpsertPlan { rows, .. }) => rows, other => panic!("expected an Insert or Upsert plan, got {other:?}"), } } diff --git a/nodedb-sql/tests/sql_suite/cases/truncate_engine_routing.rs b/nodedb-sql/tests/sql_suite/cases/truncate_engine_routing.rs index fd99a4253..afce7db54 100644 --- a/nodedb-sql/tests/sql_suite/cases/truncate_engine_routing.rs +++ b/nodedb-sql/tests/sql_suite/cases/truncate_engine_routing.rs @@ -7,7 +7,7 @@ //! with a typed error naming `DROP ARRAY`, and an unknown name is //! `UnknownTable`. -use nodedb_sql::types::{CollectionInfo, EngineType}; +use nodedb_sql::types::{CollectionInfo, EngineType, VectorPrimaryTruncatePlan}; use nodedb_sql::{SqlCatalog, SqlCatalogError, SqlError, SqlPlan, plan_sql}; use nodedb_types::DatabaseId; @@ -110,11 +110,11 @@ fn truncate_restart_identity_is_carried() { #[test] fn truncate_vector_primary_lowers_to_its_own_plan() { match plan_one("TRUNCATE vecs RESTART IDENTITY") { - SqlPlan::VectorPrimaryTruncate { + SqlPlan::VectorPrimaryTruncate(VectorPrimaryTruncatePlan { collection, field, restart_identity, - } => { + }) => { assert_eq!(collection, "vecs"); assert_eq!(field, "emb"); assert!(restart_identity); diff --git a/nodedb/src/control/planner/sql_plan_convert/convert_array_arms.rs b/nodedb/src/control/planner/sql_plan_convert/convert_array_arms.rs deleted file mode 100644 index 93bedfd75..000000000 --- a/nodedb/src/control/planner/sql_plan_convert/convert_array_arms.rs +++ /dev/null @@ -1,129 +0,0 @@ -// SPDX-License-Identifier: BUSL-1.1 - -use nodedb_sql::types::SqlPlan; - -use nodedb_physical::physical_task::PhysicalTask; -use crate::types::TenantId; - -use super::ConvertContext; - -/// Convert array-related `SqlPlan` variants to `PhysicalTask`s. -/// Called from `convert_one` in `convert.rs`. -pub(super) fn convert_array_plans( - plan: &SqlPlan, - tenant_id: TenantId, - ctx: &ConvertContext, -) -> Option>> { - match plan { - SqlPlan::CreateArray { - name, - dims, - attrs, - tile_extents, - cell_order, - tile_order, - prefix_bits, - audit_retain_ms, - minimum_audit_retain_ms, - } => Some(super::super::array_convert::convert_create_array( - super::super::array_convert::CreateArrayArgs { - name, - dims, - attrs, - tile_extents, - cell_order: *cell_order, - tile_order: *tile_order, - prefix_bits: *prefix_bits, - audit_retain_ms: *audit_retain_ms, - minimum_audit_retain_ms: *minimum_audit_retain_ms, - tenant_id, - ctx, - }, - )), - - SqlPlan::DropArray { name, if_exists } => Some( - super::super::array_convert::convert_drop_array(name, *if_exists, tenant_id, ctx), - ), - - SqlPlan::AlterArray { - name, - audit_retain_ms, - minimum_audit_retain_ms, - } => Some(super::super::array_alter_convert::convert_alter_array( - name, - *audit_retain_ms, - *minimum_audit_retain_ms, - tenant_id, - ctx, - )), - - SqlPlan::InsertArray { name, rows } => Some( - super::super::array_convert::convert_insert_array(name, rows, tenant_id, ctx), - ), - - SqlPlan::DeleteArray { name, coords } => Some( - super::super::array_convert::convert_delete_array(name, coords, tenant_id, ctx), - ), - - SqlPlan::ArraySlice { - name, - slice, - attr_projection, - limit, - temporal, - } => Some(super::super::array_fn_convert::convert_slice( - name, - slice, - attr_projection, - *limit, - *temporal, - tenant_id, - ctx, - )), - - SqlPlan::ArrayProject { - name, - attr_projection, - } => Some(super::super::array_fn_convert::convert_project( - name, - attr_projection, - tenant_id, - ctx, - )), - - SqlPlan::ArrayAgg { - name, - attr, - reducer, - group_by_dim, - temporal, - } => Some(super::super::array_fn_convert::convert_agg( - name, - attr, - *reducer, - group_by_dim.as_deref(), - *temporal, - tenant_id, - ctx, - )), - - SqlPlan::ArrayElementwise { - left, - right, - op, - attr, - } => Some(super::super::array_fn_convert::convert_elementwise( - left, right, *op, attr, tenant_id, ctx, - )), - - SqlPlan::ArrayFlush { name } => Some(super::super::array_fn_convert::convert_flush( - name, tenant_id, ctx, - )), - - SqlPlan::ArrayCompact { name } => Some(super::super::array_fn_convert::convert_compact( - name, tenant_id, ctx, - )), - - _ => None, - } -} diff --git a/nodedb/src/control/planner/sql_plan_convert/dml/merge.rs b/nodedb/src/control/planner/sql_plan_convert/dml/merge.rs index f5d83986c..6295ec511 100644 --- a/nodedb/src/control/planner/sql_plan_convert/dml/merge.rs +++ b/nodedb/src/control/planner/sql_plan_convert/dml/merge.rs @@ -2,7 +2,9 @@ //! Lower `SqlPlan::Merge` to `DocumentOp::Merge` physical task. -use nodedb_sql::types::{MergeClauseKind, MergePlanAction, MergePlanClause, SqlExpr, SqlPlan}; +use nodedb_sql::types::{ + DocumentIndexLookupPlan, MergeClauseKind, MergePlanAction, MergePlanClause, SqlExpr, SqlPlan, +}; use crate::bridge::envelope::PhysicalPlan; use crate::types::TenantId; @@ -54,7 +56,7 @@ pub(in super::super) fn convert_merge( SqlPlan::Scan { collection, .. } => { nodedb_types::QualifiedCollection::new(ctx.database_id, collection) } - SqlPlan::DocumentIndexLookup { collection, .. } => { + SqlPlan::DocumentIndexLookup(DocumentIndexLookupPlan { collection, .. }) => { nodedb_types::QualifiedCollection::new(ctx.database_id, collection) } other => { diff --git a/nodedb/src/control/planner/sql_plan_convert/dml/surrogate_keys.rs b/nodedb/src/control/planner/sql_plan_convert/dml/surrogate_keys.rs index 5eebac5aa..3b66b8eee 100644 --- a/nodedb/src/control/planner/sql_plan_convert/dml/surrogate_keys.rs +++ b/nodedb/src/control/planner/sql_plan_convert/dml/surrogate_keys.rs @@ -7,7 +7,11 @@ //! A key this module does not derive is still correct: conversion records it //! as a miss, and the bind step resolves it and converts the plan again. -use nodedb_sql::types::{EngineType, SqlPlan, SqlValue}; +use nodedb_sql::types::{ + CtePlan, EngineType, InsertArrayPlan, InsertPlan, KvInsertPlan, SqlPlan, SqlValue, + TimeseriesIngestPlan, UpsertPlan, VectorPrimaryDeletePlan, VectorPrimaryInsertPlan, + VectorPrimaryUpdatePlan, +}; use super::super::convert::ConvertContext; use super::super::value::{sql_value_to_bytes, sql_value_to_string}; @@ -55,7 +59,7 @@ fn collect_key_batches<'p>( batches: &mut Vec>, ) { match plan { - SqlPlan::Cte { definitions, outer } => { + SqlPlan::Cte(CtePlan { definitions, outer }) => { for (_, definition) in definitions { collect_key_batches(definition, ctx, batches); } @@ -88,24 +92,24 @@ fn plan_key_batch<'p>( ctx: &ConvertContext, ) -> crate::Result>> { let batch = match plan { - SqlPlan::Insert { + SqlPlan::Insert(InsertPlan { collection, rows, primary_key: Some(primary_key), .. - } - | SqlPlan::Upsert { + }) + | SqlPlan::Upsert(UpsertPlan { collection, rows, primary_key: Some(primary_key), .. - } => row_identity_batch(ctx, collection, primary_key, rows)?, - SqlPlan::VectorPrimaryInsert { + }) => row_identity_batch(ctx, collection, primary_key, rows)?, + SqlPlan::VectorPrimaryInsert(VectorPrimaryInsertPlan { collection, rows, primary_key: Some(primary_key), .. - } => { + }) => { let rows: Vec> = rows .iter() .map(|row| { @@ -117,11 +121,11 @@ fn plan_key_batch<'p>( .collect(); row_identity_batch(ctx, collection, primary_key, &rows)? } - SqlPlan::KvInsert { + SqlPlan::KvInsert(KvInsertPlan { collection, entries, .. - } => KeyBatch { + }) => KeyBatch { collection, pks: entries .iter() @@ -192,31 +196,31 @@ fn plan_key_batch<'p>( fresh: 0, }, }, - SqlPlan::VectorPrimaryDelete { + SqlPlan::VectorPrimaryDelete(VectorPrimaryDeletePlan { collection, target_keys, .. - } - | SqlPlan::VectorPrimaryUpdate { + }) + | SqlPlan::VectorPrimaryUpdate(VectorPrimaryUpdatePlan { collection, target_keys, .. - } => KeyBatch { + }) => KeyBatch { collection, pks: document_keys(target_keys), binds: false, fresh: 0, }, // Every timeseries row takes a fresh identity. - SqlPlan::TimeseriesIngest { + SqlPlan::TimeseriesIngest(TimeseriesIngestPlan { collection, rows, .. - } => KeyBatch { + }) => KeyBatch { collection, pks: Vec::new(), binds: true, fresh: rows.len(), }, - SqlPlan::InsertArray { name, rows } => KeyBatch { + SqlPlan::InsertArray(InsertArrayPlan { name, rows }) => KeyBatch { collection: name, pks: super::super::array_convert::insert_array_cell_pks( name, @@ -275,7 +279,9 @@ fn kv_keys(target_keys: &[SqlValue]) -> Vec> { #[cfg(test)] mod tests { - use nodedb_sql::types::{EngineType, SqlPlan, SqlValue, WriteRoute}; + use nodedb_sql::types::{ + CtePlan, EngineType, InsertPlan, SqlPlan, SqlValue, TimeseriesIngestPlan, WriteRoute, + }; use super::plan_key_batches; use crate::control::planner::sql_plan_convert::{ConvertContext, PlanningPurpose}; @@ -313,7 +319,7 @@ mod tests { } fn insert(collection: &str, primary_key: &str, rows: Vec>) -> SqlPlan { - SqlPlan::Insert { + SqlPlan::Insert(InsertPlan { collection: collection.to_string(), engine: EngineType::DocumentSchemaless, route: WriteRoute::Document, @@ -322,7 +328,7 @@ mod tests { if_absent: false, column_schema: Vec::new(), primary_key: Some(primary_key.to_string()), - } + }) } fn keys(pks: &[Vec]) -> Vec<&[u8]> { @@ -393,10 +399,10 @@ mod tests { #[test] fn cte_outer_write_is_collected() { let ctx = ctx(); - let plans = vec![SqlPlan::Cte { + let plans = vec![SqlPlan::Cte(CtePlan { definitions: Vec::new(), outer: Box::new(insert("users", "id", vec![row(Some("c"))])), - }]; + })]; let batches = plan_key_batches(&plans, &ctx); assert_eq!(batches.len(), 1); assert_eq!(keys(&batches[0].pks), vec![&b"c"[..]]); @@ -434,11 +440,11 @@ mod tests { #[test] fn timeseries_ingest_draws_one_fresh_identity_per_row() { let ctx = ctx(); - let plans = vec![SqlPlan::TimeseriesIngest { + let plans = vec![SqlPlan::TimeseriesIngest(TimeseriesIngestPlan { collection: "metrics".to_string(), rows: vec![Vec::new(), Vec::new(), Vec::new()], volatile_defaults: false, - }]; + })]; let batches = plan_key_batches(&plans, &ctx); assert_eq!(batches.len(), 1); assert_eq!(batches[0].collection, "metrics"); diff --git a/nodedb/src/control/planner/sql_plan_convert/dml/update_delete/update_from.rs b/nodedb/src/control/planner/sql_plan_convert/dml/update_delete/update_from.rs index 602f6f6d0..bc3d88ebd 100644 --- a/nodedb/src/control/planner/sql_plan_convert/dml/update_delete/update_from.rs +++ b/nodedb/src/control/planner/sql_plan_convert/dml/update_delete/update_from.rs @@ -6,7 +6,7 @@ //! Assignments are converted with table-qualified column references so the Data //! Plane can resolve `src.col` against the merged `{target + "src.col": ...}` doc. -use nodedb_sql::types::{Filter, SqlExpr, SqlPlan}; +use nodedb_sql::types::{DocumentIndexLookupPlan, Filter, SqlExpr, SqlPlan}; use crate::bridge::envelope::PhysicalPlan; use crate::types::TenantId; @@ -62,9 +62,9 @@ pub(in crate::control::planner::sql_plan_convert) fn convert_update_from( let alias_str = alias.as_deref().unwrap_or(collection.as_str()).to_string(); (qualified, alias_str) } - SqlPlan::DocumentIndexLookup { + SqlPlan::DocumentIndexLookup(DocumentIndexLookupPlan { collection, alias, .. - } => { + }) => { let qualified = nodedb_types::QualifiedCollection::new(ctx.database_id, collection); let alias_str = alias.as_deref().unwrap_or(collection.as_str()).to_string(); (qualified, alias_str) diff --git a/nodedb/src/control/planner/sql_plan_convert/lateral.rs b/nodedb/src/control/planner/sql_plan_convert/lateral.rs index 4d5508663..0835595a5 100644 --- a/nodedb/src/control/planner/sql_plan_convert/lateral.rs +++ b/nodedb/src/control/planner/sql_plan_convert/lateral.rs @@ -6,7 +6,7 @@ //! inside the `QueryOp` so the Data Plane executor can materialise outer rows //! in-process before iterating over them. -use nodedb_sql::types::{Filter, Projection, SortKey, SqlPlan}; +use nodedb_sql::types::{DocumentIndexLookupPlan, Filter, Projection, SortKey, SqlPlan}; use crate::bridge::envelope::PhysicalPlan; use crate::types::TenantId; @@ -160,7 +160,7 @@ pub(super) fn convert_lateral_loop( pub(super) fn collection_name_from_plan(plan: &SqlPlan) -> Option { match plan { SqlPlan::Scan { collection, .. } - | SqlPlan::DocumentIndexLookup { collection, .. } + | SqlPlan::DocumentIndexLookup(DocumentIndexLookupPlan { collection, .. }) | SqlPlan::PointGet { collection, .. } => Some(collection.clone()), _ => None, } @@ -169,7 +169,8 @@ pub(super) fn collection_name_from_plan(plan: &SqlPlan) -> Option { /// Extract base filters from a scan-like SqlPlan. fn inner_filters_from_plan(plan: &SqlPlan) -> crate::Result> { match plan { - SqlPlan::Scan { filters, .. } | SqlPlan::DocumentIndexLookup { filters, .. } => { + SqlPlan::Scan { filters, .. } + | SqlPlan::DocumentIndexLookup(DocumentIndexLookupPlan { filters, .. }) => { serialize_filters(filters) } _ => Ok(Vec::new()), diff --git a/nodedb/src/control/planner/sql_plan_convert/surrogate_prefetch/bind.rs b/nodedb/src/control/planner/sql_plan_convert/surrogate_prefetch/bind.rs index 0087272a9..f92706054 100644 --- a/nodedb/src/control/planner/sql_plan_convert/surrogate_prefetch/bind.rs +++ b/nodedb/src/control/planner/sql_plan_convert/surrogate_prefetch/bind.rs @@ -89,7 +89,7 @@ mod tests { use futures::future::BoxFuture; use nodedb_physical::physical_plan::{KvOp, TimeseriesOp}; use nodedb_physical::physical_task::PhysicalTask; - use nodedb_sql::types::{EngineType, SqlExpr, SqlPlan, SqlValue}; + use nodedb_sql::types::{EngineType, SqlExpr, SqlPlan, SqlValue, TimeseriesIngestPlan}; use nodedb_types::{CollectionKey, Surrogate}; use super::{convert_bound, convert_resolving_misses}; @@ -135,11 +135,11 @@ mod tests { } fn timeseries_ingest(rows: usize) -> Vec { - vec![SqlPlan::TimeseriesIngest { + vec![SqlPlan::TimeseriesIngest(TimeseriesIngestPlan { collection: "metrics".to_string(), rows: vec![vec![("value".to_string(), SqlValue::Int(1))]; rows], volatile_defaults: false, - }] + })] } fn ingest_surrogates(tasks: &[PhysicalTask]) -> Vec { diff --git a/nodedb/src/control/server/shared/returning/clause.rs b/nodedb/src/control/server/shared/returning/clause.rs index 22689d7fc..430f7af80 100644 --- a/nodedb/src/control/server/shared/returning/clause.rs +++ b/nodedb/src/control/server/shared/returning/clause.rs @@ -14,9 +14,12 @@ use nodedb_physical::physical_plan::{ReturningColumns, ReturningItem, ReturningS use crate::Error; use crate::control::planner::plan_error_map::map_plan_error; use nodedb_sql::catalog::SqlCatalog; -use nodedb_sql::types::SqlPlan; use nodedb_sql::types::plan::referenced_columns; use nodedb_sql::types::query::Projection; +use nodedb_sql::types::{ + InsertPlan, KvInsertPlan, MergePlan, SqlPlan, TimeseriesIngestPlan, UpsertPlan, + VectorPrimaryDeletePlan, VectorPrimaryInsertPlan, VectorPrimaryUpdatePlan, +}; /// A resolved RETURNING clause: the Control-Plane projection and the /// Data-Plane spec derived from it. @@ -98,17 +101,19 @@ fn spec_from_projection(projection: &[Projection]) -> ReturningSpec { pub fn returning_target_collection(plans: &[SqlPlan]) -> Option { let plan = plans.first()?; match plan { - SqlPlan::Insert { collection, .. } - | SqlPlan::KvInsert { collection, .. } - | SqlPlan::Upsert { collection, .. } + SqlPlan::Insert(InsertPlan { collection, .. }) + | SqlPlan::KvInsert(KvInsertPlan { collection, .. }) + | SqlPlan::Upsert(UpsertPlan { collection, .. }) | SqlPlan::Update { collection, .. } | SqlPlan::UpdateFrom { collection, .. } | SqlPlan::Delete { collection, .. } - | SqlPlan::TimeseriesIngest { collection, .. } - | SqlPlan::VectorPrimaryInsert { collection, .. } - | SqlPlan::VectorPrimaryDelete { collection, .. } - | SqlPlan::VectorPrimaryUpdate { collection, .. } => Some(collection.clone()), - SqlPlan::Merge { target, .. } | SqlPlan::InsertSelect { target, .. } => { + | SqlPlan::TimeseriesIngest(TimeseriesIngestPlan { collection, .. }) + | SqlPlan::VectorPrimaryInsert(VectorPrimaryInsertPlan { collection, .. }) + | SqlPlan::VectorPrimaryDelete(VectorPrimaryDeletePlan { collection, .. }) + | SqlPlan::VectorPrimaryUpdate(VectorPrimaryUpdatePlan { collection, .. }) => { + Some(collection.clone()) + } + SqlPlan::Merge(MergePlan { target, .. }) | SqlPlan::InsertSelect { target, .. } => { Some(target.clone()) } SqlPlan::ConstantResult { .. } diff --git a/nodedb/tests/inproc/cases/executor_tests/test_group_by_alias.rs b/nodedb/tests/inproc/cases/executor_tests/test_group_by_alias.rs index e76a9376d..bda9a3f86 100644 --- a/nodedb/tests/inproc/cases/executor_tests/test_group_by_alias.rs +++ b/nodedb/tests/inproc/cases/executor_tests/test_group_by_alias.rs @@ -8,7 +8,7 @@ use nodedb::control::planner::sql_plan_convert::{ConvertContext, convert}; use nodedb_physical::physical_plan::{PhysicalPlan, TimeseriesOp}; -use nodedb_sql::types::{CollectionInfo, EngineType, SqlCatalog, SqlPlan}; +use nodedb_sql::types::{CollectionInfo, EngineType, SqlCatalog, SqlPlan, TimeseriesScanPlan}; use nodedb_types::QualifiedCollection; use super::helpers::*; @@ -152,9 +152,9 @@ fn plan_group_by_full_expression_has_bucket_interval() { "SELECT time_bucket('1 hour', timestamp) AS b, COUNT(*) FROM dns_bench \ GROUP BY time_bucket('1 hour', timestamp)", ); - let SqlPlan::TimeseriesScan { + let SqlPlan::TimeseriesScan(TimeseriesScanPlan { bucket_interval_ms, .. - } = plan + }) = plan else { panic!("expected TimeseriesScan, got {plan:?}"); }; @@ -170,9 +170,9 @@ fn plan_group_by_alias_has_bucket_interval() { let plan = plan_sql( "SELECT time_bucket('1 hour', timestamp) AS b, COUNT(*) FROM dns_bench GROUP BY b", ); - let SqlPlan::TimeseriesScan { + let SqlPlan::TimeseriesScan(TimeseriesScanPlan { bucket_interval_ms, .. - } = plan + }) = plan else { panic!("expected TimeseriesScan, got {plan:?}"); }; @@ -188,9 +188,9 @@ fn plan_group_by_positional_has_bucket_interval() { let plan = plan_sql( "SELECT time_bucket('1 hour', timestamp) AS b, COUNT(*) FROM dns_bench GROUP BY 1", ); - let SqlPlan::TimeseriesScan { + let SqlPlan::TimeseriesScan(TimeseriesScanPlan { bucket_interval_ms, .. - } = plan + }) = plan else { panic!("expected TimeseriesScan, got {plan:?}"); }; From fb93eaa21c562fdabd93b351cba4fecc4a2396bd Mon Sep 17 00:00:00 2001 From: Farhan Syah Date: Thu, 8 Oct 2026 16:51:47 +0800 Subject: [PATCH 09/24] fix(document): render row identity under the declared key column A collection keyed by a declared primary key renders a row's identity under that column on every scan, filter, RLS image and RETURNING path, never as an extra id beside it. RLS injection resolves the identity column through the catalog. An assignment to the key column (UPDATE SET, ON CONFLICT DO UPDATE, MERGE, native document put) is refused unless it keeps the row's key. A WHERE on any declared key is a point lookup, a computed item over it is evaluated on the Control Plane, and to_jsonb(*) returns the whole row. --- nodedb-query/src/expr/eval.rs | 41 ++- nodedb-query/src/expr/mod.rs | 2 +- nodedb-query/src/expr/types.rs | 7 + nodedb-query/src/functions/json/dispatch.rs | 7 + nodedb-sql/src/functions/arg_types.rs | 10 +- .../src/functions/builtins/scalars/pg_json.rs | 10 + nodedb-sql/src/functions/registry.rs | 1 + nodedb-sql/src/lib.rs | 2 +- nodedb-sql/src/optimizer/pipeline.rs | 6 +- nodedb-sql/src/optimizer/point_get.rs | 347 +++++++++++++++--- nodedb-sql/src/planner/cp_projection.rs | 33 +- nodedb-sql/src/planner/returning.rs | 34 +- nodedb-types/src/id/qualified_collection.rs | 18 + nodedb-types/src/row_identity.rs | 13 + .../control/planner/catalog_adapter/mod.rs | 4 +- .../control/planner/context/query/planning.rs | 17 +- .../control/planner/rls_injection/context.rs | 81 ++-- .../src/control/planner/rls_injection/plan.rs | 12 +- .../sql_plan_convert/dml/key_assignment.rs | 117 ++++++ .../planner/sql_plan_convert/dml/kv_insert.rs | 70 +++- .../planner/sql_plan_convert/dml/mod.rs | 1 + .../dml/update_delete/update.rs | 36 +- .../planner/sql_plan_convert/dml/upsert.rs | 22 +- .../sql_plan_convert/dml/vector_primary.rs | 5 + .../sql_plan_convert/expr/bridge_expr.rs | 3 +- .../server/http/routes/promql/remote.rs | 1 + .../src/control/server/ilp_batch/dispatch.rs | 7 +- .../server/native/dispatch/direct_ops.rs | 1 + .../server/native/dispatch/graph_match.rs | 1 + .../native/dispatch/plan_builder/document.rs | 33 +- .../plan_builder/document_identity.rs | 239 ++++++++++++ .../server/pgwire/handler/routing/execute.rs | 1 + .../control/server/resp/gateway_dispatch.rs | 1 + .../server/response_shape/compose/kernel.rs | 57 ++- .../response_shape/compose/materialized.rs | 7 +- .../control/server/response_shape/project.rs | 82 ++++- .../server/response_shape/stamp/mod.rs | 2 + .../server/response_shape/stamp/no_access.rs | 59 +++ .../server/response_translate/text_hybrid.rs | 61 ++- .../server/shared/clone_write/probes.rs | 1 + .../ddl/neutral/query_functions/helpers.rs | 11 +- .../server/shared/ddl/neutral/read_gate.rs | 1 + .../server/shared/ddl/user_dispatch.rs | 1 + nodedb/src/control/server/sync/kv_handler.rs | 1 + nodedb/src/control/trigger/dml_hook.rs | 1 + .../data/executor/core_loop/filter_match.rs | 42 ++- .../handlers/bulk_dml/update_project.rs | 22 +- .../handlers/control/calvin_reply/images.rs | 15 +- .../handlers/control/range_scan_versioned.rs | 57 +-- .../executor/handlers/document/read/fetch.rs | 14 +- .../document/read/materialize_scan.rs | 21 +- .../handlers/document/resolve/apply.rs | 5 + .../handlers/document/resolve/bulk.rs | 8 +- .../handlers/document/resolve/context.rs | 22 +- .../handlers/document/resolve/point.rs | 13 +- .../handlers/document/resolve/upsert.rs | 3 +- .../data/executor/handlers/identity_guard.rs | 205 +++++++++++ nodedb/src/data/executor/handlers/kv/rls.rs | 12 +- .../merge_orchestrated/apply/insert_rows.rs | 3 +- .../merge_orchestrated/apply/orchestrate.rs | 25 +- .../handlers/merge_orchestrated/plan.rs | 36 +- .../src/data/executor/handlers/point/get.rs | 33 +- .../executor/handlers/point/update/exec.rs | 14 +- .../src/data/executor/handlers/recursive.rs | 11 +- .../data/executor/handlers/returning_doc.rs | 46 ++- .../data/executor/handlers/returning_rows.rs | 24 +- .../executor/handlers/spatial/full_scan.rs | 25 +- .../executor/handlers/spatial/rtree_scan.rs | 32 +- .../data/executor/handlers/spatial_refine.rs | 41 ++- .../handlers/transaction/overlay/merge.rs | 10 +- .../transaction/overlay/spatial_merge.rs | 45 +-- .../overlay/vector_primary_merge.rs | 22 +- .../transaction/stage_write/dispatch.rs | 1 + .../stage_write/stage_returning.rs | 24 +- .../stage_write/stage_vector_targets.rs | 9 +- .../handlers/update_from_join_collect.rs | 31 +- .../handlers/update_from_join_write.rs | 24 +- .../executor/handlers/vector_direct_delete.rs | 17 +- .../vector_direct_resolve/resolve_delete.rs | 4 +- .../vector_direct_resolve/resolve_update.rs | 4 +- .../vector_direct_resolve/resolve_upsert.rs | 10 +- .../handlers/vector_direct_targets.rs | 12 +- .../executor/handlers/vector_direct_update.rs | 17 +- .../data/executor/handlers/vector_upsert.rs | 10 +- nodedb/src/data/executor/identity_column.rs | 38 ++ nodedb/src/data/executor/mod.rs | 1 + nodedb/src/data/executor/row_shape.rs | 237 +++++++----- nodedb/src/data/executor/scan_normalize.rs | 150 ++++---- nodedb/src/data/executor/scan_versioned.rs | 19 +- .../cases/native_document_identity_update.rs | 104 ++++++ .../wire/cases/declared_key_scan_rows.rs | 96 +++++ .../sql_on_conflict_primary_key_identity.rs | 219 +++++++++++ .../cases/sql_update_primary_key_identity.rs | 161 ++++++++ 93 files changed, 2780 insertions(+), 681 deletions(-) create mode 100644 nodedb/src/control/planner/sql_plan_convert/dml/key_assignment.rs create mode 100644 nodedb/src/control/server/native/dispatch/plan_builder/document_identity.rs create mode 100644 nodedb/src/control/server/response_shape/stamp/no_access.rs create mode 100644 nodedb/src/data/executor/handlers/identity_guard.rs create mode 100644 nodedb/src/data/executor/identity_column.rs create mode 100644 nodedb/tests/native/cases/native_document_identity_update.rs create mode 100644 nodedb/tests/wire/cases/declared_key_scan_rows.rs create mode 100644 nodedb/tests/wire/cases/sql_on_conflict_primary_key_identity.rs create mode 100644 nodedb/tests/wire/cases/sql_update_primary_key_identity.rs 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-sql/src/functions/arg_types.rs b/nodedb-sql/src/functions/arg_types.rs index 652aa8f17..c07ddafe4 100644 --- a/nodedb-sql/src/functions/arg_types.rs +++ b/nodedb-sql/src/functions/arg_types.rs @@ -17,10 +17,7 @@ const ANY: &[ColumnType] = &[]; const NUMERIC: &[ColumnType] = &[ ColumnType::Int64, ColumnType::Float64, - ColumnType::Decimal { - precision: 38, - scale: 10, - }, + ColumnType::Decimal(None), ]; const TEXT: &[ColumnType] = &[ColumnType::String]; @@ -120,7 +117,7 @@ pub static BM25_SCORE_ARGS: &[ArgTypeSpec] = &[any("column"), typed("query", TEX pub static SEARCH_SCORE_ARGS: &[ArgTypeSpec] = &[any("column"), typed("query", TEXT)]; -pub static TEXT_MATCH_ARGS: &[ArgTypeSpec] = &[any("column"), typed("query", TEXT), any("options")]; +pub static TEXT_MATCH_ARGS: &[ArgTypeSpec] = &[any("column"), typed("query", TEXT)]; // ── Hybrid search ───────────────────────────────────────────────────────────── @@ -392,6 +389,9 @@ pub static JSON_CONTAINS_ARGS: &[ArgTypeSpec] = &[any("container"), any("needle" /// `json_merge(base, overlay)` / `json_patch(base, overlay)`. pub static JSON_MERGE_ARGS: &[ArgTypeSpec] = &[any("base"), any("overlay")]; +/// `to_jsonb(value)` — any value, or `*` for the whole row. +pub static TO_JSONB_ARGS: &[ArgTypeSpec] = &[any("value")]; + // ── Sequence accessors ─────────────────────────────────────────────────────── /// `nextval(name)` / `currval(name)` — one sequence name. diff --git a/nodedb-sql/src/functions/builtins/scalars/pg_json.rs b/nodedb-sql/src/functions/builtins/scalars/pg_json.rs index ec51faf73..531d9a644 100644 --- a/nodedb-sql/src/functions/builtins/scalars/pg_json.rs +++ b/nodedb-sql/src/functions/builtins/scalars/pg_json.rs @@ -119,5 +119,15 @@ pub(super) fn pg_json_functions() -> Vec { Some(ColumnType::Bool), arg_types::JSON_EXISTS_ARGS, ), + // `to_jsonb(v)` is `v` as JSON. `to_jsonb(*)` is the whole row. + m( + "to_jsonb", + Scalar, + 1, + 1, + no_trigger(), + Some(ColumnType::Json), + arg_types::TO_JSONB_ARGS, + ), ] } diff --git a/nodedb-sql/src/functions/registry.rs b/nodedb-sql/src/functions/registry.rs index 8284a2ed1..9b681a511 100644 --- a/nodedb-sql/src/functions/registry.rs +++ b/nodedb-sql/src/functions/registry.rs @@ -440,6 +440,7 @@ mod tests { "json_contains", "json_merge", "json_patch", + "to_jsonb", ], ); diff --git a/nodedb-sql/src/lib.rs b/nodedb-sql/src/lib.rs index 0f2ee0ad0..e7071703f 100644 --- a/nodedb-sql/src/lib.rs +++ b/nodedb-sql/src/lib.rs @@ -141,7 +141,7 @@ fn plan_statements( StatementKind::Select(query) => { let plan = planner::select::plan_statement_query(query, catalog, &functions, temporal)?; - let plan = optimizer::optimize(plan, catalog); + let plan = optimizer::optimize(plan, catalog)?; plans.push(plan); } StatementKind::Insert(ins) => { diff --git a/nodedb-sql/src/optimizer/pipeline.rs b/nodedb-sql/src/optimizer/pipeline.rs index 2b7249693..403d38734 100644 --- a/nodedb-sql/src/optimizer/pipeline.rs +++ b/nodedb-sql/src/optimizer/pipeline.rs @@ -8,7 +8,7 @@ use crate::types::SqlPlan; use super::{point_get, predicate_pushdown}; /// Apply all optimization passes to a plan. -pub fn optimize(plan: SqlPlan, catalog: &dyn SqlCatalog) -> SqlPlan { - let plan = point_get::optimize(plan, catalog); - predicate_pushdown::optimize(plan) +pub fn optimize(plan: SqlPlan, catalog: &dyn SqlCatalog) -> crate::Result { + let plan = point_get::optimize(plan, catalog)?; + Ok(predicate_pushdown::optimize(plan)) } diff --git a/nodedb-sql/src/optimizer/point_get.rs b/nodedb-sql/src/optimizer/point_get.rs index fbc446dc7..0c4df8853 100644 --- a/nodedb-sql/src/optimizer/point_get.rs +++ b/nodedb-sql/src/optimizer/point_get.rs @@ -5,6 +5,7 @@ use nodedb_types::DatabaseId; use crate::catalog::SqlCatalog; +use crate::planner::cp_projection::to_cp_computed; use crate::types::*; /// If a Scan has a single equality filter on the collection's primary key, @@ -15,9 +16,17 @@ use crate::types::*; /// explicit `PRIMARY KEY` gets an auto-generated `_rowid` key, so a filter on a /// regular `id` column is NOT a point lookup and must stay a scan (routing it /// to PointGet would resolve a surrogate for the wrong key and return zero -/// rows). When the catalog cannot resolve a primary key (unknown collection), -/// fall back to the legacy convention so nothing regresses. -pub fn optimize(plan: SqlPlan, catalog: &dyn SqlCatalog) -> SqlPlan { +/// rows). A declared key of any name is a point lookup: `WHERE sku = 'p1'` +/// reads the one row keyed `p1`, as `WHERE id = 'p1'` does. When the catalog +/// cannot resolve a primary key (unknown collection), the conventional names +/// apply. A catalog error fails the plan: `RetryableSchemaChanged` must reach +/// the caller. +/// +/// A point lookup returns the whole stored row and runs no expression. A +/// computed SELECT item over it becomes [`Projection::CpComputed`]: the +/// Control Plane evaluates it once on the returned row. `SELECT to_jsonb(*) +/// ... WHERE = $1` stays a point lookup that way. +pub fn optimize(plan: SqlPlan, catalog: &dyn SqlCatalog) -> crate::Result { match plan { SqlPlan::Scan { ref collection, @@ -27,86 +36,308 @@ pub fn optimize(plan: SqlPlan, catalog: &dyn SqlCatalog) -> SqlPlan { ref projection, ref temporal, .. - } if filters.len() == 1 - && !temporal.is_temporal() - && !projection.iter().any(|p| { - matches!( - p, - Projection::Computed { .. } | Projection::CpComputed { .. } - ) - }) => - { + } if filters.len() == 1 && !temporal.is_temporal() && has_point_lookup(engine) => { let pk = catalog - .get_collection(DatabaseId::DEFAULT, collection) - .ok() - .flatten() + .get_collection(DatabaseId::DEFAULT, collection)? .and_then(|info| info.primary_key); if let Some((key_col, key_val)) = extract_pk_equality(&filters[0], pk.as_deref()) { - return SqlPlan::PointGet { + return Ok(SqlPlan::PointGet { collection: collection.clone(), alias: alias.clone(), engine: *engine, key_column: key_col, key_value: key_val, - projection: projection.clone(), - }; + projection: projection.iter().cloned().map(to_cp_computed).collect(), + }); } - plan + Ok(plan) } - _ => plan, + _ => Ok(plan), + } +} + +/// Whether the engine answers a key equality with a point lookup. +/// +/// Timeseries and array refuse point lookups in their engine rules. A key +/// equality on them stays a scan. +fn has_point_lookup(engine: &EngineType) -> bool { + match engine { + EngineType::DocumentSchemaless + | EngineType::DocumentStrict + | EngineType::KeyValue + | EngineType::Columnar + | EngineType::Spatial => true, + EngineType::Timeseries | EngineType::Array => false, } } /// Extract a simple equality filter eligible for the point-get rewrite. /// -/// Only the conventional document-key columns (`id` / `document_id` / `key`) -/// are candidates — this pass deliberately does not promote arbitrary declared -/// primary keys (e.g. `sku`) to point lookups. The refinement over the legacy -/// name-only check: when the catalog resolves a real primary key, a candidate -/// column is eligible only if it actually IS that key. This excludes a regular -/// `id` column on a collection whose real key is the auto-generated `_rowid` -/// (where a point-get would resolve a surrogate for the wrong key and return -/// zero rows) while leaving every previously-optimized case unchanged. When the -/// catalog can't resolve a key, fall back to the name-only convention. +/// When the catalog resolves the primary key, the candidate column is that +/// key, whatever its name. A regular `id` column on a collection keyed by +/// another column stays a scan: a point lookup would resolve a surrogate for +/// the wrong key and return zero rows. When the catalog resolves no key, the +/// conventional document-key names (`id` / `document_id` / `key`) are the +/// candidates. fn extract_pk_equality(filter: &Filter, pk: Option<&str>) -> Option<(String, SqlValue)> { - let is_pk_column = |col: &str| -> bool { - let conventional = col == "id" || col == "document_id" || col == "key"; - match pk { - Some(pk) => conventional && col.eq_ignore_ascii_case(pk), - None => conventional, - } - }; - match &filter.expr { + let (column, value) = match &filter.expr { FilterExpr::Comparison { field, op: CompareOp::Eq, value, - } => { - let f = field.to_lowercase(); - if is_pk_column(&f) { - Some((f, value.clone())) - } else { - None - } - } + } => (field.as_str(), value), FilterExpr::Expr(SqlExpr::BinaryOp { left, op: BinaryOp::Eq, right, - }) => { - let col = match left.as_ref() { - SqlExpr::Column { name, .. } => name.to_lowercase(), - _ => return None, - }; - if !is_pk_column(&col) { - return None; + }) => match (left.as_ref(), right.as_ref()) { + (SqlExpr::Column { name, .. }, SqlExpr::Literal(value)) => (name.as_str(), value), + _ => return None, + }, + _ => return None, + }; + key_column_named(column, pk).map(|key| (key, value.clone())) +} + +/// The key column `column` names, spelled as the catalog declares it, or +/// `None` when `column` is not the key. +fn key_column_named(column: &str, pk: Option<&str>) -> Option { + match pk { + Some(pk) => column.eq_ignore_ascii_case(pk).then(|| pk.to_string()), + None => { + let column = column.to_lowercase(); + matches!(column.as_str(), "id" | "document_id" | "key").then_some(column) + } + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::catalog::SqlCatalogError; + use crate::temporal::TemporalScope; + + struct ChangingCatalog; + + impl SqlCatalog for ChangingCatalog { + fn get_collection( + &self, + _: DatabaseId, + _: &str, + ) -> std::result::Result, SqlCatalogError> { + Err(SqlCatalogError::RetryableSchemaChanged { + descriptor: "collection users".into(), + }) + } + } + + fn id_scan() -> SqlPlan { + SqlPlan::Scan { + collection: "users".into(), + alias: None, + engine: EngineType::DocumentSchemaless, + filters: vec![Filter { + expr: FilterExpr::Comparison { + field: "id".into(), + op: CompareOp::Eq, + value: SqlValue::String("u1".into()), + }, + }], + projection: Vec::new(), + sort_keys: Vec::new(), + limit: None, + offset: 0, + distinct: false, + window_functions: Vec::new(), + temporal: TemporalScope::default(), + } + } + + /// A schemaless `users` collection keyed by `id`. + struct UsersCatalog; + + impl SqlCatalog for UsersCatalog { + fn get_collection( + &self, + _: DatabaseId, + name: &str, + ) -> std::result::Result, SqlCatalogError> { + Ok((name == "users").then(|| CollectionInfo { + name: "users".to_string(), + engine: EngineType::DocumentSchemaless, + columns: Vec::new(), + primary_key: Some("id".to_string()), + has_auto_tier: false, + indexes: Vec::new(), + bitemporal: false, + primary: nodedb_types::PrimaryEngine::Document, + vector_primary: None, + partition_strategy: nodedb_types::PartitionStrategy::CollectionHomed, + open_schema: CollectionInfo::open_schema_for(EngineType::DocumentSchemaless), + })) + } + } + + #[test] + fn a_computed_item_over_a_key_lookup_is_control_plane_computed() { + let SqlPlan::Scan { + collection, + alias, + engine, + filters, + sort_keys, + limit, + offset, + distinct, + window_functions, + temporal, + .. + } = id_scan() + else { + panic!("id_scan is a scan"); + }; + let scan = SqlPlan::Scan { + collection, + alias, + engine, + filters, + projection: vec![Projection::Computed { + expr: SqlExpr::Function { + name: "to_jsonb".into(), + args: vec![SqlExpr::Wildcard], + distinct: false, + }, + alias: "document".into(), + }], + sort_keys, + limit, + offset, + distinct, + window_functions, + temporal, + }; + match optimize(scan, &UsersCatalog).unwrap() { + SqlPlan::PointGet { projection, .. } => assert!(matches!( + projection.as_slice(), + [Projection::CpComputed { + expr: SqlExpr::Function { name, args, .. }, + alias, + }] if name == "to_jsonb" + && matches!(args.as_slice(), [SqlExpr::Wildcard]) + && alias == "document" + )), + other => panic!("expected a point lookup, got {other:?}"), + } + } + + #[test] + fn a_catalog_error_fails_the_point_get_rewrite() { + let error = optimize(id_scan(), &ChangingCatalog).unwrap_err(); + assert_eq!( + error, + crate::SqlError::from(SqlCatalogError::RetryableSchemaChanged { + descriptor: "collection users".into(), + }) + ); + } + + /// A collection keyed by the declared column `Sku`. + struct SkuCatalog; + + impl SqlCatalog for SkuCatalog { + fn get_collection( + &self, + _: DatabaseId, + name: &str, + ) -> std::result::Result, SqlCatalogError> { + Ok(Some(CollectionInfo { + name: name.to_string(), + engine: EngineType::DocumentSchemaless, + columns: Vec::new(), + primary_key: Some("Sku".to_string()), + has_auto_tier: false, + indexes: Vec::new(), + bitemporal: false, + primary: nodedb_types::PrimaryEngine::Document, + vector_primary: None, + partition_strategy: nodedb_types::PartitionStrategy::CollectionHomed, + open_schema: CollectionInfo::open_schema_for(EngineType::DocumentSchemaless), + })) + } + } + + fn equality_scan(field: &str, engine: EngineType) -> SqlPlan { + let SqlPlan::Scan { + collection, + alias, + projection, + sort_keys, + limit, + offset, + distinct, + window_functions, + temporal, + .. + } = id_scan() + else { + panic!("id_scan is a scan"); + }; + SqlPlan::Scan { + collection, + alias, + engine, + filters: vec![Filter { + expr: FilterExpr::Comparison { + field: field.into(), + op: CompareOp::Eq, + value: SqlValue::String("p1".into()), + }, + }], + projection, + sort_keys, + limit, + offset, + distinct, + window_functions, + temporal, + } + } + + #[test] + fn a_declared_key_equality_is_a_point_lookup_on_that_key() { + match optimize( + equality_scan("sku", EngineType::DocumentSchemaless), + &SkuCatalog, + ) + .unwrap() + { + SqlPlan::PointGet { + key_column, + key_value, + .. + } => { + assert_eq!(key_column, "Sku", "the catalog spelling names the key"); + assert_eq!(key_value, SqlValue::String("p1".into())); } - let val = match right.as_ref() { - SqlExpr::Literal(v) => v.clone(), - _ => return None, - }; - Some((col, val)) + other => panic!("expected a point lookup, got {other:?}"), + } + } + + #[test] + fn an_id_column_beside_a_declared_key_stays_a_scan() { + let plan = optimize( + equality_scan("id", EngineType::DocumentSchemaless), + &SkuCatalog, + ) + .unwrap(); + assert!(matches!(plan, SqlPlan::Scan { .. }), "{plan:?}"); + } + + #[test] + fn a_key_equality_on_an_engine_without_point_lookups_stays_a_scan() { + for engine in [EngineType::Timeseries, EngineType::Array] { + let plan = optimize(equality_scan("sku", engine), &SkuCatalog).unwrap(); + assert!(matches!(plan, SqlPlan::Scan { .. }), "{plan:?}"); } - _ => None, } } diff --git a/nodedb-sql/src/planner/cp_projection.rs b/nodedb-sql/src/planner/cp_projection.rs index efd9c0a82..33fbd3173 100644 --- a/nodedb-sql/src/planner/cp_projection.rs +++ b/nodedb-sql/src/planner/cp_projection.rs @@ -18,8 +18,8 @@ use crate::functions::sequence_accessor::is_sequence_accessor; use crate::resolver::ColumnScope; use crate::resolver::columns::TableScope; use crate::resolver::expr::convert_expr; -use crate::types::Projection; use crate::types::plan::contains_sequence_accessor; +use crate::types::{Projection, SqlExpr}; /// A copy of `scope` that resolves sequence accessors, when `scope` iterates /// rows and backs the statement's output SELECT. A FROM-less SELECT gets @@ -93,11 +93,40 @@ pub fn convert_cp_item( })) } +/// A computed item as a Control-Plane computed item. +/// +/// A computed item that is a bare column under an alias stays `Computed`: the +/// value is the stored column, looked up under the source name and displayed +/// under the alias, so nothing is evaluated. Every other item passes through. +pub fn to_cp_computed(projection: Projection) -> Projection { + match projection { + Projection::Computed { expr, alias } => match expr { + SqlExpr::Column { .. } => Projection::Computed { expr, alias }, + SqlExpr::Function { .. } + | SqlExpr::Literal(_) + | SqlExpr::BinaryOp { .. } + | SqlExpr::UnaryOp { .. } + | SqlExpr::Case { .. } + | SqlExpr::Cast { .. } + | SqlExpr::Subquery(_) + | SqlExpr::Wildcard + | SqlExpr::IsNull { .. } + | SqlExpr::InList { .. } + | SqlExpr::Between { .. } + | SqlExpr::Like { .. } + | SqlExpr::ArrayLiteral(_) => Projection::CpComputed { expr, alias }, + }, + Projection::Column(_) + | Projection::Star + | Projection::QualifiedStar(_) + | Projection::CpComputed { .. } => projection, + } +} + #[cfg(test)] mod tests { use super::*; use crate::resolver::columns::test_support::open_scope; - use crate::types::SqlExpr; use sqlparser::dialect::GenericDialect; use sqlparser::parser::Parser; diff --git a/nodedb-sql/src/planner/returning.rs b/nodedb-sql/src/planner/returning.rs index df6a5cbab..1a7b50925 100644 --- a/nodedb-sql/src/planner/returning.rs +++ b/nodedb-sql/src/planner/returning.rs @@ -16,9 +16,10 @@ use nodedb_types::DatabaseId; use crate::error::{Result, SqlError}; use crate::parser::statement::parse_sql; +use crate::planner::cp_projection::to_cp_computed; use crate::planner::select::helpers::convert_projection; use crate::resolver::columns::{ResolvedTable, TableScope}; -use crate::types::{Projection, SqlCatalog, SqlExpr}; +use crate::types::{Projection, SqlCatalog}; /// Resolve a RETURNING item list against the DML target collection. /// @@ -84,40 +85,11 @@ fn parse_items(items_sql: &str) -> Result> { Ok(select.projection) } -/// A computed item becomes Control-Plane computed. A computed item that is a -/// bare column under an alias stays `Computed`: the value is the stored -/// column, looked up under the source name and displayed under the alias, so -/// nothing is evaluated. -fn to_cp_computed(projection: Projection) -> Projection { - match projection { - Projection::Computed { expr, alias } => match expr { - SqlExpr::Column { .. } => Projection::Computed { expr, alias }, - SqlExpr::Function { .. } - | SqlExpr::Literal(_) - | SqlExpr::BinaryOp { .. } - | SqlExpr::UnaryOp { .. } - | SqlExpr::Case { .. } - | SqlExpr::Cast { .. } - | SqlExpr::Subquery(_) - | SqlExpr::Wildcard - | SqlExpr::IsNull { .. } - | SqlExpr::InList { .. } - | SqlExpr::Between { .. } - | SqlExpr::Like { .. } - | SqlExpr::ArrayLiteral(_) => Projection::CpComputed { expr, alias }, - }, - Projection::Column(_) - | Projection::Star - | Projection::QualifiedStar(_) - | Projection::CpComputed { .. } => projection, - } -} - #[cfg(test)] mod tests { use super::*; use crate::catalog::SqlCatalogError; - use crate::types::{CollectionInfo, ColumnInfo, EngineType, SqlDataType}; + use crate::types::{CollectionInfo, ColumnInfo, EngineType, SqlDataType, SqlExpr}; /// One strict `items` collection: `id TEXT`, `score BIGINT`. struct ItemsCatalog; diff --git a/nodedb-types/src/id/qualified_collection.rs b/nodedb-types/src/id/qualified_collection.rs index 3f8a2e86e..f614f4c91 100644 --- a/nodedb-types/src/id/qualified_collection.rs +++ b/nodedb-types/src/id/qualified_collection.rs @@ -58,6 +58,16 @@ impl QualifiedCollection { pub fn as_str(&self) -> &str { &self.0 } + + /// The collection name [`Self::new`] qualified for `database_id`: the + /// catalog key a collection lookup in that database takes. + pub fn collection_name(&self, database_id: DatabaseId) -> &str { + if database_id == DatabaseId::DEFAULT { + return &self.0; + } + let prefix = format!("{}/", database_id.as_u64()); + self.0.strip_prefix(prefix.as_str()).unwrap_or(&self.0) + } } impl fmt::Display for QualifiedCollection { @@ -83,6 +93,14 @@ mod tests { assert_eq!(q.as_str(), "7/users"); } + #[test] + fn collection_name_inverts_new() { + let q = QualifiedCollection::new(DatabaseId::new(7), "users"); + assert_eq!(q.collection_name(DatabaseId::new(7)), "users"); + let bare = QualifiedCollection::new(DatabaseId::DEFAULT, "users"); + assert_eq!(bare.collection_name(DatabaseId::DEFAULT), "users"); + } + #[test] fn serde_roundtrips_as_plain_string() { let q = QualifiedCollection::new(DatabaseId::new(7), "users"); diff --git a/nodedb-types/src/row_identity.rs b/nodedb-types/src/row_identity.rs index fd97daa92..45a2a1e95 100644 --- a/nodedb-types/src/row_identity.rs +++ b/nodedb-types/src/row_identity.rs @@ -178,6 +178,19 @@ impl std::fmt::Display for RowIdentity { } } +/// The declared key among a collection's resolved primary key, or `None`. +/// +/// `id` and `_rowid` name no declared key: a row keyed by either renders its +/// identity under [`DEFAULT_IDENTITY_COLUMN`]. Every other key is declared, +/// and a row's identity renders under it. Both planes apply this rule, so a +/// scan row, a filter image, and a write image name the identity alike. +pub fn declared_key(primary_key: Option<&str>) -> Option<&str> { + primary_key.filter(|key| { + !key.eq_ignore_ascii_case(DEFAULT_IDENTITY_COLUMN) + && !key.eq_ignore_ascii_case(ROWID_COLUMN) + }) +} + /// Extract the stringified value of `field` from a MessagePack row body. /// /// Returns `None` when the body is not an object, lacks `field`, or the diff --git a/nodedb/src/control/planner/catalog_adapter/mod.rs b/nodedb/src/control/planner/catalog_adapter/mod.rs index 00ed1e9dc..820fbdf05 100644 --- a/nodedb/src/control/planner/catalog_adapter/mod.rs +++ b/nodedb/src/control/planner/catalog_adapter/mod.rs @@ -36,4 +36,6 @@ mod sql_catalog_impl; mod type_convert; pub use adapter::OriginCatalog; -pub(crate) use type_convert::{convert_collection_type, declared_column_info}; +pub(crate) use type_convert::{ + convert_collection_type, declared_column_info, document_declared_key, +}; diff --git a/nodedb/src/control/planner/context/query/planning.rs b/nodedb/src/control/planner/context/query/planning.rs index 77aac1272..239ef8ee5 100644 --- a/nodedb/src/control/planner/context/query/planning.rs +++ b/nodedb/src/control/planner/context/query/planning.rs @@ -41,7 +41,8 @@ fn resolve_returning_and_output_schema( returning .as_ref() .map(|clause| clause.projection.as_slice()), - ); + ) + .map_err(|error| map_plan_error(error, tenant_id))?; Ok((output_schema, returning)) } @@ -287,7 +288,12 @@ impl QueryContext { let rls_version = sec.rls_store.tenant_version(tenant_id.as_u64()); // Inject RLS predicates. - crate::control::planner::rls_injection::inject_rls(&mut tasks, sec.rls_store, sec.auth)?; + crate::control::planner::rls_injection::inject_rls( + &mut tasks, + sec.rls_store, + self.catalog_inputs.credentials.catalog(), + sec.auth, + )?; // Refuse what column redaction cannot cover (aggregates over a // redacted column, graph traversals), before anything is dispatched. @@ -406,7 +412,12 @@ impl QueryContext { let rls_version = sec.rls_store.tenant_version(tenant_id.as_u64()); // Inject RLS predicates. - crate::control::planner::rls_injection::inject_rls(&mut tasks, sec.rls_store, sec.auth)?; + crate::control::planner::rls_injection::inject_rls( + &mut tasks, + sec.rls_store, + self.catalog_inputs.credentials.catalog(), + sec.auth, + )?; // Refuse what column redaction cannot cover (aggregates over a // redacted column, graph traversals), before anything is dispatched. diff --git a/nodedb/src/control/planner/rls_injection/context.rs b/nodedb/src/control/planner/rls_injection/context.rs index 3fb02aea3..55f3cdace 100644 --- a/nodedb/src/control/planner/rls_injection/context.rs +++ b/nodedb/src/control/planner/rls_injection/context.rs @@ -9,6 +9,7 @@ //! silent no-op, indistinguishable from no policy at all. use crate::control::security::auth_context::AuthContext; +use crate::control::security::catalog::SystemCatalog; use crate::control::security::rls::{PolicyType, RlsPolicyStore}; use super::filters::{get_rls, get_rls_write, merge_filters}; @@ -23,6 +24,8 @@ pub(super) struct RlsCtx<'a> { /// Database bare op-carried collection names (e.g. /// `AlgoParams.collection`) must be qualified against before a lookup. pub(super) database_id: nodedb_types::DatabaseId, + /// The catalog a write image's identity column resolves against. + pub(super) catalog: &'a SystemCatalog, } impl RlsCtx<'_> { @@ -48,15 +51,17 @@ impl RlsCtx<'_> { Ok(()) } - /// Store the collection's read policy in a dedicated post-fetch slot. + /// AND the collection's read policy into a post-fetch slot. The slot can + /// already hold the statement's own WHERE predicates (a vector search + /// carries them there), so the policy joins them and never replaces them. pub(super) fn set_post_filters( &self, collection: &nodedb_types::QualifiedCollection, rls_filters: &mut Vec, ) -> crate::Result<()> { - let rls = self.read_filters(collection)?; - if !rls.is_empty() { - *rls_filters = rls; + let policy = self.read_filters(collection)?; + if !policy.is_empty() { + merge_filters(rls_filters, &policy)?; } Ok(()) } @@ -111,29 +116,63 @@ impl RlsCtx<'_> { } /// Admit a document write, injecting the row's client-visible identity - /// as `id`. + /// under the collection's identity column. /// - /// A schemaless row with no declared `id` column carries its identity - /// only in its storage key, never in `image`. The caller decides the - /// encoding: a minted key renders as the surrogate's decimal string, - /// never the hex storage key, matching what the read paths inject via - /// `sparse_row_to_doc` — a policy naming `id` judges the write against - /// the value a later read returns. `inject_str_field` is a no-op when - /// `image` already carries `id`. + /// A row that lacks its identity column carries its identity only in its + /// storage key, never in `image`. The caller decides the encoding: a + /// minted key renders as the surrogate's decimal string, never the hex + /// storage key, matching what the read paths inject via + /// `sparse_row_to_doc` — a policy judges the write against the row a + /// later read returns. A declared-key image holds its key and gains no + /// `id`. `inject_str_field` is a no-op when `image` already carries the + /// column. pub(super) fn admit_document_write_image( &self, collection: &nodedb_types::QualifiedCollection, identity: &crate::engine::document::store::RowIdentity, image: &[u8], ) -> crate::Result<()> { - let with_id = nodedb_query::msgpack_scan::inject_str_field(image, "id", identity.as_str()); - self.admit_write_image(collection, &with_id) + let check = get_rls_write(self.store, self.tenant_id, collection.as_str(), self.auth)?; + if check.is_empty() { + return Ok(()); + } + let identity_column = self.identity_column(collection)?; + let with_id = nodedb_query::msgpack_scan::inject_str_field( + image, + &identity_column, + identity.as_str(), + ); + crate::control::security::rls::admit_compiled_write_image( + &check, + &with_id, + self.tenant_id, + collection.as_str(), + ) } - /// Admit a write whose post-image is a JSON object (a graph edge's - /// `PROPERTIES`). Non-object bytes, including an empty `PROPERTIES`, - /// deny rather than admit by omission. - pub(super) fn admit_write_json_image( + /// The column `collection`'s rows render their identity under: its + /// declared key, per `document_declared_key`, else `id`. The register + /// config the Data Plane reads comes from the same function, so a write + /// image and a read row name the identity alike. An unknown collection + /// has no declared key. + fn identity_column( + &self, + collection: &nodedb_types::QualifiedCollection, + ) -> crate::Result { + let name = collection.collection_name(self.database_id); + let declared_key = self + .catalog + .get_collection(self.database_id, self.tenant_id, name)? + .and_then(|stored| { + crate::control::planner::catalog_adapter::document_declared_key(&stored) + }); + Ok(declared_key.unwrap_or_else(|| nodedb_types::DEFAULT_IDENTITY_COLUMN.to_string())) + } + + /// Admit a write whose post-image is a plain-MessagePack property map (a + /// graph edge's `PROPERTIES`). Non-map bytes, including an empty + /// `PROPERTIES`, deny rather than admit by omission. + pub(super) fn admit_write_property_image( &self, collection: &nodedb_types::QualifiedCollection, image: &[u8], @@ -142,8 +181,8 @@ impl RlsCtx<'_> { if check.is_empty() { return Ok(()); } - let decoded = sonic_rs::from_slice::(image).ok(); - let Some(object @ serde_json::Value::Object(_)) = decoded else { + let decoded = nodedb_types::json_msgpack::value_from_msgpack(image).ok(); + let Some(nodedb_types::Value::Object(_)) = decoded else { return Err(crate::Error::RejectedAuthz { tenant_id: crate::types::TenantId::new(self.tenant_id), resource: format!( @@ -154,7 +193,7 @@ impl RlsCtx<'_> { }; crate::control::security::rls::admit_compiled_write_image( &check, - &nodedb_types::json_to_msgpack_or_empty(&object), + image, self.tenant_id, collection.as_str(), ) diff --git a/nodedb/src/control/planner/rls_injection/plan.rs b/nodedb/src/control/planner/rls_injection/plan.rs index a97fd4d09..12395ca1d 100644 --- a/nodedb/src/control/planner/rls_injection/plan.rs +++ b/nodedb/src/control/planner/rls_injection/plan.rs @@ -16,6 +16,7 @@ use crate::bridge::envelope::PhysicalPlan; use crate::control::security::auth_context::AuthContext; +use crate::control::security::catalog::SystemCatalog; use crate::control::security::rls::RlsPolicyStore; use nodedb_physical::physical_task::PhysicalTask; @@ -23,10 +24,13 @@ use super::context::RlsCtx; /// Inject RLS predicates into physical tasks after plan conversion: reads /// get filters injected, a write's policy admits its image or refuses. -/// `Err` on a missing `$auth` field or an uncoverable read/write shape. +/// `catalog` resolves the identity column a write image names its row by. +/// `Err` on a missing `$auth` field, an uncoverable read/write shape, or a +/// catalog lookup error. pub fn inject_rls( tasks: &mut [PhysicalTask], rls_store: &RlsPolicyStore, + catalog: &SystemCatalog, auth: &AuthContext, ) -> crate::Result<()> { for task in tasks.iter_mut() { @@ -35,6 +39,7 @@ pub fn inject_rls( tenant_id: task.tenant_id.as_u64(), auth, database_id: task.database_id, + catalog, }; walk(&ctx, &mut task.plan)?; refuse_undecided_write_check(&task.plan)?; @@ -48,6 +53,7 @@ pub fn inject_rls_for_single_plan( database_id: nodedb_types::DatabaseId, plan: &mut PhysicalPlan, rls_store: &RlsPolicyStore, + catalog: &SystemCatalog, auth: &AuthContext, ) -> crate::Result<()> { let ctx = RlsCtx { @@ -55,6 +61,7 @@ pub fn inject_rls_for_single_plan( tenant_id, auth, database_id, + catalog, }; walk(&ctx, plan)?; refuse_undecided_write_check(plan) @@ -105,6 +112,7 @@ pub(super) fn walk(ctx: &RlsCtx<'_>, plan: &mut PhysicalPlan) -> crate::Result<( #[cfg(test)] pub(super) mod test_support { use crate::control::security::auth_context::AuthContext; + use crate::control::security::catalog::SystemCatalog; use crate::control::security::predicate::{CompareOp, PredicateValue, RlsPredicate}; use crate::control::security::rls::{PolicyType, RlsPolicy, RlsPolicyStore}; use crate::types::TenantId; @@ -265,6 +273,7 @@ pub(super) mod test_support { nodedb_types::DatabaseId::DEFAULT, plan, store, + &SystemCatalog::open_in_memory().expect("in-memory catalog"), ®ular_auth(), ) } @@ -278,6 +287,7 @@ pub(super) mod test_support { nodedb_types::DatabaseId::DEFAULT, plan, &RlsPolicyStore::new(), + &SystemCatalog::open_in_memory().expect("in-memory catalog"), ®ular_auth(), ) } diff --git a/nodedb/src/control/planner/sql_plan_convert/dml/key_assignment.rs b/nodedb/src/control/planner/sql_plan_convert/dml/key_assignment.rs new file mode 100644 index 000000000..084e64301 --- /dev/null +++ b/nodedb/src/control/planner/sql_plan_convert/dml/key_assignment.rs @@ -0,0 +1,117 @@ +// SPDX-License-Identifier: BUSL-1.1 + +//! A plan-time assignment to a document's identity column keeps its key. +//! +//! A document row stays stored under its key when a write rewrites it in +//! place, so its identity column must keep naming that key. A changed column +//! makes SQL reads and key reads name different documents. Changing a key is +//! a DELETE plus an INSERT under the new key. + +use nodedb_sql::types::{SqlExpr, SqlValue}; + +use super::super::value::sql_value_to_string; + +/// Refuse an assignment that moves the document stored under `doc_id` off its +/// key. +/// +/// An assignment to `identity_column` keeps the key when its right-hand side +/// is the column itself or `EXCLUDED.` (both hold the row's key), or a +/// literal that renders as `doc_id`. A NULL literal is a NOT NULL violation. +pub(super) fn check_assignments_keep_key( + collection: &str, + identity_column: &str, + assignments: &[(String, SqlExpr)], + doc_id: &str, +) -> crate::Result<()> { + for (field, expr) in assignments { + if !field.eq_ignore_ascii_case(identity_column) { + continue; + } + let keeps_key = match expr { + SqlExpr::Column { name, .. } => name.eq_ignore_ascii_case(identity_column), + SqlExpr::Literal(SqlValue::Null) => { + return Err(crate::Error::RejectedConstraint { + collection: collection.to_string(), + constraint: "not_null".to_string(), + detail: format!("primary key '{identity_column}' cannot be set to NULL"), + }); + } + SqlExpr::Literal(value) => sql_value_to_string(value) == doc_id, + _ => false, + }; + if !keeps_key { + return Err(crate::Error::RejectedConstraint { + collection: collection.to_string(), + constraint: "primary_key_immutable".to_string(), + detail: format!( + "cannot change primary key '{identity_column}' of document '{doc_id}' \ + in '{collection}': the row stays stored under its current key; DELETE \ + the row and INSERT it under the new key" + ), + }); + } + } + Ok(()) +} + +#[cfg(test)] +mod tests { + use super::*; + + fn assign(rhs: SqlExpr) -> Vec<(String, SqlExpr)> { + vec![("id".to_string(), rhs)] + } + + fn constraint_of(result: crate::Result<()>) -> String { + match result { + Err(crate::Error::RejectedConstraint { constraint, .. }) => constraint, + other => panic!("expected RejectedConstraint, got {other:?}"), + } + } + + #[test] + fn an_assignment_naming_the_key_keeps_it() { + for rhs in [ + SqlExpr::Column { + table: Some("excluded".into()), + name: "id".into(), + }, + SqlExpr::Column { + table: None, + name: "id".into(), + }, + SqlExpr::Literal(SqlValue::String("k1".into())), + ] { + check_assignments_keep_key("docs", "id", &assign(rhs), "k1") + .expect("the assignment names the row's key"); + } + let other_column = vec![( + "v".to_string(), + SqlExpr::Literal(SqlValue::String("x".into())), + )]; + check_assignments_keep_key("docs", "id", &other_column, "k1") + .expect("a non-key assignment is unchecked"); + } + + #[test] + fn an_assignment_moving_the_key_is_refused() { + let moved = assign(SqlExpr::Literal(SqlValue::String("k2".into()))); + assert_eq!( + constraint_of(check_assignments_keep_key("docs", "id", &moved, "k1")), + "primary_key_immutable" + ); + let computed = assign(SqlExpr::Column { + table: Some("excluded".into()), + name: "v".into(), + }); + assert_eq!( + constraint_of(check_assignments_keep_key("docs", "id", &computed, "k1")), + "primary_key_immutable" + ); + let nulled = assign(SqlExpr::Literal(SqlValue::Null)); + assert_eq!( + constraint_of(check_assignments_keep_key("docs", "id", &nulled, "k1")), + "not_null" + ); + } +} diff --git a/nodedb/src/control/planner/sql_plan_convert/dml/kv_insert.rs b/nodedb/src/control/planner/sql_plan_convert/dml/kv_insert.rs index e01cce8b6..f2c1d84db 100644 --- a/nodedb/src/control/planner/sql_plan_convert/dml/kv_insert.rs +++ b/nodedb/src/control/planner/sql_plan_convert/dml/kv_insert.rs @@ -3,6 +3,7 @@ //! `SqlPlan::KvInsert` → `PhysicalTask` lowering. use nodedb_sql::types::{KvInsertIntent, SqlExpr, SqlValue}; +use nodedb_types::CollectionType; use crate::bridge::envelope::PhysicalPlan; use crate::types::TenantId; @@ -10,10 +11,11 @@ use nodedb_physical::physical_plan::*; use super::super::convert::ConvertContext; use super::super::value::{ - assignments_to_update_values, sql_value_to_bytes, write_msgpack_map_header, write_msgpack_str, - write_msgpack_value, + assignments_to_update_values, sql_value_to_bytes, sql_value_to_string, + write_msgpack_map_header, write_msgpack_str, write_msgpack_value, }; use super::insert::assign_for_pk; +use super::key_assignment::check_assignments_keep_key; use nodedb_physical::physical_task::{PhysicalTask, PostSetOp}; pub(in super::super) fn convert_kv_insert( @@ -29,10 +31,28 @@ pub(in super::super) fn convert_kv_insert( let coll_qualified = super::super::convert::db_qualified(ctx.database_id, collection); let qualified_collection = nodedb_types::QualifiedCollection::new(ctx.database_id, collection); let collection = coll_qualified.as_str(); - let update_values = if on_conflict_updates.is_empty() { - Vec::new() + // The conflict branch merges the assignments into the body stored under + // the row's key, and a named key column is a copy of that key in the + // body. An assignment to the key column must name the key it keeps. + let key_column = if on_conflict_updates.is_empty() { + None } else { - assignments_to_update_values(on_conflict_updates)? + Some(kv_key_column(ctx, collection)?) + }; + // The built-in `key` column is not in the body, so an assignment that + // keeps it writes nothing there. + let update_values = match key_column.as_deref() { + None => Vec::new(), + Some(key_column) => { + let body_assignments: Vec<(String, SqlExpr)> = on_conflict_updates + .iter() + .filter(|(field, _)| { + !(key_column == KV_KEY_COLUMN && field.eq_ignore_ascii_case(KV_KEY_COLUMN)) + }) + .cloned() + .collect(); + assignments_to_update_values(&body_assignments)? + } }; let vshard = collection_key.vshard(); let ttl_ms = ttl_secs * 1000; @@ -48,6 +68,14 @@ pub(in super::super) fn convert_kv_insert( detail: "primary key cannot be NULL or omitted".to_string(), }); } + if let Some(key_column) = key_column.as_deref() { + check_assignments_keep_key( + collection, + key_column, + on_conflict_updates, + &sql_value_to_string(key_val), + )?; + } let key = sql_value_to_bytes(key_val)?; let value = if value_cols.len() == 1 && value_cols[0].0 == "value" { sql_value_to_bytes(&value_cols[0].1)? @@ -86,7 +114,9 @@ pub(in super::super) fn convert_kv_insert( returning: None, rls_filters: Vec::new(), }, - KvInsertIntent::Put if !update_values.is_empty() => KvOp::InsertOnConflictUpdate { + // A conflict clause whose assignments all keep the built-in key + // still merges: an empty merge keeps the stored body. + KvInsertIntent::Put if key_column.is_some() => KvOp::InsertOnConflictUpdate { collection: qualified_collection.clone(), key, value, @@ -127,3 +157,31 @@ pub(in super::super) fn convert_kv_insert( } Ok(tasks) } + +/// The column a KV collection's key is read from: its schema's primary-key +/// column, else the built-in `key` column. The SQL planner extracts the key +/// by the same rule. +fn kv_key_column(ctx: &ConvertContext, collection: &str) -> crate::Result { + let Some(credentials) = ctx.credentials.as_ref() else { + return Ok(KV_KEY_COLUMN.to_string()); + }; + let bare = + crate::control::target_identity::naming::bare_collection_name(ctx.database_id, collection); + let stored = + credentials + .catalog() + .get_collection(ctx.database_id, ctx.tenant_id.as_u64(), &bare)?; + let declared = stored.and_then(|stored| match stored.collection_type { + CollectionType::KeyValue(config) => config + .schema + .columns + .into_iter() + .find(|column| column.primary_key) + .map(|column| column.name), + CollectionType::Document(_) | CollectionType::Columnar(_) => None, + }); + Ok(declared.unwrap_or_else(|| KV_KEY_COLUMN.to_string())) +} + +/// The key column of a KV collection with no declared primary key. +const KV_KEY_COLUMN: &str = "key"; diff --git a/nodedb/src/control/planner/sql_plan_convert/dml/mod.rs b/nodedb/src/control/planner/sql_plan_convert/dml/mod.rs index 286e78e92..aaec18631 100644 --- a/nodedb/src/control/planner/sql_plan_convert/dml/mod.rs +++ b/nodedb/src/control/planner/sql_plan_convert/dml/mod.rs @@ -3,6 +3,7 @@ mod balanced_gate; mod crdt_gate; mod insert; +mod key_assignment; mod kv_insert; mod merge; mod surrogate_keys; diff --git a/nodedb/src/control/planner/sql_plan_convert/dml/update_delete/update.rs b/nodedb/src/control/planner/sql_plan_convert/dml/update_delete/update.rs index fc316bc1d..ef6da81f9 100644 --- a/nodedb/src/control/planner/sql_plan_convert/dml/update_delete/update.rs +++ b/nodedb/src/control/planner/sql_plan_convert/dml/update_delete/update.rs @@ -162,24 +162,7 @@ pub(in crate::control::planner::sql_plan_convert) fn convert_update( // No document store: route to ColumnarOp::Update regardless of PK-reduced WHERE. if matches!(engine, EngineType::Columnar | EngineType::Spatial) { - // Literals only: expressions need row-context eval, not wired into the columnar handler. - use nodedb_physical::physical_plan::UpdateValue; - let mut columnar_updates: Vec<(String, Vec)> = Vec::with_capacity(updates.len()); - for (field, update_val) in &updates { - match update_val { - UpdateValue::Literal(bytes) => { - columnar_updates.push((field.clone(), bytes.clone())) - } - UpdateValue::Expr(_) => { - return Err(crate::Error::BadRequest { - detail: format!( - "UPDATE with non-literal RHS on columnar/spatial engine \ - (field '{field}') is not yet supported; use a literal value" - ), - }); - } - } - } + // An expression assignment evaluates per matched row on the Data Plane. // PK-targeted WHERE: convert target_keys to an Eq filter on the PK column. let effective_filter = pk_effective_filter(filter_bytes, target_keys)?; return Ok(vec![PhysicalTask { @@ -189,7 +172,7 @@ pub(in crate::control::planner::sql_plan_convert) fn convert_update( plan: PhysicalPlan::Columnar(ColumnarOp::Update { collection: qualified_collection.clone(), filters: effective_filter, - updates: columnar_updates, + updates, rls_write_check: nodedb_types::RlsWriteCheck::pending_injection(), }), post_set_op: PostSetOp::None, @@ -209,6 +192,19 @@ pub(in crate::control::planner::sql_plan_convert) fn convert_update( } // CRDT partial-update payload, built once from the literal SET assignments. let crdt_fields_json = if is_crdt { + // The partial upsert merges these fields into the row stored under + // each target key, so an identity assignment must name that key. + let identity_column = declared_primary_key + .as_deref() + .unwrap_or(nodedb_types::DEFAULT_IDENTITY_COLUMN); + for key in target_keys { + super::super::key_assignment::check_assignments_keep_key( + collection, + identity_column, + assignments, + &sql_value_to_string(key), + )?; + } Some(super::super::crdt_gate::literal_assignments_to_fields_json( assignments, )?) @@ -234,7 +230,7 @@ pub(in crate::control::planner::sql_plan_convert) fn convert_update( if edge_bearing || multi_key { // Reject `Expr` RHS to a reserved edge field: reconciliation diffs against - // literal SET values only (mirrors the KV/columnar `Expr`-RHS rejection). + // literal SET values only (mirrors the KV `Expr`-RHS rejection). if edge_bearing && let Some((field, _)) = assignments.iter().find(|(field, expr)| { matches!(field.as_str(), "_from" | "_to" | "_type") diff --git a/nodedb/src/control/planner/sql_plan_convert/dml/upsert.rs b/nodedb/src/control/planner/sql_plan_convert/dml/upsert.rs index b632d3e76..b647688c6 100644 --- a/nodedb/src/control/planner/sql_plan_convert/dml/upsert.rs +++ b/nodedb/src/control/planner/sql_plan_convert/dml/upsert.rs @@ -16,9 +16,10 @@ use nodedb_physical::physical_plan::*; use super::super::convert::ConvertContext; use super::super::value::{assignments_to_update_values, row_to_msgpack, rows_to_msgpack_array}; use super::insert::{ - build_schema_bytes, columnar_row_surrogates, declared_primary_key_name, is_auto_rowid_pk, - resolve_doc_identity_with_declared, + build_schema_bytes, columnar_row_surrogates, declared_primary_key_name, doc_identity_keys, + is_auto_rowid_pk, resolve_doc_identity_with_declared, }; +use super::key_assignment::check_assignments_keep_key; use nodedb_physical::physical_task::{PhysicalTask, PostSetOp}; /// Bundled arguments for [`convert_upsert`]. @@ -97,6 +98,13 @@ pub(in super::super) fn convert_upsert( declared_pk.as_deref(), row, )?; + // The conflicting row stays stored under `doc_id`. + check_assignments_keep_key( + collection, + declared_pk.as_deref().unwrap_or(primary_key), + on_conflict_updates, + &doc_id, + )?; let plan = if is_crdt { PhysicalPlan::Crdt(CrdtOp::DocUpsert { collection: qualified_collection.clone(), @@ -141,6 +149,16 @@ pub(in super::super) fn convert_upsert( } if !columnar_rows.is_empty() { + // A columnar row's key is read from its merged values, and the write + // deletes only the row stored under that key. A merge that moves the + // key leaves the conflicting row in place and inserts a second row + // under the conflicting row's surrogate. An assignment to the key + // column must name the key it keeps. A row that names no key mints a + // fresh one, so it never conflicts. + let key_column = declared_pk.as_deref().unwrap_or(primary_key); + for doc_id in doc_identity_keys(primary_key, declared_pk.as_deref(), &columnar_rows) { + check_assignments_keep_key(collection, key_column, on_conflict_updates, &doc_id)?; + } let payload = rows_to_msgpack_array(&columnar_rows)?; let surrogates = columnar_row_surrogates( ctx, diff --git a/nodedb/src/control/planner/sql_plan_convert/dml/vector_primary.rs b/nodedb/src/control/planner/sql_plan_convert/dml/vector_primary.rs index f0ca30c23..4bf525b26 100644 --- a/nodedb/src/control/planner/sql_plan_convert/dml/vector_primary.rs +++ b/nodedb/src/control/planner/sql_plan_convert/dml/vector_primary.rs @@ -24,6 +24,7 @@ use super::super::value::{ use super::insert::{ declared_primary_key_name, is_auto_rowid_pk, resolve_doc_identity_with_declared, }; +use super::key_assignment::check_assignments_keep_key; /// Collection-level settings every vector-primary write carries. pub(in super::super) struct VectorPrimaryCfg<'a> { @@ -159,6 +160,10 @@ pub(in super::super) fn convert_vector_primary_insert( &row_fields, )?; let key_column = declared.as_deref().unwrap_or(primary_key); + // The conflict branch patches the sidecar of the node stored under + // `doc_id`, which keeps its surrogate. An assignment to the key + // column must name that key. + check_assignments_keep_key(collection, key_column, on_conflict_updates, &doc_id)?; let mut fields = row.payload_fields.clone(); if !is_auto_rowid_pk(primary_key) && !fields diff --git a/nodedb/src/control/planner/sql_plan_convert/expr/bridge_expr.rs b/nodedb/src/control/planner/sql_plan_convert/expr/bridge_expr.rs index f90ab6949..83e47f5a4 100644 --- a/nodedb/src/control/planner/sql_plan_convert/expr/bridge_expr.rs +++ b/nodedb/src/control/planner/sql_plan_convert/expr/bridge_expr.rs @@ -111,7 +111,8 @@ fn convert_expr_inner(expr: &SqlExpr, qualify: bool) -> crate::bridge::expr_eval to_type: cast_type, } } - SqlExpr::Wildcard => BExpr::Column("*".into()), + // `*` is the whole row: `to_jsonb(*)` reads the row document. + SqlExpr::Wildcard => BExpr::Column(nodedb_query::expr::WHOLE_ROW_COLUMN.into()), // NOT e / -e → evaluator's Negate (handles both bool and numeric). SqlExpr::UnaryOp { expr, .. } => BExpr::Negate(Box::new(convert_expr_inner(expr, qualify))), diff --git a/nodedb/src/control/server/http/routes/promql/remote.rs b/nodedb/src/control/server/http/routes/promql/remote.rs index 4a734addf..65431bf62 100644 --- a/nodedb/src/control/server/http/routes/promql/remote.rs +++ b/nodedb/src/control/server/http/routes/promql/remote.rs @@ -135,6 +135,7 @@ pub async fn remote_write( if let Err(error) = crate::control::planner::rls_injection::inject_rls( std::slice::from_mut(&mut task), &state.shared.rls, + state.shared.credentials.catalog(), scope.auth(), ) { tracing::warn!(error = ?error, collection = %collection, "remote write denied by row policy"); diff --git a/nodedb/src/control/server/ilp_batch/dispatch.rs b/nodedb/src/control/server/ilp_batch/dispatch.rs index 875eafa36..877d1fba4 100644 --- a/nodedb/src/control/server/ilp_batch/dispatch.rs +++ b/nodedb/src/control/server/ilp_batch/dispatch.rs @@ -198,7 +198,12 @@ async fn flush_ilp_batch_inner( let scope = ClientRequestScope::for_database(identity, state.auth_stores(), database_id, peer_addr) .into_resolved_scope(); - crate::control::planner::rls_injection::inject_rls(&mut tasks, &state.rls, scope.auth())?; + crate::control::planner::rls_injection::inject_rls( + &mut tasks, + &state.rls, + state.credentials.catalog(), + scope.auth(), + )?; // A spent hard quota refuses the batch before any of it is staged. The // charge below runs once the atomic Calvin write has committed, so it can diff --git a/nodedb/src/control/server/native/dispatch/direct_ops.rs b/nodedb/src/control/server/native/dispatch/direct_ops.rs index b357719e4..8a2c0d3f9 100644 --- a/nodedb/src/control/server/native/dispatch/direct_ops.rs +++ b/nodedb/src/control/server/native/dispatch/direct_ops.rs @@ -102,6 +102,7 @@ pub(crate) async fn handle_direct_op( ctx.database_id(), &mut plan, &ctx.state.rls, + ctx.state.credentials.catalog(), ctx.auth_context(), ) { return error_to_native_with_sqlstate(seq, "42501", &e); diff --git a/nodedb/src/control/server/native/dispatch/graph_match.rs b/nodedb/src/control/server/native/dispatch/graph_match.rs index f6f3749a0..fb70e0cea 100644 --- a/nodedb/src/control/server/native/dispatch/graph_match.rs +++ b/nodedb/src/control/server/native/dispatch/graph_match.rs @@ -51,6 +51,7 @@ pub(crate) async fn handle_graph_match( ctx.database_id(), &mut plan, &ctx.state.rls, + ctx.state.credentials.catalog(), ctx.auth_context(), ) { return error_to_native_with_sqlstate(seq, "42501", &error); diff --git a/nodedb/src/control/server/native/dispatch/plan_builder/document.rs b/nodedb/src/control/server/native/dispatch/plan_builder/document.rs index 545327475..14dfe4d30 100644 --- a/nodedb/src/control/server/native/dispatch/plan_builder/document.rs +++ b/nodedb/src/control/server/native/dispatch/plan_builder/document.rs @@ -12,6 +12,9 @@ use crate::bridge::envelope::PhysicalPlan; use nodedb_physical::physical_plan::{DocumentOp, KvOp, TimeseriesOp}; use super::super::DispatchCtx; +use super::document_identity::{ + identified_body, identified_json_body, identity_column, stores_schemaless_bodies, +}; use super::{collection_type, declared_primary_key, require_doc_id}; pub(crate) async fn build_point_get( @@ -62,7 +65,9 @@ pub(crate) async fn build_point_put( ) -> crate::Result { let doc_id = require_doc_id(fields)?; let value = fields.data.clone().unwrap_or_default(); - match collection_type(ctx, collection)? { + let coll_type = collection_type(ctx, collection)?; + let schemaless = stores_schemaless_bodies(coll_type.as_ref()); + match coll_type { Some(CollectionType::KeyValue(_)) => { let key = doc_id.into_bytes(); let surrogate = super::helpers::assign_surrogate(ctx, collection, &key).await?; @@ -108,6 +113,12 @@ pub(crate) async fn build_point_put( .to_string(), }), Some(CollectionType::Document(_)) | None => { + let value = if schemaless { + let column = identity_column(declared_primary_key(ctx, collection)?); + identified_body(&value, &doc_id, &column)? + } else { + value + }; let pk_bytes = doc_id.as_bytes().to_vec(); let surrogate = super::helpers::assign_surrogate(ctx, collection, &pk_bytes).await?; Ok(PhysicalPlan::Document(DocumentOp::PointPut { @@ -258,12 +269,18 @@ pub(crate) async fn build_batch_insert( // Every row's identity in one batch at the collection's home. let pks: Vec<&[u8]> = batch_docs.iter().map(|d| d.id.as_bytes()).collect(); let surrogates = super::helpers::assign_surrogates(ctx, collection, &pks).await?; + let schemaless = stores_schemaless_bodies(collection_type(ctx, collection)?.as_ref()); + let column = identity_column(declared_primary_key(ctx, collection)?); let mut documents: Vec<(String, Vec)> = Vec::with_capacity(batch_docs.len()); for d in batch_docs { - let value_bytes = sonic_rs::to_vec(&d.fields).map_err(|e| crate::Error::Serialization { - format: "json".into(), - detail: format!("failed to serialize document '{}': {e}", d.id), - })?; + let value_bytes = if schemaless { + identified_json_body(d.fields.clone(), &d.id, &column)? + } else { + sonic_rs::to_vec(&d.fields).map_err(|e| crate::Error::Serialization { + format: "json".into(), + detail: format!("failed to serialize document '{}': {e}", d.id), + })? + }; documents.push((d.id.clone(), value_bytes)); } Ok(PhysicalPlan::Document(DocumentOp::BatchInsert { @@ -363,6 +380,12 @@ pub(crate) async fn build_upsert( ) -> crate::Result { let doc_id = require_doc_id(fields)?; let value = fields.data.clone().unwrap_or_default(); + let value = if stores_schemaless_bodies(collection_type(ctx, collection)?.as_ref()) { + let column = identity_column(declared_primary_key(ctx, collection)?); + identified_body(&value, &doc_id, &column)? + } else { + value + }; let surrogate = super::helpers::assign_surrogate(ctx, collection, doc_id.as_bytes()).await?; Ok(PhysicalPlan::Document(DocumentOp::Upsert { collection: QualifiedCollection::new(ctx.database_id(), collection), diff --git a/nodedb/src/control/server/native/dispatch/plan_builder/document_identity.rs b/nodedb/src/control/server/native/dispatch/plan_builder/document_identity.rs new file mode 100644 index 000000000..f01a7bd63 --- /dev/null +++ b/nodedb/src/control/server/native/dispatch/plan_builder/document_identity.rs @@ -0,0 +1,239 @@ +// SPDX-License-Identifier: BUSL-1.1 + +//! The stored body of a schemaless document written over the native protocol. +//! +//! SQL reads a schemaless row's identity from its identity column: the +//! declared primary key, else `id`. A native write keys the row by the +//! document id it carries, so the stored body names that id under the +//! identity column. + +use nodedb_types::{CollectionType, DEFAULT_IDENTITY_COLUMN, Value}; + +/// Whether a collection stores schemaless document bodies. A collection the +/// catalog does not hold defaults to a schemaless document collection. +pub(super) fn stores_schemaless_bodies(collection_type: Option<&CollectionType>) -> bool { + collection_type.is_none_or(CollectionType::is_schemaless) +} + +/// The identity column of a schemaless collection: its declared primary key, +/// else `id`. The planner resolves a schemaless collection's key by the same +/// rule. +pub(super) fn identity_column(declared_primary_key: Option) -> String { + declared_primary_key.unwrap_or_else(|| DEFAULT_IDENTITY_COLUMN.to_string()) +} + +/// `body` as a MessagePack map that carries `doc_id` under `column`. +/// +/// `body` is a MessagePack map or JSON object text. A body `column` field +/// that is not the string `doc_id` is refused: the row is keyed by `doc_id`, +/// so SQL and key reads would disagree on which document the row is. +pub(super) fn identified_body(body: &[u8], doc_id: &str, column: &str) -> crate::Result> { + if nodedb_query::msgpack_scan::map_header(body, 0).is_some() { + if let Some((start, end)) = nodedb_query::msgpack_scan::extract_field(body, 0, column) { + let cell = body + .get(start..end) + .ok_or_else(|| crate::Error::BadRequest { + detail: format!("document '{doc_id}': body field '{column}' is truncated"), + })?; + let value = + nodedb_types::value_from_msgpack(cell).map_err(|e| crate::Error::BadRequest { + detail: format!( + "document '{doc_id}': body field '{column}' does not decode: {e}" + ), + })?; + check_body_id(&value, doc_id, column)?; + return Ok(body.to_vec()); + } + return Ok(nodedb_query::msgpack_scan::inject_str_field( + body, column, doc_id, + )); + } + let json: serde_json::Value = + sonic_rs::from_slice(body).map_err(|e| crate::Error::BadRequest { + detail: format!( + "document '{doc_id}': body is neither a MessagePack map nor JSON text: {e}" + ), + })?; + identified_json_body(json, doc_id, column) +} + +/// A JSON object as a MessagePack map that carries `doc_id` under `column`. +/// +/// A non-object is refused: a document's fields are a map. A body `column` +/// field that is not the string `doc_id` is refused. Numbers keep their exact +/// value: an unsigned integer above `i64::MAX` stays an unsigned MessagePack +/// integer. +pub(super) fn identified_json_body( + json: serde_json::Value, + doc_id: &str, + column: &str, +) -> crate::Result> { + let Some(object) = json.as_object() else { + return Err(crate::Error::BadRequest { + detail: format!("document '{doc_id}': body must be an object, got {json}"), + }); + }; + if let Some(id) = object.get(column) { + match id { + serde_json::Value::String(s) if s == doc_id => {} + other => { + return Err(id_mismatch(doc_id, column, &other.to_string())); + } + } + } + let map = nodedb_types::json_to_msgpack(&json).map_err(|e| crate::Error::Serialization { + format: "msgpack".into(), + detail: format!("document '{doc_id}': body encode failed: {e}"), + })?; + Ok(nodedb_query::msgpack_scan::inject_str_field( + &map, column, doc_id, + )) +} + +/// Refuse a decoded body identity that is not the string `doc_id`. +fn check_body_id(value: &Value, doc_id: &str, column: &str) -> crate::Result<()> { + match value { + Value::String(s) if s == doc_id => Ok(()), + other => Err(id_mismatch(doc_id, column, &format!("{other:?}"))), + } +} + +fn id_mismatch(doc_id: &str, column: &str, body_id: &str) -> crate::Error { + crate::Error::BadRequest { + detail: format!( + "document '{doc_id}': body field '{column}' holds {body_id}, which differs from \ + the document id '{doc_id}'; '{column}' is the collection's identity column; \ + remove the field or set it to the document id" + ), + } +} + +#[cfg(test)] +mod tests { + use super::*; + + fn decoded(body: &[u8]) -> Value { + nodedb_types::value_from_msgpack(body).expect("stored body decodes") + } + + fn field<'a>(value: &'a Value, name: &str) -> Option<&'a Value> { + match value { + Value::Object(map) => map.get(name), + other => panic!("stored body must be a map, got {other:?}"), + } + } + + fn packed(entries: &[(&str, Value)]) -> Vec { + let map = entries + .iter() + .map(|(k, v)| ((*k).to_string(), v.clone())) + .collect(); + nodedb_types::value_to_msgpack(&Value::Object(map)).expect("encode") + } + + fn assert_refused_naming_both(result: crate::Result>, body_id: &str) { + match result { + Err(crate::Error::BadRequest { detail }) => { + assert!(detail.contains("'d1'"), "names the document id: {detail}"); + assert!(detail.contains(body_id), "names the body id: {detail}"); + } + other => panic!("expected BadRequest, got {other:?}"), + } + } + + #[test] + fn the_identity_column_is_the_declared_key_else_id() { + assert_eq!(identity_column(None), "id"); + assert_eq!(identity_column(Some("sku".into())), "sku"); + } + + #[test] + fn json_body_gains_the_document_id() { + let body = identified_body(br#"{"body":"hello"}"#, "d1", "id").expect("json object"); + let value = decoded(&body); + assert_eq!(field(&value, "id"), Some(&Value::String("d1".into()))); + assert_eq!(field(&value, "body"), Some(&Value::String("hello".into()))); + } + + #[test] + fn msgpack_body_gains_the_document_id() { + let body = packed(&[("n", Value::Integer(7))]); + let value = decoded(&identified_body(&body, "d2", "id").expect("msgpack map")); + assert_eq!(field(&value, "id"), Some(&Value::String("d2".into()))); + assert_eq!(field(&value, "n"), Some(&Value::Integer(7))); + } + + #[test] + fn a_declared_key_carries_the_document_id_instead_of_id() { + let body = identified_body(br#"{"name":"pen"}"#, "d1", "sku").expect("json object"); + let value = decoded(&body); + assert_eq!(field(&value, "sku"), Some(&Value::String("d1".into()))); + assert_eq!(field(&value, "id"), None); + + let body = packed(&[("name", Value::String("pen".into()))]); + let value = decoded(&identified_body(&body, "d1", "sku").expect("msgpack map")); + assert_eq!(field(&value, "sku"), Some(&Value::String("d1".into()))); + assert_eq!(field(&value, "id"), None); + } + + #[test] + fn a_declared_key_naming_another_id_is_refused() { + assert_refused_naming_both(identified_body(br#"{"sku":"other"}"#, "d1", "sku"), "other"); + let body = packed(&[("sku", Value::String("other".into()))]); + assert_refused_naming_both(identified_body(&body, "d1", "sku"), "other"); + // `id` is an ordinary field once another column is the key. + assert!(identified_body(br#"{"id":"x"}"#, "d1", "sku").is_ok()); + } + + #[test] + fn a_body_naming_another_id_is_refused() { + assert_refused_naming_both(identified_body(br#"{"id":"mine"}"#, "d1", "id"), "mine"); + assert_refused_naming_both(identified_body(br#"{"id":5}"#, "d1", "id"), "5"); + let body = packed(&[("id", Value::String("mine".into()))]); + assert_refused_naming_both(identified_body(&body, "d1", "id"), "mine"); + let body = packed(&[("id", Value::Integer(5))]); + assert_refused_naming_both(identified_body(&body, "d1", "id"), "5"); + } + + #[test] + fn a_body_naming_its_own_id_is_accepted() { + let body = identified_body(br#"{"id":"d1","n":1}"#, "d1", "id").expect("json object"); + assert_eq!( + field(&decoded(&body), "id"), + Some(&Value::String("d1".into())) + ); + let body = packed(&[("id", Value::String("d1".into()))]); + let stored = identified_body(&body, "d1", "id").expect("msgpack map"); + assert_eq!( + field(&decoded(&stored), "id"), + Some(&Value::String("d1".into())) + ); + } + + #[test] + fn a_u64_above_i64_max_keeps_its_number() { + let body = + identified_body(br#"{"big":18446744073709551615}"#, "d1", "id").expect("json object"); + assert_eq!( + field(&decoded(&body), "big"), + Some(&Value::Decimal(rust_decimal::Decimal::from(u64::MAX))) + ); + let json = nodedb_types::json_from_msgpack(&body).expect("json read"); + assert_eq!(json["big"].as_u64(), Some(u64::MAX)); + } + + #[test] + fn a_non_object_body_is_refused() { + assert!(identified_body(b"[1,2]", "d1", "id").is_err()); + assert!(identified_body(b"", "d1", "id").is_err()); + } + + #[test] + fn only_schemaless_documents_store_identified_bodies() { + use nodedb_types::DocumentMode; + assert!(stores_schemaless_bodies(None)); + assert!(stores_schemaless_bodies(Some(&CollectionType::Document( + DocumentMode::Schemaless + )))); + } +} diff --git a/nodedb/src/control/server/pgwire/handler/routing/execute.rs b/nodedb/src/control/server/pgwire/handler/routing/execute.rs index b76fe619b..35bdf6000 100644 --- a/nodedb/src/control/server/pgwire/handler/routing/execute.rs +++ b/nodedb/src/control/server/pgwire/handler/routing/execute.rs @@ -166,6 +166,7 @@ impl NodeDbPgHandler { columns: described.columns.clone(), is_star: described.is_star, cp_computed: output_schema.cp_computed, + declared_key: output_schema.declared_key, }, None => output_schema, }; diff --git a/nodedb/src/control/server/resp/gateway_dispatch.rs b/nodedb/src/control/server/resp/gateway_dispatch.rs index a77515cae..a78cf5009 100644 --- a/nodedb/src/control/server/resp/gateway_dispatch.rs +++ b/nodedb/src/control/server/resp/gateway_dispatch.rs @@ -240,6 +240,7 @@ fn authorize_resp_task( database_id, &mut plan, &state.rls, + state.credentials.catalog(), scope.auth(), )?; diff --git a/nodedb/src/control/server/response_shape/compose/kernel.rs b/nodedb/src/control/server/response_shape/compose/kernel.rs index 2489009db..450c24a40 100644 --- a/nodedb/src/control/server/response_shape/compose/kernel.rs +++ b/nodedb/src/control/server/response_shape/compose/kernel.rs @@ -17,7 +17,7 @@ use crate::control::sequence::SequenceAccess; use super::super::project::push_flat_rows; use super::super::redaction::RedactionCtx; use super::super::schema::OutputSchema; -use super::super::stamp::stamp_rows; +use super::super::stamp::{NoSequenceAccess, stamp_rows}; use super::super::types::{DdlColType, ShapedRow, ShapedRows}; /// Pure shaping core: given an already-decoded Data-Plane value, unwrap the @@ -31,8 +31,8 @@ use super::super::types::{DdlColType, ShapedRow, ShapedRows}; /// or vector-translate but still needs the same envelope-unwrap + projection /// logic applied per batch, so streaming callers call this directly. /// -/// `sequences` resolves the projection's Control-Plane computed columns -/// (`cp_computed`). A projection that carries any and a caller that passes +/// `sequences` resolves the sequence accessors of the projection's +/// Control-Plane computed columns (`cp_computed`). An accessor call with /// `None` is an error: the alias would otherwise project as NULL. pub fn shape_decoded_rows( decoded: Value, @@ -41,7 +41,11 @@ pub fn shape_decoded_rows( sequences: Option<&dyn SequenceAccess>, ) -> crate::Result { let mut rows = Vec::new(); - push_flat_rows(decoded, &mut rows)?; + push_flat_rows( + decoded, + OutputSchema::identity_column(projection), + &mut rows, + )?; // Column-level redaction runs on the flat row maps, AFTER the scan // envelope is unwrapped and BEFORE any projection or column derivation. @@ -100,9 +104,10 @@ pub fn shape_decoded_rows( /// Stamp the projection's Control-Plane computed columns onto the flat rows. /// /// A projection with no computed column is a no-op whatever `sequences` is. -/// One that carries any needs session sequence access: a caller with none -/// (a per-batch stream, gateway forwarding, a clone merge) cannot answer -/// the statement, and says so rather than shipping NULL under the alias. +/// A caller with no session sequence state (a per-batch stream, gateway +/// forwarding, a clone merge) passes `None`. A column that calls no sequence +/// accessor evaluates there all the same. A sequence accessor call is +/// refused, never shipped as NULL under the alias. pub(in crate::control::server::response_shape) fn stamp_computed_columns( projection: Option<&OutputSchema>, rows: &mut [ShapedRow], @@ -111,11 +116,9 @@ pub(in crate::control::server::response_shape) fn stamp_computed_columns( let Some(schema) = projection.filter(|s| !s.cp_computed.is_empty()) else { return Ok(()); }; - let Some(access) = sequences else { - return Err(crate::Error::FeatureNotSupported { - detail: "Control-Plane computed columns need session sequence access on this path" - .to_string(), - }); + let access: &dyn SequenceAccess = match sequences { + Some(access) => access, + None => &NoSequenceAccess, }; stamp_rows(rows, &schema.cp_computed, access) } @@ -319,6 +322,7 @@ mod tests { .collect(), is_star: false, cp_computed: Vec::new(), + declared_key: None, } } @@ -584,4 +588,33 @@ mod tests { shape_decoded_rows(decoded, Some(&projection), None, None).expect_err("must refuse"); assert!(matches!(err, crate::Error::FeatureNotSupported { .. })); } + + /// `SELECT to_jsonb(*) AS document FROM t WHERE id = 'r1'`: the whole row + /// is one typed object, and no sequence access is needed to build it. + #[test] + fn a_whole_row_column_needs_no_sequence_access() { + use crate::bridge::expr_eval::SqlExpr; + let decoded = one_row(&[("id", text("r1")), ("n", Value::Integer(5))]); + let mut projection = named_projection(&[("document", "document")]); + projection.cp_computed = vec![ + crate::control::server::response_shape::schema::CpComputedColumn { + alias: "document".to_string(), + expr: SqlExpr::Function { + name: "to_jsonb".to_string(), + args: vec![SqlExpr::Column( + nodedb_query::expr::WHOLE_ROW_COLUMN.to_string(), + )], + }, + }, + ]; + + let shaped = shape_decoded_rows(decoded, Some(&projection), None, None) + .expect("an accessor-free column evaluates without sequence access"); + assert_eq!(shaped.columns, vec!["document".to_string()]); + let Value::Object(document) = &shaped.rows[0]["document"] else { + panic!("the whole row is an object: {:?}", shaped.rows[0]); + }; + assert_eq!(document.get("id"), Some(&text("r1"))); + assert_eq!(document.get("n"), Some(&Value::Integer(5))); + } } diff --git a/nodedb/src/control/server/response_shape/compose/materialized.rs b/nodedb/src/control/server/response_shape/compose/materialized.rs index ab8b1c10f..61d9144fc 100644 --- a/nodedb/src/control/server/response_shape/compose/materialized.rs +++ b/nodedb/src/control/server/response_shape/compose/materialized.rs @@ -112,9 +112,10 @@ pub fn shape_response_materialized( /// plan-dependent `apply_kv_wrap` / `translate_search_response` transforms those /// callers never ran. /// -/// `sequences` resolves the projection's Control-Plane computed columns; a -/// caller with no session in scope passes `None`, and a projection that -/// carries computed columns then fails rather than shipping NULL. +/// `sequences` resolves the sequence accessors of the projection's +/// Control-Plane computed columns. A caller with no session in scope passes +/// `None`. A column that calls an accessor then fails rather than shipping +/// NULL. A column that calls none evaluates. pub fn shape_payload_no_plan( payload: &[u8], plan_kind: PlanKind, diff --git a/nodedb/src/control/server/response_shape/project.rs b/nodedb/src/control/server/response_shape/project.rs index 969db72af..864f87bbb 100644 --- a/nodedb/src/control/server/response_shape/project.rs +++ b/nodedb/src/control/server/response_shape/project.rs @@ -19,11 +19,18 @@ use super::types::ShapedRow; /// The envelope `id` is a rendered [`StorageKey`](crate::engine::document::store::StorageKey) /// by construction. /// A value that fails `StorageKey::parse` is surfaced as `Err`, never accommodated. -pub fn push_flat_rows(value: Value, out: &mut Vec) -> crate::Result<()> { +/// +/// `identity_column` is the column a scan row's identity renders under: the +/// scanned collection's declared key, else the implicit `id`. +pub fn push_flat_rows( + value: Value, + identity_column: &str, + out: &mut Vec, +) -> crate::Result<()> { match value { Value::Array(items) => { for item in items { - push_flat_rows(item, out)?; + push_flat_rows(item, identity_column, out)?; } } Value::Object(mut map) => { @@ -32,24 +39,25 @@ pub fn push_flat_rows(value: Value, out: &mut Vec) -> crate::Result<( { let mut inner: ShapedRow = inner.into_iter().collect(); // The envelope carries the row's storage key, which is - // internal. A body with no `id` field renders its identity at - // this boundary, by the rule the Data Plane's row shaping + // internal. A body that lacks its identity column renders the + // identity there, by the rule the Data Plane's row shaping // applies: the body's `_rowid`, else the storage key. A copy // under a new surrogate keeps its `_rowid`, so it keeps its - // identity. `or_insert` leaves a declared primary key as the - // authority. + // identity. A body that holds its identity column keeps the + // stored value and gains no column: a declared key stays the + // authority, and no `id` appears beside it. if let Some(Value::String(key)) = map.remove("id") { let storage_key = crate::engine::document::store::StorageKey::parse(&key) .ok_or_else(|| crate::Error::Internal { detail: format!("scan envelope id is not a storage key: '{key}'"), })?; - let identity = match inner.get(nodedb_types::ROWID_COLUMN) { - Some(Value::Integer(rowid)) => rowid.to_string(), - _ => storage_key.to_identity().into_string(), - }; - inner - .entry("id".to_string()) - .or_insert(Value::String(identity)); + if !inner.contains_key(identity_column) { + let identity = match inner.get(nodedb_types::ROWID_COLUMN) { + Some(Value::Integer(rowid)) => rowid.to_string(), + _ => storage_key.to_identity().into_string(), + }; + inner.insert(identity_column.to_string(), Value::String(identity)); + } } out.push(inner); return Ok(()); @@ -138,7 +146,53 @@ pub fn cell_keys(columns: &[String]) -> Vec { #[cfg(test)] mod tests { - use super::cell_keys; + use super::{Value, cell_keys, push_flat_rows}; + + /// A scan envelope around `fields`, keyed by surrogate 7. + fn envelope(fields: &[(&str, &str)]) -> Value { + let key = nodedb_types::StorageKey::for_surrogate(nodedb_types::Surrogate::new(7)); + let data = fields + .iter() + .map(|(k, v)| ((*k).to_string(), Value::String((*v).to_string()))) + .collect(); + Value::Object( + [ + ("id".to_string(), Value::String(key.to_string())), + ("data".to_string(), Value::Object(data)), + ] + .into_iter() + .collect(), + ) + } + + fn flat(value: Value, identity_column: &str) -> Vec<(String, Value)> { + let mut rows = Vec::new(); + push_flat_rows(value, identity_column, &mut rows).expect("a well-formed envelope"); + assert_eq!(rows.len(), 1); + rows.remove(0).into_iter().collect() + } + + #[test] + fn a_row_holding_its_declared_key_gains_no_column() { + let row = flat(envelope(&[("sku", "p1"), ("name", "pen")]), "sku"); + assert_eq!( + row, + vec![ + ("name".to_string(), Value::String("pen".into())), + ("sku".to_string(), Value::String("p1".into())), + ] + ); + } + + #[test] + fn a_row_lacking_its_identity_column_renders_the_identity_there() { + let row = flat(envelope(&[("name", "pen")]), "sku"); + assert!(row.contains(&("sku".to_string(), Value::String("7".into())))); + assert!(!row.iter().any(|(k, _)| k == "id"), "{row:?}"); + + let row = flat(envelope(&[("name", "pen")]), "id"); + assert!(row.contains(&("id".to_string(), Value::String("7".into())))); + } fn keys(cols: &[&str]) -> Vec { cell_keys(&cols.iter().map(|s| s.to_string()).collect::>()) diff --git a/nodedb/src/control/server/response_shape/stamp/mod.rs b/nodedb/src/control/server/response_shape/stamp/mod.rs index 894b2f929..57b70fe95 100644 --- a/nodedb/src/control/server/response_shape/stamp/mod.rs +++ b/nodedb/src/control/server/response_shape/stamp/mod.rs @@ -3,6 +3,8 @@ //! Control-Plane evaluation of per-row computed SELECT-list columns. pub mod evaluate; +pub mod no_access; mod substitute; pub use evaluate::stamp_rows; +pub use no_access::NoSequenceAccess; diff --git a/nodedb/src/control/server/response_shape/stamp/no_access.rs b/nodedb/src/control/server/response_shape/stamp/no_access.rs new file mode 100644 index 000000000..ef93fa1b6 --- /dev/null +++ b/nodedb/src/control/server/response_shape/stamp/no_access.rs @@ -0,0 +1,59 @@ +// SPDX-License-Identifier: BUSL-1.1 + +//! The sequence access of a shaping path that holds no session. + +use crate::control::sequence::SequenceAccess; + +/// [`SequenceAccess`] for a path with no session sequence state: a per-batch +/// stream, gateway forwarding, a clone merge. +/// +/// A computed column that calls no sequence accessor evaluates without it. +/// A call to `nextval`, `currval` or `setval` is refused, never answered +/// with a guessed value. +pub struct NoSequenceAccess; + +impl NoSequenceAccess { + fn refuse(function: &str, name: &str) -> crate::Error { + crate::Error::FeatureNotSupported { + detail: format!( + "{function}('{name}') in a Control-Plane computed column needs session \ + sequence access, which this path does not hold" + ), + } + } +} + +impl SequenceAccess for NoSequenceAccess { + fn nextval(&self, name: &str) -> crate::Result { + Err(Self::refuse("nextval", name)) + } + + fn currval(&self, name: &str) -> crate::Result { + Err(Self::refuse("currval", name)) + } + + fn setval(&self, name: &str, _value: i64) -> crate::Result { + Err(Self::refuse("setval", name)) + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn every_accessor_is_refused_naming_the_sequence() { + for err in [ + NoSequenceAccess.nextval("s"), + NoSequenceAccess.currval("s"), + NoSequenceAccess.setval("s", 1), + ] { + match err { + Err(crate::Error::FeatureNotSupported { detail }) => { + assert!(detail.contains("('s')"), "{detail}"); + } + other => panic!("expected FeatureNotSupported, got {other:?}"), + } + } + } +} diff --git a/nodedb/src/control/server/response_translate/text_hybrid.rs b/nodedb/src/control/server/response_translate/text_hybrid.rs index c2c307386..f327bc187 100644 --- a/nodedb/src/control/server/response_translate/text_hybrid.rs +++ b/nodedb/src/control/server/response_translate/text_hybrid.rs @@ -5,12 +5,11 @@ //! search responses. //! //! `TextOp::Search` hits carry the standard `{id, data}` document-scan -//! envelope, keyed by `StorageKey::for_surrogate(surrogate)` hex — the document -//! body itself already carries the user's PK as an ordinary field (it was -//! written verbatim from the user's INSERT), so the resolved value only -//! needs injecting when the body has no `id` field of its own (a headless -//! FTS-indexed row with no document ever written, in which case there is no -//! PK to resolve either). +//! envelope, keyed by `StorageKey::for_surrogate(surrogate)` hex — the Data +//! Plane already gives every body its identity under the collection's +//! identity column (the declared key, else `id`), so the resolved value only +//! needs injecting, under that same column, when the row has no body (a +//! headless FTS-indexed row with no document ever written). //! //! `TextOp::HybridSearch` / `HybridSearchTriple` hits never fetch a document //! body at all — the row is just `{doc_id, , vector_rank?, @@ -31,12 +30,15 @@ use super::hit_key::parse_surrogate_hex; use super::vector::resolve_surrogate_pk; /// Decode the DP-side JSON/msgpack array of `TextOp::Search` / -/// `PhraseSearch`-shaped hits (`{id: , data: {...}}`), and for -/// any row whose `data` object has no `id` field of its own, resolve the -/// surrogate to the user PK via the catalog and inject it into `data`. Rows -/// whose body already carries an `id` (the common case) are left untouched; -/// an unresolved or headless surrogate is left untouched (no fabricated PK). -/// On any decode failure the payload is returned unchanged. +/// `PhraseSearch` / `BM25ScoreScan` rows (`{id: , data: {...}}`), and for +/// any row whose `data` object lacks the collection's identity column, resolve +/// the surrogate to the user PK via the catalog and inject it into `data` +/// under that column. A declared-key row never gains an `id` beside its key. +/// Rows that hold the column (the common case) are left untouched. An +/// unresolved or headless surrogate is left untouched (no fabricated PK). +/// On any decode failure, or a catalog error resolving the identity column, +/// the payload is returned unchanged: a row never gains a wrongly named +/// column. pub fn translate_text_search_payload( payload: &[u8], state: &SharedState, @@ -52,6 +54,9 @@ pub fn translate_text_search_payload( let Ok(JsonValue::Array(mut rows)) = sonic_rs::from_str::(&text) else { return payload.to_vec(); }; + let Ok(identity_column) = identity_column(state, database_id, tenant_id, collection) else { + return payload.to_vec(); + }; for row in &mut rows { let JsonValue::Object(map) = row else { @@ -60,11 +65,11 @@ pub fn translate_text_search_payload( let Some(JsonValue::String(hex_id)) = map.get("id").cloned() else { continue; }; - let data_has_id = matches!( + let data_has_identity = matches!( map.get("data"), - Some(JsonValue::Object(inner)) if inner.contains_key("id") + Some(JsonValue::Object(inner)) if inner.contains_key(&identity_column) ); - if data_has_id { + if data_has_identity { continue; } let Some(surrogate) = parse_surrogate_hex(&hex_id) else { @@ -73,7 +78,7 @@ pub fn translate_text_search_payload( if let Some(pk) = resolve_surrogate_pk(state, database_id, tenant_id, collection, surrogate) && let Some(JsonValue::Object(inner)) = map.get_mut("data") { - inner.insert("id".to_string(), JsonValue::String(pk)); + inner.insert(identity_column.clone(), JsonValue::String(pk)); } } @@ -83,6 +88,30 @@ pub fn translate_text_search_payload( } } +/// The column `collection`'s rows render their identity under: its declared +/// key, per `document_declared_key`, else `id`. The Data Plane's register +/// config comes from the same function. +fn identity_column( + state: &SharedState, + database_id: DatabaseId, + tenant_id: TenantId, + collection: &str, +) -> crate::Result { + let qualified = nodedb_types::QualifiedCollection::from_stored(collection.to_string()); + let declared_key = state + .credentials + .catalog() + .get_collection( + database_id, + tenant_id.as_u64(), + qualified.collection_name(database_id), + )? + .and_then(|stored| { + crate::control::planner::catalog_adapter::document_declared_key(&stored) + }); + Ok(declared_key.unwrap_or_else(|| nodedb_types::DEFAULT_IDENTITY_COLUMN.to_string())) +} + /// Decode the DP-side JSON/msgpack array of `HybridSearchHit`-shaped rows /// (`{doc_id: , : f64, /// vector_rank?, text_rank?}`), resolve each row's `doc_id` surrogate to the diff --git a/nodedb/src/control/server/shared/clone_write/probes.rs b/nodedb/src/control/server/shared/clone_write/probes.rs index 8489dab95..de59fdc6a 100644 --- a/nodedb/src/control/server/shared/clone_write/probes.rs +++ b/nodedb/src/control/server/shared/clone_write/probes.rs @@ -228,6 +228,7 @@ fn with_caller_rls( database_id, &mut plan, &state.rls, + state.credentials.catalog(), scope.auth(), )?; crate::control::planner::redaction_refusal::refuse_unredactable_plan( diff --git a/nodedb/src/control/server/shared/ddl/neutral/query_functions/helpers.rs b/nodedb/src/control/server/shared/ddl/neutral/query_functions/helpers.rs index 5f1ef98c9..5bfdf7394 100644 --- a/nodedb/src/control/server/shared/ddl/neutral/query_functions/helpers.rs +++ b/nodedb/src/control/server/shared/ddl/neutral/query_functions/helpers.rs @@ -89,10 +89,19 @@ pub fn single_result(value: &str) -> Vec { /// typed value for the unwrap and its rows rendered back to JSON for the /// callers, which read fields as JSON. Rows that are not `{id, data}` /// wrapped (already-flat producers) pass through unchanged. +/// +/// The callers read a row's identity under `id`, so a row that lacks `id` +/// gains it here. These rows feed the query functions only and never reach +/// a client as a result set. pub fn unwrap_scan_docs(docs: Vec) -> Result>, DdlError> { let mut rows = Vec::with_capacity(docs.len()); for doc in docs { - push_flat_rows(Value::from(doc), &mut rows).map_err(|e| DdlError::from_error(&e))?; + push_flat_rows( + Value::from(doc), + nodedb_types::DEFAULT_IDENTITY_COLUMN, + &mut rows, + ) + .map_err(|e| DdlError::from_error(&e))?; } Ok(rows.iter().map(row_to_wire_json).collect()) } diff --git a/nodedb/src/control/server/shared/ddl/neutral/read_gate.rs b/nodedb/src/control/server/shared/ddl/neutral/read_gate.rs index 8e8b9a2e6..74957696a 100644 --- a/nodedb/src/control/server/shared/ddl/neutral/read_gate.rs +++ b/nodedb/src/control/server/shared/ddl/neutral/read_gate.rs @@ -160,6 +160,7 @@ impl<'a> CollectionReadGate<'a> { self.scope.database_id(), plan, &self.state.rls, + self.state.credentials.catalog(), self.scope.auth(), ) .map_err(|error| match &error { diff --git a/nodedb/src/control/server/shared/ddl/user_dispatch.rs b/nodedb/src/control/server/shared/ddl/user_dispatch.rs index cd49946ac..c2a817eb2 100644 --- a/nodedb/src/control/server/shared/ddl/user_dispatch.rs +++ b/nodedb/src/control/server/shared/ddl/user_dispatch.rs @@ -200,6 +200,7 @@ async fn authorize_for_identity( scope.database_id(), &mut plan, &state.rls, + state.credentials.catalog(), scope.auth(), )?; crate::control::planner::redaction_refusal::refuse_unredactable_plan( diff --git a/nodedb/src/control/server/sync/kv_handler.rs b/nodedb/src/control/server/sync/kv_handler.rs index 345ca1e72..e9fccb723 100644 --- a/nodedb/src/control/server/sync/kv_handler.rs +++ b/nodedb/src/control/server/sync/kv_handler.rs @@ -133,6 +133,7 @@ impl KvPushDispatcher for SharedStateKvDispatcher<'_> { database_id, &mut plan, &self.shared.rls, + self.shared.credentials.catalog(), request.scope().auth(), )?; diff --git a/nodedb/src/control/trigger/dml_hook.rs b/nodedb/src/control/trigger/dml_hook.rs index 5efdc2ca6..83c8261eb 100644 --- a/nodedb/src/control/trigger/dml_hook.rs +++ b/nodedb/src/control/trigger/dml_hook.rs @@ -377,6 +377,7 @@ pub async fn fetch_old_row( database_id, &mut plan, &state.rls, + state.credentials.catalog(), auth, )?; crate::control::planner::redaction_refusal::refuse_unredactable_plan( diff --git a/nodedb/src/data/executor/core_loop/filter_match.rs b/nodedb/src/data/executor/core_loop/filter_match.rs index 18764b1f6..e00f371d8 100644 --- a/nodedb/src/data/executor/core_loop/filter_match.rs +++ b/nodedb/src/data/executor/core_loop/filter_match.rs @@ -39,31 +39,32 @@ use super::CoreLoop; /// the behavior-flip rule applies: the query fails instead of the row being /// silently excluded. /// -/// `row_key` is the row's storage key. A schemaless collection with no -/// declared `id` field carries its identity only in that key, never in the -/// body, so the body is matched with `id` injected — the same injection +/// `row_key` is the row's storage key. A row that lacks its identity column +/// carries its identity only in that key, so the image is matched with the +/// identity injected under `identity_column` — the same injection /// [`super::super::row_shape::sparse_row_to_doc`] applies to a materialized -/// row, so `WHERE id ...` sees the identity a reader of the same row sees. A -/// strict row already surfaces `id` as a real tuple column, so no injection -/// runs on that arm. +/// row, so a predicate on the identity column sees the identity a reader of +/// the same row sees. A declared-key row holds its key and gains no `id`. pub(in crate::data::executor) fn matches_with_resolved_schema( strict_schema: Option<&StrictSchema>, filters: &[ScanFilter], row_key: &StorageKey, body: &[u8], + identity_column: &str, ) -> Result { - match strict_schema { + let decoded; + let stored: &[u8] = match strict_schema { Some(schema) => match strict_format::binary_tuple_to_msgpack(body, schema) { - Some(msgpack) => ScanFilter::all_match_binary(filters, &msgpack), - None => Ok(false), + Some(msgpack) => { + decoded = msgpack; + &decoded + } + None => return Ok(false), }, - None => { - let identity = row_key.to_identity(); - let with_id = - nodedb_query::msgpack_scan::inject_str_field(body, "id", identity.as_str()); - ScanFilter::all_match_binary(filters, &with_id) - } - } + None => body, + }; + let image = super::super::row_shape::inject_row_identity(stored, row_key, identity_column); + ScanFilter::all_match_binary(filters, &image) } impl CoreLoop { @@ -114,8 +115,15 @@ impl CoreLoop { filters: &'a [ScanFilter], ) -> impl Fn(&StorageKey, &[u8]) -> Result + 'a { let strict_schema = self.resolve_strict_schema(database_id, tid, collection); + let identity_column = self.identity_column(database_id, tid, collection); move |row_key: &StorageKey, body: &[u8]| { - matches_with_resolved_schema(strict_schema.as_ref(), filters, row_key, body) + matches_with_resolved_schema( + strict_schema.as_ref(), + filters, + row_key, + body, + &identity_column, + ) } } } diff --git a/nodedb/src/data/executor/handlers/bulk_dml/update_project.rs b/nodedb/src/data/executor/handlers/bulk_dml/update_project.rs index a864e85b3..4fa1f3483 100644 --- a/nodedb/src/data/executor/handlers/bulk_dml/update_project.rs +++ b/nodedb/src/data/executor/handlers/bulk_dml/update_project.rs @@ -13,6 +13,7 @@ use nodedb_types::columnar::StrictSchema; use crate::data::executor::core_loop::CoreLoop; use crate::data::executor::doc_format; +use crate::data::executor::handlers::identity_guard::IdentitySnapshot; use crate::engine::document::store::StorageKey; use crate::types::{DatabaseId, TenantId}; @@ -136,9 +137,9 @@ impl CoreLoop { let doc_id = doc_id_owned.as_str(); // Decode current value — format depends on storage mode, with the - // storage key attached as `id` for a schemaless row whose body - // carries none, so this image matches the one DELETE's - // write-gate judges. A row the statement matched but cannot + // row identity attached under the collection's identity column for a + // schemaless row whose body lacks it, so this image matches the one + // DELETE's write-gate judges. A row the statement matched but cannot // decode fails the statement rather than under-reporting the // affected count. let mut doc = match strict_schema { @@ -155,17 +156,20 @@ impl CoreLoop { } None => { // `key` is the storage key from the apply set. The - // decoded document's `id` is the row's client-visible - // identity, not the storage key. + // decoded document's identity column holds the row's + // client-visible identity, not the storage key. let identity = key.to_identity(); crate::data::executor::handlers::returning_doc::from_stored_json( ¤t_bytes, &identity, None, + &self.identity_column(database_id, tid, collection), )? } }; + let identity = + IdentitySnapshot::capture(strict_schema, declared_primary_key, updates, &doc); // Feeds the secondary-index SET diff for values the UPDATE drops. let old_doc = doc.clone(); // All assignments see this pre-update snapshot — they don't @@ -203,6 +207,7 @@ impl CoreLoop { declared_primary_key, )?; } + identity.check_unchanged(collection, &doc)?; // Recompute generated columns if any dependency changed. A column // the engine cannot recompute fails the statement. @@ -220,6 +225,13 @@ impl CoreLoop { .map_err(crate::Error::DataPlane)?; } + // A declared numeric column holds the value its type stores, whether + // the assignment was a literal or computed. + crate::data::executor::strict_format::coerce_declared_doc( + &mut doc, + self.declared_columns(&config_key), + )?; + // Re-encode — format depends on storage mode. An encode error // carries its own typed cause, such as a field the strict schema // does not declare. diff --git a/nodedb/src/data/executor/handlers/control/calvin_reply/images.rs b/nodedb/src/data/executor/handlers/control/calvin_reply/images.rs index e43967bb8..3a72b5c1f 100644 --- a/nodedb/src/data/executor/handlers/control/calvin_reply/images.rs +++ b/nodedb/src/data/executor/handlers/control/calvin_reply/images.rs @@ -125,8 +125,15 @@ impl CoreLoop { .iter() .map(|row| (&row.identity, row.bytes.as_slice())) .collect(); - build_stored_rows_payload(spec, rls_filters, strict_schema.as_ref(), &stored) - .map_err(ErrorCode::from) + let identity_column = self.identity_column(at.database_id, at.tid, at.collection); + build_stored_rows_payload( + spec, + rls_filters, + strict_schema.as_ref(), + &identity_column, + &stored, + ) + .map_err(ErrorCode::from) } RowEngine::Kv => { let keys = rows @@ -150,7 +157,9 @@ impl CoreLoop { .zip(rows) .map(|(key, row)| (key, row.bytes.as_slice())) .collect(); - vector_stored_rows_payload(spec, rls_filters, &stored).map_err(ErrorCode::from) + let identity_column = self.identity_column(at.database_id, at.tid, at.collection); + vector_stored_rows_payload(spec, rls_filters, &identity_column, &stored) + .map_err(ErrorCode::from) } RowEngine::Columnar | RowEngine::Timeseries => Err(ErrorCode::Internal { detail: format!( diff --git a/nodedb/src/data/executor/handlers/control/range_scan_versioned.rs b/nodedb/src/data/executor/handlers/control/range_scan_versioned.rs index 0e994b0d5..ba108486b 100644 --- a/nodedb/src/data/executor/handlers/control/range_scan_versioned.rs +++ b/nodedb/src/data/executor/handlers/control/range_scan_versioned.rs @@ -102,6 +102,8 @@ impl CoreLoop { None } }); + let identity_column = + self.identity_column(task.request.database_id.as_u64(), tid, collection); // Predicate: decode each current body, extract `field`, keep in-range // rows. `extract_index_values(_, field, false)` yields the scalar @@ -135,22 +137,17 @@ impl CoreLoop { // The filters evaluate against the normalized msgpack form, // the same encoding every other RLS site filters on — a // strict body is a Binary Tuple until it is decoded here. - // A schemaless row's identity lives only in its storage - // key when its body carries no `id` field, so the - // client-visible identity is injected before the RLS - // check, matching what a reader of the same row sees. + // A row that lacks its identity column carries its + // identity only in its storage key, so the identity is + // injected under that column before the RLS check, the + // image a reader of the same row sees. Some(filters) => match nodedb_types::json_msgpack::json_to_msgpack(&doc) { Ok(mp) => { - let mp = if strict_schema.is_none() { - let identity = doc_id.to_identity(); - nodedb_query::msgpack_scan::inject_str_field( - &mp, - "id", - identity.as_str(), - ) - } else { - mp - }; + let mp = crate::data::executor::row_shape::inject_row_identity( + &mp, + doc_id, + &identity_column, + ); crate::bridge::scan_filter::ScanFilter::all_match_binary(filters, &mp) .unwrap_or(false) } @@ -185,12 +182,7 @@ impl CoreLoop { Ok(rows) => rows, Err(e) => { warn!(core = self.core_id, error = %e, "versioned range scan failed"); - return self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ); + return self.response_error(task, ErrorCode::from(e)); } }; @@ -242,29 +234,18 @@ impl CoreLoop { // Sort ascending by `field` and cap at `limit` — the same ordering the // secondary-index range path yields (index-key order) and the same // truncation the non-bitemporal fallback applies. - if let Err(e) = sort::sort_rows( - &mut rows, - &[nodedb_physical::physical_plan::SortKeySpec::column( - field, true, - )], - ) { - return self.response_error( - task, - ErrorCode::Internal { - detail: format!("in-memory sort failed: {e}"), - }, - ); + let sort_keys = [nodedb_physical::physical_plan::SortKeySpec::column( + field, true, + )]; + let decimal_keys = sort::decimal_sort_keys(&sort_keys, strict_schema.as_ref()); + if let Err(e) = sort::sort_rows(&mut rows, &sort_keys, &decimal_keys) { + return self.response_error(task, ErrorCode::from(e)); } rows.truncate(limit); match response_codec::encode_raw_document_rows(&rows) { Ok(payload) => self.response_with_payload(task, payload), - Err(e) => self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ), + Err(e) => self.response_error(task, ErrorCode::from(e)), } } } diff --git a/nodedb/src/data/executor/handlers/document/read/fetch.rs b/nodedb/src/data/executor/handlers/document/read/fetch.rs index 0e1d2106c..66c05a989 100644 --- a/nodedb/src/data/executor/handlers/document/read/fetch.rs +++ b/nodedb/src/data/executor/handlers/document/read/fetch.rs @@ -54,6 +54,8 @@ impl CoreLoop { // prefix, not an answer. let deadline = crate::data::executor::deadline::DeadlineCheck::for_task(task); let stop = || deadline.expired(); + let identity_column = + self.identity_column(task.request.database_id.as_u64(), tid, collection); match params.mode { DocScanMode::Current => self.fetch_current(task, tid, ¶ms, &deadline), @@ -79,6 +81,7 @@ impl CoreLoop { filter_predicates, doc_id, body, + &identity_column, ) { Ok(b) => b, Err(e) => { @@ -136,6 +139,7 @@ impl CoreLoop { filter_predicates, doc_id, body, + &identity_column, ) { Ok(b) => b, Err(e) => { @@ -214,6 +218,7 @@ impl CoreLoop { self.effective_fetch_limit(params.limit, params.offset, params.full_fetch); let database_id = task.request.database_id.as_u64(); let bitemporal = self.is_bitemporal(database_id, tid, collection); + let identity_column = self.identity_column(database_id, tid, collection); // Resolved from the collection's registered kind, never from the bytes: // a tagged sidecar and a plain document body are both valid MessagePack // maps with the same header, so sniffing necessarily mis-reads one. @@ -246,7 +251,13 @@ impl CoreLoop { } else { value }; - match matches_with_resolved_schema(strict_schema, filter_predicates, key, value) { + match matches_with_resolved_schema( + strict_schema, + filter_predicates, + key, + value, + &identity_column, + ) { Ok(b) => b, Err(e) => { predicate_err.set(Some(e)); @@ -327,6 +338,7 @@ impl CoreLoop { &key, &body, SparseBodyFormatRef::VectorSidecar, + &identity_column, ); (key, mp) }) diff --git a/nodedb/src/data/executor/handlers/document/read/materialize_scan.rs b/nodedb/src/data/executor/handlers/document/read/materialize_scan.rs index 0c87606c7..31845e424 100644 --- a/nodedb/src/data/executor/handlers/document/read/materialize_scan.rs +++ b/nodedb/src/data/executor/handlers/document/read/materialize_scan.rs @@ -5,10 +5,11 @@ //! `doc_id` is the hex-encoded surrogate; `value_bytes` is always standard //! MessagePack — Binary Tuple and vector-primary sidecar sources are //! transcoded here so consumers never re-decide the source format. `value_bytes` -//! also carries an `id` field via `sparse_row_to_doc`, so a Control-Plane -//! filter naming `id` sees the same identity the read paths produce. A -//! `raw_bodies` scan adds no `id`: the clone materializer copies the stored -//! contents, which a `HASH_CHAIN` link covers and a strict schema accepts. +//! also carries the row identity under the collection's identity column via +//! `sparse_row_to_doc`, so a Control-Plane filter sees the same identity the +//! read paths produce. A `raw_bodies` scan adds no identity: the clone +//! materializer copies the stored contents, which a `HASH_CHAIN` link covers +//! and a strict schema accepts. //! Payload: `[next_cursor: bin, entries: [[doc_id, surrogate, value], ...]]`. use nodedb_types::StorageKey; @@ -136,20 +137,22 @@ impl CoreLoop { .unwrap_or_default() }; - // Normalize every body to standard msgpack and inject its `id` here — - // the one place that owns the source format — so no consumer repeats - // the decision or filters a row missing the identity its storage key - // already carries. + // Normalize every body to standard msgpack and inject its identity + // here — the one place that owns the source format — so no consumer + // repeats the decision or filters a row missing the identity its + // storage key already carries. let body_format = self.sparse_body_format(task.request.database_id, TenantId::new(tid), collection); let format_ref = body_format.as_format_ref(); + let identity_column = + self.identity_column(task.request.database_id.as_u64(), tid, collection); // A raw body goes out as stored, with no `id` added: the clone // materializer copies the exact stored contents. for (key, value) in &mut entries { *value = if raw_bodies { sparse_body_to_msgpack(value, format_ref).into_owned() } else { - sparse_row_to_doc(key, value, format_ref).1 + sparse_row_to_doc(key, value, format_ref, &identity_column).1 }; } diff --git a/nodedb/src/data/executor/handlers/document/resolve/apply.rs b/nodedb/src/data/executor/handlers/document/resolve/apply.rs index 6dcdc6d14..302842758 100644 --- a/nodedb/src/data/executor/handlers/document/resolve/apply.rs +++ b/nodedb/src/data/executor/handlers/document/resolve/apply.rs @@ -60,6 +60,11 @@ impl CoreLoop { document_id.as_str(), ), None, + &self.identity_column( + task.request.database_id.as_u64(), + tid, + collection.as_str(), + ), tid, collection.as_str(), ) diff --git a/nodedb/src/data/executor/handlers/document/resolve/bulk.rs b/nodedb/src/data/executor/handlers/document/resolve/bulk.rs index ed56170a2..3eb5b2a07 100644 --- a/nodedb/src/data/executor/handlers/document/resolve/bulk.rs +++ b/nodedb/src/data/executor/handlers/document/resolve/bulk.rs @@ -179,6 +179,7 @@ impl CoreLoop { &stored, &identity, ctx.strict_schema.as_ref(), + &ctx.identity_column, tid, collection, ) @@ -198,12 +199,7 @@ impl CoreLoop { .iter() .map(|(id, body)| (id, body.as_slice())) .collect(); - let response_payload = resolved_response_payload( - returning, - rls_filters, - ctx.strict_schema.as_ref(), - &borrowed, - )?; + let response_payload = resolved_response_payload(returning, rls_filters, &ctx, &borrowed)?; Ok(DocumentResolveOutcome { mutations, response_payload, diff --git a/nodedb/src/data/executor/handlers/document/resolve/context.rs b/nodedb/src/data/executor/handlers/document/resolve/context.rs index 496fe2907..e07b9aa26 100644 --- a/nodedb/src/data/executor/handlers/document/resolve/context.rs +++ b/nodedb/src/data/executor/handlers/document/resolve/context.rs @@ -29,6 +29,9 @@ pub(super) struct DocResolveCtx { /// `Some` exactly when the collection stores Binary Tuples. pub strict_schema: Option, pub bitemporal: bool, + /// The column a row renders its identity under, per + /// `CoreLoop::identity_column`. + pub identity_column: String, } impl CoreLoop { @@ -57,6 +60,7 @@ impl CoreLoop { tid, strict_schema, bitemporal: self.is_bitemporal(database_id, tid, collection), + identity_column: self.identity_column(database_id, tid, collection), } } @@ -154,16 +158,20 @@ pub(super) fn affected_payload(affected: usize) -> Vec { pub(super) fn resolved_response_payload( returning: Option<&ReturningSpec>, rls_filters: &[u8], - strict_schema: Option<&StrictSchema>, + ctx: &DocResolveCtx, rows: &[(&RowIdentity, &[u8])], ) -> Result, ErrorCode> { match returning { - Some(spec) => { - returning_rows::build_stored_rows_payload(spec, rls_filters, strict_schema, rows) - .map_err(|e| ErrorCode::Internal { - detail: format!("RETURNING encode: {e}"), - }) - } + Some(spec) => returning_rows::build_stored_rows_payload( + spec, + rls_filters, + ctx.strict_schema.as_ref(), + &ctx.identity_column, + rows, + ) + .map_err(|e| ErrorCode::Internal { + detail: format!("RETURNING encode: {e}"), + }), None => Ok(affected_payload(rows.len())), } } diff --git a/nodedb/src/data/executor/handlers/document/resolve/point.rs b/nodedb/src/data/executor/handlers/document/resolve/point.rs index cfe8e945f..746a9509e 100644 --- a/nodedb/src/data/executor/handlers/document/resolve/point.rs +++ b/nodedb/src/data/executor/handlers/document/resolve/point.rs @@ -152,6 +152,7 @@ impl CoreLoop { &stored_image, &row_identity, ctx.strict_schema.as_ref(), + &ctx.identity_column, tid, collection, ) @@ -160,7 +161,7 @@ impl CoreLoop { let response_payload = resolved_response_payload( returning, rls_filters, - ctx.strict_schema.as_ref(), + &ctx, &[(&document_identity, stored_image.as_slice())], )?; Ok(DocumentResolveOutcome { @@ -209,12 +210,7 @@ impl CoreLoop { let Some((surrogate, prior)) = prior else { return Ok(DocumentResolveOutcome { mutations: Vec::new(), - response_payload: resolved_response_payload( - returning, - rls_filters, - ctx.strict_schema.as_ref(), - &[], - )?, + response_payload: resolved_response_payload(returning, rls_filters, &ctx, &[])?, }); }; let row_identity = StorageKey::for_surrogate(surrogate).to_identity(); @@ -224,6 +220,7 @@ impl CoreLoop { &prior, &row_identity, ctx.strict_schema.as_ref(), + &ctx.identity_column, tid, collection, ) @@ -232,7 +229,7 @@ impl CoreLoop { let response_payload = resolved_response_payload( returning, rls_filters, - ctx.strict_schema.as_ref(), + &ctx, &[(&document_identity, prior.as_slice())], )?; Ok(DocumentResolveOutcome { diff --git a/nodedb/src/data/executor/handlers/document/resolve/upsert.rs b/nodedb/src/data/executor/handlers/document/resolve/upsert.rs index 93cf6d295..598baf457 100644 --- a/nodedb/src/data/executor/handlers/document/resolve/upsert.rs +++ b/nodedb/src/data/executor/handlers/document/resolve/upsert.rs @@ -83,6 +83,7 @@ impl CoreLoop { &body, &row_identity, None, + &ctx.identity_column, tid, collection, ) @@ -110,7 +111,7 @@ impl CoreLoop { let response_payload = resolved_response_payload( returning, rls_filters, - ctx.strict_schema.as_ref(), + &ctx, &[(&document_identity, stored_image.as_slice())], )?; diff --git a/nodedb/src/data/executor/handlers/identity_guard.rs b/nodedb/src/data/executor/handlers/identity_guard.rs new file mode 100644 index 000000000..9c5e88fdb --- /dev/null +++ b/nodedb/src/data/executor/handlers/identity_guard.rs @@ -0,0 +1,205 @@ +// SPDX-License-Identifier: BUSL-1.1 + +//! A document UPDATE keeps the row's identity column. +//! +//! A document row is stored under the key its identity column held at insert. +//! An UPDATE rewrites the body in place under that key, so a changed identity +//! column makes SQL reads (which read the column) and key reads (which read +//! the key) name different documents. Changing a document's key is a DELETE +//! plus an INSERT under the new key. + +use nodedb_physical::physical_plan::UpdateValue; +use nodedb_types::DEFAULT_IDENTITY_COLUMN; +use nodedb_types::columnar::{SchemaOps, StrictSchema}; + +/// Whether `field` is a column a document row's identity is read from: a +/// strict schema's primary-key column, else a schemaless collection's declared +/// primary key, else `id`. +fn is_identity_column( + strict_schema: Option<&StrictSchema>, + declared_primary_key: Option<&str>, + field: &str, +) -> bool { + match strict_schema { + Some(schema) => schema + .columns() + .iter() + .any(|column| column.primary_key && column.name == field), + None => field == declared_primary_key.unwrap_or(DEFAULT_IDENTITY_COLUMN), + } +} + +/// Whether `updates` assigns any identity column of the collection. +pub(in crate::data::executor) fn assigns_identity( + strict_schema: Option<&StrictSchema>, + declared_primary_key: Option<&str>, + updates: &[(String, UpdateValue)], +) -> bool { + updates + .iter() + .any(|(field, _)| is_identity_column(strict_schema, declared_primary_key, field)) +} + +/// The pre-update values of the identity columns an UPDATE assigns. Empty, +/// and allocation-free, when the UPDATE assigns none. +pub(in crate::data::executor) struct IdentitySnapshot { + columns: Vec<(String, Option)>, + /// The identity is a declared primary key, whose NOT NULL rule refuses a + /// NULL or omitted post-image value with its own constraint. + declared_key: bool, +} + +impl IdentitySnapshot { + /// Record the identity columns `updates` assigns, as `before` holds them. + pub(in crate::data::executor) fn capture( + strict_schema: Option<&StrictSchema>, + declared_primary_key: Option<&str>, + updates: &[(String, UpdateValue)], + before: &serde_json::Value, + ) -> Self { + let columns = updates + .iter() + .filter(|(field, _)| is_identity_column(strict_schema, declared_primary_key, field)) + .map(|(field, _)| (field.clone(), before.get(field).cloned())) + .collect(); + Self { + columns, + declared_key: strict_schema.is_some() || declared_primary_key.is_some(), + } + } + + /// Refuse `after` when an assigned identity column differs from its + /// captured value. A NULL or omitted declared key is left to the NOT NULL + /// rule, so that refusal keeps its own constraint. + pub(in crate::data::executor) fn check_unchanged( + &self, + collection: &str, + after: &serde_json::Value, + ) -> crate::Result<()> { + for (column, before) in &self.columns { + let value = after.get(column); + if self.declared_key && matches!(value, None | Some(serde_json::Value::Null)) { + continue; + } + if before.as_ref() != value { + return Err(identity_change_refused(collection, column)); + } + } + Ok(()) + } +} + +fn identity_change_refused(collection: &str, column: &str) -> crate::Error { + crate::Error::RejectedConstraint { + collection: collection.to_string(), + constraint: "primary_key_immutable".to_string(), + detail: format!( + "UPDATE cannot change primary key '{column}' of collection '{collection}': \ + the row stays stored under its current key; DELETE the row and INSERT it \ + under the new key" + ), + } +} + +#[cfg(test)] +mod tests { + use super::*; + use nodedb_types::columnar::{ColumnDef, ColumnType}; + use serde_json::json; + + fn literal(value: &nodedb_types::Value) -> UpdateValue { + UpdateValue::Literal(nodedb_types::value_to_msgpack(value).expect("encode")) + } + + fn set(field: &str) -> Vec<(String, UpdateValue)> { + vec![( + field.to_string(), + literal(&nodedb_types::Value::String("x".into())), + )] + } + + fn strict() -> StrictSchema { + StrictSchema::new(vec![ + ColumnDef::required("k", ColumnType::String).with_primary_key(), + ColumnDef::nullable("v", ColumnType::String), + ]) + .expect("schema") + } + + #[test] + fn schemaless_identity_is_id_unless_a_key_is_declared() { + assert!(assigns_identity(None, None, &set("id"))); + assert!(!assigns_identity(None, None, &set("v"))); + assert!(assigns_identity(None, Some("sku"), &set("sku"))); + assert!(!assigns_identity(None, Some("sku"), &set("id"))); + } + + #[test] + fn strict_identity_is_the_schema_primary_key() { + let schema = strict(); + assert!(assigns_identity(Some(&schema), None, &set("k"))); + assert!(!assigns_identity(Some(&schema), None, &set("v"))); + } + + #[test] + fn a_changed_identity_is_refused() { + let updates = set("id"); + let snapshot = IdentitySnapshot::capture(None, None, &updates, &json!({"id": "a"})); + match snapshot.check_unchanged("docs", &json!({"id": "b"})) { + Err(crate::Error::RejectedConstraint { + constraint, detail, .. + }) => { + assert_eq!(constraint, "primary_key_immutable"); + assert!(detail.contains("'id'"), "names the column: {detail}"); + } + other => panic!("expected RejectedConstraint, got {other:?}"), + } + } + + #[test] + fn an_identity_newly_written_to_an_id_less_body_is_refused() { + let updates = set("id"); + let snapshot = IdentitySnapshot::capture(None, None, &updates, &json!({"v": 1})); + assert!( + snapshot + .check_unchanged("docs", &json!({"id": "b", "v": 1})) + .is_err() + ); + } + + #[test] + fn the_same_identity_is_accepted() { + let updates = set("id"); + let snapshot = IdentitySnapshot::capture(None, None, &updates, &json!({"id": "a"})); + snapshot + .check_unchanged("docs", &json!({"id": "a", "v": 2})) + .expect("assigning the current key keeps the identity"); + } + + #[test] + fn a_null_declared_key_is_left_to_the_not_null_rule() { + let schema = strict(); + let updates = set("k"); + let snapshot = IdentitySnapshot::capture(Some(&schema), None, &updates, &json!({"k": "a"})); + snapshot + .check_unchanged("docs", &json!({"k": null})) + .expect("the NOT NULL rule owns a NULL declared key"); + let updates = set("id"); + let snapshot = IdentitySnapshot::capture(None, None, &updates, &json!({"id": "a"})); + assert!( + snapshot + .check_unchanged("docs", &json!({"id": null})) + .is_err(), + "an undeclared `id` has no NOT NULL rule, so NULL is a changed identity" + ); + } + + #[test] + fn an_update_of_other_columns_is_never_checked() { + let updates = set("v"); + let snapshot = IdentitySnapshot::capture(None, None, &updates, &json!({"id": "a"})); + snapshot + .check_unchanged("docs", &json!({"id": "changed-elsewhere"})) + .expect("only assigned identity columns are checked"); + } +} diff --git a/nodedb/src/data/executor/handlers/kv/rls.rs b/nodedb/src/data/executor/handlers/kv/rls.rs index d3d3e3c5d..2afe79781 100644 --- a/nodedb/src/data/executor/handlers/kv/rls.rs +++ b/nodedb/src/data/executor/handlers/kv/rls.rs @@ -45,5 +45,15 @@ pub(in crate::data::executor) fn admit_kv_row( // own identity, taken verbatim. let key_display = String::from_utf8_lossy(key); let identity = crate::engine::document::store::RowIdentity::from_user_key(key_display.as_ref()); - rls_write_gate::admit_stored_row(rls_write_check, body, &identity, None, tid, collection) + // A KV image names its key `id`, as the KV write policy has always + // read it. KV rows never take the document identity column. + rls_write_gate::admit_stored_row( + rls_write_check, + body, + &identity, + None, + nodedb_types::DEFAULT_IDENTITY_COLUMN, + tid, + collection, + ) } diff --git a/nodedb/src/data/executor/handlers/merge_orchestrated/apply/insert_rows.rs b/nodedb/src/data/executor/handlers/merge_orchestrated/apply/insert_rows.rs index a3389f69d..12c1b7c64 100644 --- a/nodedb/src/data/executor/handlers/merge_orchestrated/apply/insert_rows.rs +++ b/nodedb/src/data/executor/handlers/merge_orchestrated/apply/insert_rows.rs @@ -83,6 +83,7 @@ impl CoreLoop { balanced_entries, returned_docs, } = tally; + let identity_column = self.identity_column(database_id, tid, collection); // Every NOT-MATCHED INSERT row of a HASH_CHAIN target is a chain link. // The head advances in the shared transaction, and the pre-image goes @@ -205,7 +206,7 @@ impl CoreLoop { outcome.bitemporal_sys_from_ms, )); if returning { - match returning_doc(&ins.body, &storage_key) { + match returning_doc(&ins.body, &storage_key, &identity_column) { Ok(doc) => returned_docs.push(doc), Err(e) => { return Err(self.abort_merge_apply(MergeAbort { diff --git a/nodedb/src/data/executor/handlers/merge_orchestrated/apply/orchestrate.rs b/nodedb/src/data/executor/handlers/merge_orchestrated/apply/orchestrate.rs index ec2311121..99880a485 100644 --- a/nodedb/src/data/executor/handlers/merge_orchestrated/apply/orchestrate.rs +++ b/nodedb/src/data/executor/handlers/merge_orchestrated/apply/orchestrate.rs @@ -70,9 +70,14 @@ impl CoreLoop { // Gate every arm on the target's write policy BEFORE the apply // transaction opens, so a rejected row leaves nothing written and // nothing to unwind. - if let Err(e) = - gate_merge_arms(&plan, params.rls_write_check, tid, params.target_collection) - { + let identity_column = self.identity_column(database_id, tid, params.target_collection); + if let Err(e) = gate_merge_arms( + &plan, + params.rls_write_check, + &identity_column, + tid, + params.target_collection, + ) { return self.response_error(task, e); } @@ -288,23 +293,13 @@ impl CoreLoop { let mut response = if let Some(spec) = params.returning { match returning_rows::build_rows_payload(spec, params.rls_filters, &returned_docs) { Ok(payload) => self.response_with_payload(task, payload), - Err(e) => self.response_error( - task, - ErrorCode::Internal { - detail: format!("RETURNING encode: {e}"), - }, - ), + Err(e) => self.response_error(task, ErrorCode::from(e)), } } else { let result = serde_json::json!({ "affected": affected }); match encode_json_as_msgpack(&result) { Ok(payload) => self.response_with_payload(task, payload), - Err(e) => self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ), + Err(e) => self.response_error(task, ErrorCode::from(e)), } }; response.write_set = write_set; diff --git a/nodedb/src/data/executor/handlers/merge_orchestrated/plan.rs b/nodedb/src/data/executor/handlers/merge_orchestrated/plan.rs index 2a11a1eb6..2398bdab8 100644 --- a/nodedb/src/data/executor/handlers/merge_orchestrated/plan.rs +++ b/nodedb/src/data/executor/handlers/merge_orchestrated/plan.rs @@ -14,6 +14,7 @@ use nodedb_physical::physical_plan::document::merge_types::{ MergeActionOp, MergeClauseKind as MergeClauseKindOp, }; +use super::super::identity_guard::IdentitySnapshot; use super::super::merge::MergeParams; use super::super::merge_helpers::{ build_insert_doc, build_merged, build_update_doc, find_arm, json_to_str, @@ -55,21 +56,22 @@ pub(super) struct MergePlanActions { pub(super) inserts: Vec, } -/// Decode a stored target row into JSON, with `id` injected for a schemaless -/// row whose body carries none. Fails rather than skipping — a row the -/// classifier can't read is not "absent", and treating it as absent inserts a -/// duplicate of a row that already exists. +/// Decode a stored target row into JSON, with its identity injected under +/// `identity_column` for a schemaless row whose body lacks that column. Fails +/// rather than skipping — a row the classifier can't read is not "absent", +/// and treating it as absent inserts a duplicate of a row that already exists. /// -/// A schemaless collection with no declared `id` field carries its identity -/// only in the row's storage key, never in the body — so a MERGE arm's -/// `AND id ...` condition, matched by [`find_arm`], must see the row's -/// client-visible identity injected here, never the raw storage key. A -/// strict row already surfaces `id` as a real tuple column, so injection -/// only runs on the schemaless arm. +/// A schemaless row that lacks its identity column carries its identity only +/// in the row's storage key, never in the body — so a MERGE arm's condition +/// on that column, matched by [`find_arm`], must see the row's client-visible +/// identity injected here, never the raw storage key. A declared-key row +/// holds its key and gains no `id`. A strict row surfaces its key as a real +/// tuple column, so injection only runs on the schemaless arm. fn decode_target( identity: &RowIdentity, bytes: &[u8], strict_schema: &Option, + identity_column: &str, ) -> crate::Result { let mut doc = doc_format::decode_document_or_binary_tuple( bytes, @@ -78,10 +80,10 @@ fn decode_target( )?; if strict_schema.is_none() && let Some(obj) = doc.as_object_mut() - && !obj.contains_key("id") + && !obj.contains_key(identity_column) { obj.insert( - "id".to_string(), + identity_column.to_string(), serde_json::Value::String(identity.as_str().to_string()), ); } @@ -110,6 +112,7 @@ impl CoreLoop { let strict_schema = self.merge_strict_schema(database_id, tid, params.target_collection); let target_docs = self.collect_target_docs(database_id, tid, params.target_collection, txn_id)?; + let identity_column = self.identity_column(database_id, tid, params.target_collection); let mut updates: Vec = Vec::new(); let mut deletes: Vec = Vec::new(); @@ -121,7 +124,7 @@ impl CoreLoop { for (key, bytes) in &target_docs { let key = *key; let identity = key.to_identity(); - let target_doc = decode_target(&identity, bytes, &strict_schema)?; + let target_doc = decode_target(&identity, bytes, &strict_schema, &identity_column)?; let join_val = target_doc .get(params.target_join_col) .map(json_to_str) @@ -162,6 +165,13 @@ impl CoreLoop { upd, pk, )?; + IdentitySnapshot::capture( + strict_schema.as_ref(), + params.declared_primary_key, + upd, + &target_doc, + ) + .check_unchanged(params.target_collection, &updated)?; updates.push(MergeUpdate { key, body: encode_doc_body(&updated), diff --git a/nodedb/src/data/executor/handlers/point/get.rs b/nodedb/src/data/executor/handlers/point/get.rs index 3e1f65d03..ecb574d1f 100644 --- a/nodedb/src/data/executor/handlers/point/get.rs +++ b/nodedb/src/data/executor/handlers/point/get.rs @@ -84,12 +84,7 @@ impl CoreLoop { Ok(Some(data)) => data, Ok(None) => return self.response_with_payload(task, Vec::new()), Err(e) => { - return self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ); + return self.response_error(task, ErrorCode::from(e)); } } } else if let Some(overlay_data) = self.overlay_point_lookup( @@ -126,12 +121,7 @@ impl CoreLoop { Ok(None) => return self.response_with_payload(task, Vec::new()), Err(e) => { tracing::warn!(core = self.core_id, error = %e, "sparse get failed"); - return self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ); + return self.response_error(task, ErrorCode::from(e)); } } } @@ -149,15 +139,22 @@ impl CoreLoop { // actually rewritten yields an owned buffer, and only then is `data` // superseded. // - // RLS reads a second image with `id` injected: a schemaless row with - // no declared `id` field carries its identity only in `row_key`, so a - // policy naming `id` reads it as absent otherwise. The client still - // gets the stored body, which never gains a field it did not have. + // RLS reads a second image with the identity under the collection's + // identity column, the image a scan of the same row returns: a row + // that lacks the column carries its identity only in the storage key, + // so a policy naming the column reads it as absent otherwise. A + // declared-key row gains no `id`. The client still gets the stored + // body, which never gains a field it did not have. let transcoded = { let normalized = sparse_body_to_msgpack(&data, body_format.as_format_ref()); if !rls_filters.is_empty() { - let (_, gated) = - sparse_row_to_doc(&storage_key, &data, body_format.as_format_ref()); + let identity_column = self.identity_column(database_id, tid, collection); + let (_, gated) = sparse_row_to_doc( + &storage_key, + &data, + body_format.as_format_ref(), + &identity_column, + ); if !super::super::rls_eval::rls_check_msgpack_bytes(rls_filters, &gated) { return self.response_with_payload(task, Vec::new()); } diff --git a/nodedb/src/data/executor/handlers/point/update/exec.rs b/nodedb/src/data/executor/handlers/point/update/exec.rs index 01f1415b0..5e396b166 100644 --- a/nodedb/src/data/executor/handlers/point/update/exec.rs +++ b/nodedb/src/data/executor/handlers/point/update/exec.rs @@ -218,11 +218,13 @@ impl CoreLoop { // Placed after the generated columns are recomputed — a policy // may reference one — and before any store or index is touched, // so a rejected row leaves nothing behind. + let identity_column = self.identity_column(database_id, tid, collection); if let Err(e) = rls_write_gate::admit_stored_row( rls_write_check, &updated_bytes, &document_identity, strict_schema.as_ref(), + &identity_column, tid, collection, ) { @@ -288,12 +290,13 @@ impl CoreLoop { // no pre-dispatch record for a PointUpdate. let mut response = if let Some(spec) = returning { // Post-update image, decoded in the collection's - // storage mode; the user-visible key only fills in - // as `id` when the row declares none of its own. + // storage mode; the user-visible key fills in under + // the identity column only when the row lacks it. let doc = match returning_doc::from_stored( &updated_bytes, &document_identity, strict_schema.as_ref(), + &identity_column, ) { Ok(doc) => doc, Err(e) => return self.response_error(task, e), @@ -317,7 +320,7 @@ impl CoreLoop { }; // A versioned row landed at the system time the // update encoded it with, valid for all time. - response.write_set = vec![self.stored_row_image( + response.write_set = match self.stored_row_image( StoredRow { database_id, tid, @@ -327,7 +330,10 @@ impl CoreLoop { }, &updated_bytes, bitemporal.then_some(sys_from_for_encode), - )]; + ) { + Ok(image) => vec![image], + Err(e) => return self.response_error(task, e), + }; // Derived target rows live in a DIFFERENT collection // than this statement's, so each carries its own // `Some(collection)` and homes to that collection's diff --git a/nodedb/src/data/executor/handlers/recursive.rs b/nodedb/src/data/executor/handlers/recursive.rs index 574cbb96b..64e9033fe 100644 --- a/nodedb/src/data/executor/handlers/recursive.rs +++ b/nodedb/src/data/executor/handlers/recursive.rs @@ -120,15 +120,18 @@ impl CoreLoop { }; // Convert raw stored bytes to the msgpack form the CTE steps compare - // on. A schemaless row with no declared `id` field carries its - // identity only in the storage key, never in the body, so the row - // image must inject it before any predicate runs — otherwise - // `id IS NULL` and RETURNING rows both lose the identity. + // on. A row that lacks its identity column carries its identity only + // in the storage key, so the row image injects it under that column + // before any predicate runs. A declared-key row already holds its + // key and gains no `id`. + let identity_column = + self.identity_column(task.request.database_id.as_u64(), tid, collection); let to_msgpack = |doc_id: &nodedb_types::StorageKey, value: &[u8]| -> Vec { crate::data::executor::scan_normalize::sparse_row_to_doc( doc_id, value, body_format.as_format_ref(), + &identity_column, ) .1 }; diff --git a/nodedb/src/data/executor/handlers/returning_doc.rs b/nodedb/src/data/executor/handlers/returning_doc.rs index 09586282b..ef170906e 100644 --- a/nodedb/src/data/executor/handlers/returning_doc.rs +++ b/nodedb/src/data/executor/handlers/returning_doc.rs @@ -11,10 +11,11 @@ //! one: it reads the tuple's leading byte as a scalar and *succeeds*, yielding //! a document with every real column missing. A storage-mode-blind decode //! therefore ships plausible-looking garbage to the client instead of failing. -//! - **The storage key becomes `id` only when the body carries no `id`.** A -//! collection with a declared `id` primary key is authoritative for its own -//! key; overwriting it with the surrogate hex storage key returns a value the -//! client never wrote and cannot use to address the row. +//! - **The identity lands under the collection's identity column, and only +//! when the body lacks that column.** The identity column is the declared +//! key, else `id`, per `CoreLoop::identity_column`. A body that holds its key +//! is authoritative for it, and a declared-key row never gains an `id` +//! beside its key. use nodedb_types::Value; use nodedb_types::columnar::StrictSchema; @@ -23,17 +24,21 @@ use crate::data::executor::doc_format; use crate::data::executor::strict_format; use crate::engine::document::store::RowIdentity; -/// Set `id` to the row's client-visible identity unless the document already -/// carries one. +/// Set `identity_column` to the row's client-visible identity unless the +/// document already carries that column. /// /// For callers that already hold the decoded document (the update paths /// re-project the image they just built rather than re-reading storage). -pub(in crate::data::executor) fn attach_row_id(doc: &mut Value, identity: &RowIdentity) { +pub(in crate::data::executor) fn attach_row_id( + doc: &mut Value, + identity: &RowIdentity, + identity_column: &str, +) { if let Value::Object(obj) = doc - && !obj.contains_key("id") + && !obj.contains_key(identity_column) { obj.insert( - "id".to_string(), + identity_column.to_string(), Value::String(identity.as_str().to_string()), ); } @@ -54,13 +59,14 @@ pub(in crate::data::executor) fn from_stored( body: &[u8], identity: &RowIdentity, strict_schema: Option<&StrictSchema>, + identity_column: &str, ) -> crate::Result { let mut doc = match strict_schema { Some(schema) => strict_format::binary_tuple_to_row_value(body, schema) .ok_or_else(|| undecodable(identity, body.len()))?, - None => doc_format::decode_document_value(&with_identity(body, identity))?, + None => doc_format::decode_document_value(&with_identity(body, identity, identity_column))?, }; - attach_row_id(&mut doc, identity); + attach_row_id(&mut doc, identity, identity_column); Ok(doc) } @@ -75,29 +81,29 @@ pub(in crate::data::executor) fn from_stored_json( body: &[u8], identity: &RowIdentity, strict_schema: Option<&StrictSchema>, + identity_column: &str, ) -> crate::Result { let mut doc = match strict_schema { Some(schema) => strict_format::binary_tuple_to_json(body, schema) .ok_or_else(|| undecodable(identity, body.len()))?, - None => doc_format::decode_document(&with_identity(body, identity))?, + None => doc_format::decode_document(&with_identity(body, identity, identity_column))?, }; if let serde_json::Value::Object(obj) = &mut doc - && !obj.contains_key("id") + && !obj.contains_key(identity_column) { obj.insert( - "id".to_string(), + identity_column.to_string(), serde_json::Value::String(identity.as_str().to_string()), ); } Ok(doc) } -/// The schemaless body with the storage identity injected as `id`. -/// `inject_str_field` honours an existing `id` and wraps a non-map body as -/// `{id, value}`, which is the shape schemaless callers have always emitted -/// for a body that is not a document map. -fn with_identity(body: &[u8], identity: &RowIdentity) -> Vec { - nodedb_query::msgpack_scan::inject_str_field(body, "id", identity.as_str()) +/// The schemaless body with the row identity injected under +/// `identity_column`. `inject_str_field` honours an existing value and wraps +/// a non-map body as `{, value}`. +fn with_identity(body: &[u8], identity: &RowIdentity, identity_column: &str) -> Vec { + nodedb_query::msgpack_scan::inject_str_field(body, identity_column, identity.as_str()) } fn undecodable(identity: &RowIdentity, body_len: usize) -> crate::Error { diff --git a/nodedb/src/data/executor/handlers/returning_rows.rs b/nodedb/src/data/executor/handlers/returning_rows.rs index b45ca7f86..1bf54eda6 100644 --- a/nodedb/src/data/executor/handlers/returning_rows.rs +++ b/nodedb/src/data/executor/handlers/returning_rows.rs @@ -28,15 +28,18 @@ impl CoreLoop { /// Build this task's `RETURNING` response from the rows it just stored. /// The single exit every insert-family handler uses, so the read gate and /// decode can never be applied on one path and skipped on another. + /// `identity_column` is the column each row renders its identity under, + /// per `CoreLoop::identity_column`. pub(in crate::data::executor) fn stored_returning_response( &self, task: &ExecutionTask, spec: &ReturningSpec, rls_filters: &[u8], strict_schema: Option<&StrictSchema>, + identity_column: &str, rows: &[StoredRow<'_>], ) -> Response { - match build_stored_rows_payload(spec, rls_filters, strict_schema, rows) { + match build_stored_rows_payload(spec, rls_filters, strict_schema, identity_column, rows) { Ok(payload) => self.response_with_payload(task, payload), Err(e) => self.response_error( task, @@ -160,15 +163,17 @@ impl CoreLoop { /// [`SparseBodyFormatRef::VectorSidecar`] — the same converter `SELECT` /// uses. The sidecar is `zerompk` TAGGED bytes an ordinary document /// decode misreads (`"alice"` comes back as `[4,"alice"]`), so the - /// format must be this literal, not re-decided. + /// format must be this literal, not re-decided. `identity_column` is the + /// column each row renders its identity under. pub(in crate::data::executor) fn vector_stored_returning_response( &self, task: &ExecutionTask, spec: &ReturningSpec, rls_filters: &[u8], + identity_column: &str, rows: &[VectorStoredRow<'_>], ) -> Response { - match vector_stored_rows_payload(spec, rls_filters, rows) { + match vector_stored_rows_payload(spec, rls_filters, identity_column, rows) { Ok(payload) => self.response_with_payload(task, payload), Err(e) => self.response_error(task, e), } @@ -182,6 +187,7 @@ impl CoreLoop { pub(in crate::data::executor) fn vector_stored_rows_payload( spec: &ReturningSpec, rls_filters: &[u8], + identity_column: &str, rows: &[VectorStoredRow<'_>], ) -> crate::Result> { // An unreadable sidecar fails the statement: an empty row set here @@ -189,7 +195,12 @@ pub(in crate::data::executor) fn vector_stored_rows_payload( let docs: Vec = rows .iter() .map(|(row_key, sidecar)| { - let (_id, mp) = sparse_row_to_doc(row_key, sidecar, SparseBodyFormatRef::VectorSidecar); + let (_id, mp) = sparse_row_to_doc( + row_key, + sidecar, + SparseBodyFormatRef::VectorSidecar, + identity_column, + ); doc_format::decode_document_value(&mp) }) .collect::>>()?; @@ -207,11 +218,14 @@ pub(in crate::data::executor) fn build_stored_rows_payload( spec: &ReturningSpec, rls_filters: &[u8], strict_schema: Option<&StrictSchema>, + identity_column: &str, rows: &[StoredRow<'_>], ) -> crate::Result> { let docs: Vec = rows .iter() - .map(|(doc_id, body)| returning_doc::from_stored(body, doc_id, strict_schema)) + .map(|(doc_id, body)| { + returning_doc::from_stored(body, doc_id, strict_schema, identity_column) + }) .collect::>>()?; build_rows_payload(spec, rls_filters, &docs) } diff --git a/nodedb/src/data/executor/handlers/spatial/full_scan.rs b/nodedb/src/data/executor/handlers/spatial/full_scan.rs index 66c1e590c..552476ed2 100644 --- a/nodedb/src/data/executor/handlers/spatial/full_scan.rs +++ b/nodedb/src/data/executor/handlers/spatial/full_scan.rs @@ -7,7 +7,7 @@ use crate::bridge::scan_filter::ScanFilter; use crate::data::executor::core_loop::CoreLoop; use crate::data::executor::doc_format; use crate::data::executor::handlers::spatial_refine::{ - apply_predicate, extract_geometry, project_doc, + SpatialHit, apply_predicate, extract_geometry, hit_rows, project_doc, }; use crate::data::executor::handlers::transaction::overlay::SpatialOverlayMergeParams; use crate::data::executor::response_codec; @@ -61,15 +61,12 @@ impl CoreLoop { ) { Ok(e) => e, Err(e) => { - return self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ); + return self.response_error(task, ErrorCode::from(e)); } }; + let identity_column = + self.identity_column(task.request.database_id.as_u64(), tid, collection); let mut results = Vec::new(); for (doc_id, doc_bytes) in &entries { if results.len() >= limit { @@ -115,7 +112,10 @@ impl CoreLoop { } } - results.push(project_doc(&doc, doc_id, projection)); + results.push(SpatialHit::new( + doc_id, + project_doc(&doc, doc_id, projection, &identity_column), + )); } if let Some(txn_id) = task.request.txn_id { @@ -142,14 +142,9 @@ impl CoreLoop { } } - match response_codec::encode_value_vec(&results) { + match response_codec::encode_value_vec(&hit_rows(results)) { Ok(payload) => self.response_with_payload(task, payload), - Err(e) => self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ), + Err(e) => self.response_error(task, ErrorCode::from(e)), } } } diff --git a/nodedb/src/data/executor/handlers/spatial/rtree_scan.rs b/nodedb/src/data/executor/handlers/spatial/rtree_scan.rs index e8b63665a..5c77fd621 100644 --- a/nodedb/src/data/executor/handlers/spatial/rtree_scan.rs +++ b/nodedb/src/data/executor/handlers/spatial/rtree_scan.rs @@ -8,7 +8,7 @@ use crate::data::executor::core_loop::CoreLoop; use crate::data::executor::doc_format; use crate::data::executor::handlers::columnar_read::filter::decode_rls_filters; use crate::data::executor::handlers::spatial_refine::{ - apply_predicate, expand_bbox, extract_geometry, project_doc, + SpatialHit, apply_predicate, expand_bbox, extract_geometry, hit_rows, project_doc, }; use crate::data::executor::handlers::transaction::overlay::SpatialOverlayMergeParams; use crate::data::executor::response_codec; @@ -139,12 +139,7 @@ impl CoreLoop { None => { return match response_codec::encode_value_vec(&[]) { Ok(payload) => self.response_with_payload(task, payload), - Err(e) => self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ), + Err(e) => self.response_error(task, ErrorCode::from(e)), }; } }; @@ -172,6 +167,7 @@ impl CoreLoop { // routing and is correct regardless of which engine backs the row. let database_id = db_id.as_u64(); let body_format = self.sparse_body_format(db_id, tid_id, collection); + let identity_column = self.identity_column(database_id, tid, collection); // Lazily-built columnar id → document map, populated on the first // candidate that is absent from the sparse store (i.e. a columnar-family @@ -238,6 +234,7 @@ impl CoreLoop { &key, &raw, body_format.as_format_ref(), + &identity_column, ); // A candidate skipped here silently drops out of the // spatial result set, which reads as "no row matched the @@ -258,12 +255,7 @@ impl CoreLoop { columnar_docs = Some(rows.into_iter().collect()); } Err(e) => { - return self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ); + return self.response_error(task, ErrorCode::from(e)); } } } @@ -310,7 +302,10 @@ impl CoreLoop { } } - results.push(project_doc(&doc, &doc_id, projection)); + results.push(SpatialHit::new( + &doc_id, + project_doc(&doc, &doc_id, projection, &identity_column), + )); } if let Some(txn_id) = task.request.txn_id @@ -332,14 +327,9 @@ impl CoreLoop { return self.response_error(task, e); } - match response_codec::encode_value_vec(&results) { + match response_codec::encode_value_vec(&hit_rows(results)) { Ok(payload) => self.response_with_payload(task, payload), - Err(e) => self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ), + Err(e) => self.response_error(task, ErrorCode::from(e)), } } } diff --git a/nodedb/src/data/executor/handlers/spatial_refine.rs b/nodedb/src/data/executor/handlers/spatial_refine.rs index c5a1bd148..67e9faba3 100644 --- a/nodedb/src/data/executor/handlers/spatial_refine.rs +++ b/nodedb/src/data/executor/handlers/spatial_refine.rs @@ -56,24 +56,57 @@ pub(in crate::data::executor) fn apply_predicate( } } +/// One spatial scan result: the projected row and the surrogate it was read +/// under. The overlay merge keys rows by `surrogate`, never by a row field, +/// so the row names its identity by the collection's identity column alone. +pub(in crate::data::executor) struct SpatialHit { + /// The row's surrogate, `None` when its `doc_id` is no hex surrogate. + pub surrogate: Option, + pub row: Value, +} + +impl SpatialHit { + /// A hit read under `doc_id`: a document row's hex storage key, or a + /// columnar-family row's `id` value. + pub(in crate::data::executor) fn new(doc_id: &str, row: Value) -> Self { + Self { + surrogate: u32::from_str_radix(doc_id, 16).ok(), + row, + } + } +} + +/// The rows of `hits`, in order, for the response encoder. +pub(in crate::data::executor) fn hit_rows(hits: Vec) -> Vec { + hits.into_iter().map(|hit| hit.row).collect() +} + /// Apply projection to a document, returning `nodedb_types::Value`. +/// +/// The row names its identity under `identity_column`: the collection's +/// declared key, else `id`. A document that holds that column keeps its own +/// value. One that lacks it gains `doc_id` there. A declared-key row never +/// gains an `id` beside its key. pub(in crate::data::executor) fn project_doc( doc: &Value, doc_id: &str, projection: &[String], + identity_column: &str, ) -> Value { + let identity = doc + .get(identity_column) + .cloned() + .unwrap_or_else(|| Value::String(doc_id.to_string())); if projection.is_empty() { - // Add id if not present. if let Value::Object(mut map) = doc.clone() { - map.entry("id".to_string()) - .or_insert(Value::String(doc_id.to_string())); + map.entry(identity_column.to_string()).or_insert(identity); Value::Object(map) } else { doc.clone() } } else { let mut map = std::collections::HashMap::new(); - map.insert("id".to_string(), Value::String(doc_id.to_string())); + map.insert(identity_column.to_string(), identity); for col in projection { if let Some(v) = doc.get(col) { map.insert(col.clone(), v.clone()); diff --git a/nodedb/src/data/executor/handlers/transaction/overlay/merge.rs b/nodedb/src/data/executor/handlers/transaction/overlay/merge.rs index 75022a00f..f65e1a1aa 100644 --- a/nodedb/src/data/executor/handlers/transaction/overlay/merge.rs +++ b/nodedb/src/data/executor/handlers/transaction/overlay/merge.rs @@ -353,11 +353,19 @@ impl CoreLoop { // passes finish. let predicate_err: std::cell::Cell> = std::cell::Cell::new(None); + let identity_column = + self.identity_column(coll_key.0.as_u64(), coll_key.1.as_u64(), &coll_key.2); let residual_matches = |row_key: &StorageKey, body: &[u8]| -> bool { if residual.is_empty() { return true; } - match matches_with_resolved_schema(strict_schema, residual, row_key, body) { + match matches_with_resolved_schema( + strict_schema, + residual, + row_key, + body, + &identity_column, + ) { Ok(b) => b, Err(e) => { predicate_err.set(Some(e)); diff --git a/nodedb/src/data/executor/handlers/transaction/overlay/spatial_merge.rs b/nodedb/src/data/executor/handlers/transaction/overlay/spatial_merge.rs index b25304305..b21810cf7 100644 --- a/nodedb/src/data/executor/handlers/transaction/overlay/spatial_merge.rs +++ b/nodedb/src/data/executor/handlers/transaction/overlay/spatial_merge.rs @@ -5,9 +5,9 @@ //! ST_Within/ST_DWithin(...)` scan observes the transaction's own //! uncommitted spatial-row writes (read-your-own-writes). //! -//! Row identity is the hex-encoded surrogate carried in each result row's -//! `"id"` field (`project_doc`'s shape) — the same identity -//! [`super::columnar_merge`] uses, since a mainstream SQL `INSERT INTO +//! Row identity is the surrogate each [`SpatialHit`] carries beside its row — +//! the same identity [`super::columnar_merge`] uses, since a mainstream SQL +//! `INSERT INTO //! VALUES(...)` stages through `ColumnarOp::Insert` //! (`stage_columnar_insert`), not `SpatialOp::Insert`. //! @@ -39,7 +39,7 @@ use crate::bridge::scan_filter::ScanFilter; use crate::data::executor::core_loop::CoreLoop; use crate::data::executor::handlers::columnar_read::convert::row_to_projected_value; use crate::data::executor::handlers::spatial_refine::{ - apply_predicate, extract_geometry, project_doc, + SpatialHit, apply_predicate, extract_geometry, project_doc, }; use crate::data::executor::handlers::transaction::overlay::Staged; use crate::engine::document::store::StorageKey; @@ -84,26 +84,14 @@ fn decode_staged_spatial_row( }) } -/// Extract the hex-surrogate identity from a projected spatial result row -/// (`project_doc`'s `{"id": "", ...}` shape). -fn row_surrogate(row: &Value) -> Option { - match row { - Value::Object(map) => match map.get("id") { - Some(Value::String(s)) => u32::from_str_radix(s, 16).ok(), - _ => None, - }, - _ => None, - } -} - impl CoreLoop { /// Merge the overlay for `params.txn_id` into `results` (base spatial - /// scan rows, already projected via `project_doc`). No-op when the + /// scan hits, already projected via `project_doc`). No-op when the /// transaction has no overlay entries for this collection. pub(in crate::data::executor) fn merge_overlay_into_spatial_scan( &self, params: SpatialOverlayMergeParams<'_>, - results: &mut Vec, + results: &mut Vec, ) -> crate::Result<()> { let SpatialOverlayMergeParams { txn_id, @@ -144,7 +132,9 @@ impl CoreLoop { }; // Surrogates already represented in the base result. - let mut seen: HashSet = results.iter().filter_map(row_surrogate).collect(); + let mut seen: HashSet = results.iter().filter_map(|hit| hit.surrogate).collect(); + let identity_column = + self.identity_column(coll_key.0.as_u64(), coll_key.1.as_u64(), &coll_key.2); // Base-minus-superseded: a tombstoned row is dropped; a staged put // replaces the row with the re-projected staged geometry and is @@ -158,11 +148,11 @@ impl CoreLoop { // finishes, aborting the merge before the overlay-addition pass runs. let mut first_err: Option = None; let base_visible = overlay.base_visible(coll_key); - results.retain_mut(|row| { + results.retain_mut(|hit| { if first_err.is_some() { return true; } - let Some(raw) = row_surrogate(row) else { + let Some(raw) = hit.surrogate else { return base_visible; }; match overlay.get(coll_key, raw) { @@ -186,8 +176,10 @@ impl CoreLoop { return true; } } - let doc_id = StorageKey::for_surrogate(Surrogate(raw)).to_string(); - *row = project_doc(&doc, &doc_id, projection); + // A staged row that lacks its identity column renders the + // client-visible identity there, never the storage key. + let identity = StorageKey::for_surrogate(Surrogate(raw)).to_identity(); + hit.row = project_doc(&doc, identity.as_str(), projection, &identity_column); true } None => base_visible, @@ -214,8 +206,11 @@ impl CoreLoop { if !row_matches(&doc)? { continue; } - let doc_id = StorageKey::for_surrogate(Surrogate(surrogate)).to_string(); - results.push(project_doc(&doc, &doc_id, projection)); + let identity = StorageKey::for_surrogate(Surrogate(surrogate)).to_identity(); + results.push(SpatialHit { + surrogate: Some(surrogate), + row: project_doc(&doc, identity.as_str(), projection, &identity_column), + }); seen.insert(surrogate); } Ok(()) diff --git a/nodedb/src/data/executor/handlers/transaction/overlay/vector_primary_merge.rs b/nodedb/src/data/executor/handlers/transaction/overlay/vector_primary_merge.rs index 44a55b4ea..4993b0472 100644 --- a/nodedb/src/data/executor/handlers/transaction/overlay/vector_primary_merge.rs +++ b/nodedb/src/data/executor/handlers/transaction/overlay/vector_primary_merge.rs @@ -45,17 +45,23 @@ pub(in crate::data::executor) enum SidecarRowShape { } /// The scan-row body of a staged vector-primary put in `shape`. +/// `identity_column` is the column a normalized row's identity renders under. fn staged_vector_scan_body( shape: SidecarRowShape, key: &StorageKey, body: &[u8], + identity_column: &str, ) -> crate::Result> { let sidecar = staged_vector_sidecar(body)?; match shape { SidecarRowShape::Stored => Ok(sidecar), SidecarRowShape::Normalized => { - let (_, normalized) = - sparse_row_to_doc(key, &sidecar, SparseBodyFormatRef::VectorSidecar); + let (_, normalized) = sparse_row_to_doc( + key, + &sidecar, + SparseBodyFormatRef::VectorSidecar, + identity_column, + ); Ok(normalized) } } @@ -78,8 +84,10 @@ impl CoreLoop { let Some(overlay) = self.txn_overlays.get(&txn_id) else { return Ok(()); }; + let identity_column = + self.identity_column(coll_key.0.as_u64(), coll_key.1.as_u64(), &coll_key.2); merge_staged_rows(overlay, coll_key, rows, matches, &|key, body| { - staged_vector_scan_body(shape, key, body) + staged_vector_scan_body(shape, key, body, &identity_column) }) } @@ -258,16 +266,16 @@ mod tests { .to_bytes() .expect("encode"); let key = StorageKey::for_surrogate(Surrogate::new(9)); - let body = - staged_vector_scan_body(SidecarRowShape::Normalized, &key, &staged).expect("decode"); + let body = staged_vector_scan_body(SidecarRowShape::Normalized, &key, &staged, "id") + .expect("decode"); let doc: serde_json::Value = nodedb_types::json_from_msgpack(&body).expect("json"); assert_eq!(doc.get("owner").and_then(|v| v.as_str()), Some("carol")); assert_eq!(doc.get("id").and_then(|v| v.as_str()), Some("r1")); assert_eq!( - staged_vector_scan_body(SidecarRowShape::Stored, &key, &staged).expect("decode"), + staged_vector_scan_body(SidecarRowShape::Stored, &key, &staged, "id").expect("decode"), stored_sidecar, "the stored shape is the sidecar byte-for-byte" ); - assert!(staged_vector_scan_body(SidecarRowShape::Normalized, &key, &[0xc1]).is_err()); + assert!(staged_vector_scan_body(SidecarRowShape::Normalized, &key, &[0xc1], "id").is_err()); } } diff --git a/nodedb/src/data/executor/handlers/transaction/stage_write/dispatch.rs b/nodedb/src/data/executor/handlers/transaction/stage_write/dispatch.rs index 394ad2158..45f7d45d8 100644 --- a/nodedb/src/data/executor/handlers/transaction/stage_write/dispatch.rs +++ b/nodedb/src/data/executor/handlers/transaction/stage_write/dispatch.rs @@ -457,6 +457,7 @@ impl CoreLoop { body, identity, schema.as_ref(), + &self.identity_column(database_id, tid, collection), tid, collection, ) diff --git a/nodedb/src/data/executor/handlers/transaction/stage_write/stage_returning.rs b/nodedb/src/data/executor/handlers/transaction/stage_write/stage_returning.rs index 00eb72f66..66133d96a 100644 --- a/nodedb/src/data/executor/handlers/transaction/stage_write/stage_returning.rs +++ b/nodedb/src/data/executor/handlers/transaction/stage_write/stage_returning.rs @@ -145,15 +145,21 @@ impl CoreLoop { rows: &[StoredRow<'_>], ) -> Response { let schema = self.resolve_strict_schema(of.database_id, of.tid, of.collection); - let reply = - build_stored_rows_payload(returning.spec, returning.rls_filters, schema.as_ref(), rows) - .and_then(|rows| { - zerompk::to_msgpack_vec(&StagedReturningReply { affected, rows }).map_err( - |error| crate::Error::Codec { - detail: format!("staged RETURNING reply: {error}"), - }, - ) - }); + let identity_column = self.identity_column(of.database_id, of.tid, of.collection); + let reply = build_stored_rows_payload( + returning.spec, + returning.rls_filters, + schema.as_ref(), + &identity_column, + rows, + ) + .and_then(|rows| { + zerompk::to_msgpack_vec(&StagedReturningReply { affected, rows }).map_err(|error| { + crate::Error::Codec { + detail: format!("staged RETURNING reply: {error}"), + } + }) + }); match reply { Ok(payload) => self.response_with_payload(of.task, payload), Err(e) => self.response_error( diff --git a/nodedb/src/data/executor/handlers/transaction/stage_write/stage_vector_targets.rs b/nodedb/src/data/executor/handlers/transaction/stage_write/stage_vector_targets.rs index fc47f1639..021213db9 100644 --- a/nodedb/src/data/executor/handlers/transaction/stage_write/stage_vector_targets.rs +++ b/nodedb/src/data/executor/handlers/transaction/stage_write/stage_vector_targets.rs @@ -117,6 +117,8 @@ impl CoreLoop { // Staged puts whose sidecar matches; a staged tombstone is a // row that is gone. if let Some(overlay) = overlay { + let identity_column = + self.identity_column(scope.database_id, scope.tid, scope.collection); for (surrogate, staged) in overlay.iter_for_collection(&scope.coll_key) { let Staged::Put(_) = staged else { continue; @@ -128,7 +130,12 @@ impl CoreLoop { continue; }; let key = StorageKey::for_surrogate(surrogate); - if vector_sidecar_matches(&key, &row.sidecar.bytes, &filters)? { + if vector_sidecar_matches( + &key, + &row.sidecar.bytes, + &filters, + &identity_column, + )? { out.push((surrogate, row)); } } diff --git a/nodedb/src/data/executor/handlers/update_from_join_collect.rs b/nodedb/src/data/executor/handlers/update_from_join_collect.rs index 2bd18e743..d018f8105 100644 --- a/nodedb/src/data/executor/handlers/update_from_join_collect.rs +++ b/nodedb/src/data/executor/handlers/update_from_join_collect.rs @@ -159,6 +159,12 @@ impl CoreLoop { } } let merged_ndb: nodedb_types::Value = merged.clone().into(); + let identity = super::identity_guard::IdentitySnapshot::capture( + strict_schema, + declared_primary_key, + updates, + &target_doc, + ); // Apply SET assignments evaluated against the merged document. if let Some(target_obj) = target_doc.as_object_mut() { @@ -191,6 +197,7 @@ impl CoreLoop { declared_primary_key, )?; } + identity.check_unchanged(target_collection, &target_doc)?; // Recompute generated columns if any dependency changed. A column // the engine cannot recompute fails the statement. @@ -208,6 +215,13 @@ impl CoreLoop { .map_err(crate::Error::DataPlane)?; } + // A declared numeric column holds the value its type stores, + // whether the assignment was a literal or computed. + super::super::strict_format::coerce_declared_doc( + &mut target_doc, + self.declared_columns(config_key), + )?; + // Re-encode the post-image (strict Binary Tuple or MessagePack). // An encode error carries its own typed cause, such as a field the // strict schema does not declare. @@ -275,13 +289,26 @@ impl CoreLoop { target_coll_key, bitemporal, } = args; + let identity_column = self.identity_column(database_id, tid, target_collection); let mut rows = if bitemporal { self.scan_current_versions(database_id, tid, target_collection, |key, body| { - matches_with_resolved_schema(strict_schema, target_filters, key, body) + matches_with_resolved_schema( + strict_schema, + target_filters, + key, + body, + &identity_column, + ) })? } else { self.scan_plain_rows(database_id, tid, target_collection, |key, body| { - matches_with_resolved_schema(strict_schema, target_filters, key, body) + matches_with_resolved_schema( + strict_schema, + target_filters, + key, + body, + &identity_column, + ) })? }; diff --git a/nodedb/src/data/executor/handlers/update_from_join_write.rs b/nodedb/src/data/executor/handlers/update_from_join_write.rs index 2411f74bd..d5b7fa7e7 100644 --- a/nodedb/src/data/executor/handlers/update_from_join_write.rs +++ b/nodedb/src/data/executor/handlers/update_from_join_write.rs @@ -254,7 +254,7 @@ impl CoreLoop { // The row's post-image, journalled after apply: this plan // carries no pre-dispatch record of it. Then one entry per // moved target row, naming the TARGET collection. - write_set.push(self.stored_row_image( + let image = self.stored_row_image( StoredRow { database_id, tid, @@ -265,7 +265,15 @@ impl CoreLoop { &updated_bytes, // A versioned row landed at the statement's system time. bitemporal_sys_from_ms, - )); + ); + match image { + Ok(image) => write_set.push(image), + // This row committed above, so it counts as landed. + Err(e) => { + let code = refusal_after_rows(affected + 1, e); + return Err(self.refusal_with_landed_rows(task, code, write_set)); + } + } write_set.extend(write_hook::target_write_set(&target_writes)); self.doc_cache.put( database_id, @@ -314,11 +322,15 @@ impl CoreLoop { } affected += 1; if want_returning { - // `row_identity` only stands in as `id` for a row that - // declares no primary key of its own — overwriting a - // declared key would return a value the client never wrote. + // `row_identity` fills the identity column only for a row + // that lacks it. A declared key keeps the value the + // client wrote. let mut row = nodedb_types::Value::from(doc); - returning_doc::attach_row_id(&mut row, &row_identity); + returning_doc::attach_row_id( + &mut row, + &row_identity, + &self.identity_column(database_id, tid, target_collection), + ); returned_docs.push(row); } } diff --git a/nodedb/src/data/executor/handlers/vector_direct_delete.rs b/nodedb/src/data/executor/handlers/vector_direct_delete.rs index 274a6d124..9aaaa1183 100644 --- a/nodedb/src/data/executor/handlers/vector_direct_delete.rs +++ b/nodedb/src/data/executor/handlers/vector_direct_delete.rs @@ -50,11 +50,18 @@ impl CoreLoop { debug!(core = self.core_id, %collection, %field, "vector direct delete"); let database_id = task.request.database_id.as_u64(); let index_key = CoreLoop::vector_index_key(database_id, tid, collection, field); + let identity_column = self.identity_column(database_id, tid, collection); // No index means no row was ever written: nothing to remove. if !self.vector_collections.contains_key(&index_key) { if let Some(spec) = returning { - return self.vector_stored_returning_response(task, spec, rls_filters, &[]); + return self.vector_stored_returning_response( + task, + spec, + rls_filters, + &identity_column, + &[], + ); } return self.response_affected(task, 0); } @@ -105,7 +112,13 @@ impl CoreLoop { if let Some(spec) = returning { let rows: Vec<(&StorageKey, &[u8])> = removed.iter().map(|(k, b)| (k, b.as_slice())).collect(); - return self.vector_stored_returning_response(task, spec, rls_filters, &rows); + return self.vector_stored_returning_response( + task, + spec, + rls_filters, + &identity_column, + &rows, + ); } self.response_affected(task, removed.len() as u64) } diff --git a/nodedb/src/data/executor/handlers/vector_direct_resolve/resolve_delete.rs b/nodedb/src/data/executor/handlers/vector_direct_resolve/resolve_delete.rs index 2d0ac82f9..feb937d83 100644 --- a/nodedb/src/data/executor/handlers/vector_direct_resolve/resolve_delete.rs +++ b/nodedb/src/data/executor/handlers/vector_direct_resolve/resolve_delete.rs @@ -73,7 +73,9 @@ impl CoreLoop { .zip(rows.iter()) .map(|(key, (_, bytes))| (key, bytes.as_slice())) .collect(); - vector_stored_rows_payload(spec, rls_filters, &stored).map_err(ErrorCode::from)? + let identity_column = self.identity_column(database_id, tid, collection); + vector_stored_rows_payload(spec, rls_filters, &identity_column, &stored) + .map_err(ErrorCode::from)? } None => response_codec::encode_affected(rows.len() as u64), }; diff --git a/nodedb/src/data/executor/handlers/vector_direct_resolve/resolve_update.rs b/nodedb/src/data/executor/handlers/vector_direct_resolve/resolve_update.rs index 65fe57de8..b7b2eaa62 100644 --- a/nodedb/src/data/executor/handlers/vector_direct_resolve/resolve_update.rs +++ b/nodedb/src/data/executor/handlers/vector_direct_resolve/resolve_update.rs @@ -91,7 +91,9 @@ impl CoreLoop { .zip(planned.iter()) .map(|(key, row)| (key, row.sidecar.as_slice())) .collect(); - vector_stored_rows_payload(spec, rls_filters, &stored).map_err(ErrorCode::from)? + let identity_column = self.identity_column(database_id, tid, collection); + vector_stored_rows_payload(spec, rls_filters, &identity_column, &stored) + .map_err(ErrorCode::from)? } None => response_codec::encode_affected(planned.len() as u64), }; diff --git a/nodedb/src/data/executor/handlers/vector_direct_resolve/resolve_upsert.rs b/nodedb/src/data/executor/handlers/vector_direct_resolve/resolve_upsert.rs index e63a3c16a..88c5221eb 100644 --- a/nodedb/src/data/executor/handlers/vector_direct_resolve/resolve_upsert.rs +++ b/nodedb/src/data/executor/handlers/vector_direct_resolve/resolve_upsert.rs @@ -122,8 +122,14 @@ impl CoreLoop { let response_payload = match returning { Some(spec) => { let key = StorageKey::for_surrogate(surrogate); - vector_stored_rows_payload(spec, rls_filters, &[(&key, sidecar.as_slice())]) - .map_err(ErrorCode::from)? + let identity_column = self.identity_column(database_id, tid, collection); + vector_stored_rows_payload( + spec, + rls_filters, + &identity_column, + &[(&key, sidecar.as_slice())], + ) + .map_err(ErrorCode::from)? } None if !on_conflict_updates.is_empty() => { let op = if existing { "update" } else { "insert" }; diff --git a/nodedb/src/data/executor/handlers/vector_direct_targets.rs b/nodedb/src/data/executor/handlers/vector_direct_targets.rs index cd5bb3d93..d9ee08be5 100644 --- a/nodedb/src/data/executor/handlers/vector_direct_targets.rs +++ b/nodedb/src/data/executor/handlers/vector_direct_targets.rs @@ -37,15 +37,22 @@ pub(in crate::data::executor) fn decode_vector_write_filters( /// through the vector sidecar format, the same converter `SELECT` uses. A /// predicate that divides by zero fails the statement, the same as it fails /// the equivalent `SELECT`. Empty `filters` match every row. +/// `identity_column` is the column the row's identity renders under. pub(in crate::data::executor) fn vector_sidecar_matches( key: &StorageKey, sidecar: &[u8], filters: &[ScanFilter], + identity_column: &str, ) -> Result { if filters.is_empty() { return Ok(true); } - let (_id, mp) = sparse_row_to_doc(key, sidecar, SparseBodyFormatRef::VectorSidecar); + let (_id, mp) = sparse_row_to_doc( + key, + sidecar, + SparseBodyFormatRef::VectorSidecar, + identity_column, + ); ScanFilter::all_match_binary(filters, &mp).map_err(ErrorCode::from) } @@ -113,6 +120,7 @@ impl CoreLoop { let range = table .range(prefix.as_str()..end.as_str()) .map_err(|e| storage_err("range", &e))?; + let identity_column = self.identity_column(database_id, tid, collection); let mut out = Vec::new(); for entry in range { @@ -128,7 +136,7 @@ impl CoreLoop { rest, )) })?; - if !vector_sidecar_matches(&key, value_guard.value(), filters)? { + if !vector_sidecar_matches(&key, value_guard.value(), filters, &identity_column)? { continue; } out.push(key.surrogate()); diff --git a/nodedb/src/data/executor/handlers/vector_direct_update.rs b/nodedb/src/data/executor/handlers/vector_direct_update.rs index 55954d124..c8e17d309 100644 --- a/nodedb/src/data/executor/handlers/vector_direct_update.rs +++ b/nodedb/src/data/executor/handlers/vector_direct_update.rs @@ -141,12 +141,19 @@ impl CoreLoop { "vector direct update" ); let database_id = task.request.database_id.as_u64(); + let identity_column = self.identity_column(database_id, tid, collection); // No index means no row was ever written: nothing to rewrite. let probe_key = CoreLoop::vector_index_key(database_id, tid, collection, field); if !self.vector_collections.contains_key(&probe_key) { if let Some(spec) = returning { - return self.vector_stored_returning_response(task, spec, rls_filters, &[]); + return self.vector_stored_returning_response( + task, + spec, + rls_filters, + &identity_column, + &[], + ); } return self.response_affected(task, 0); } @@ -215,7 +222,13 @@ impl CoreLoop { if let Some(spec) = returning { let rows: Vec<(&StorageKey, &[u8])> = written.iter().map(|(k, b)| (k, b.as_slice())).collect(); - return self.vector_stored_returning_response(task, spec, rls_filters, &rows); + return self.vector_stored_returning_response( + task, + spec, + rls_filters, + &identity_column, + &rows, + ); } self.response_affected(task, written.len() as u64) } diff --git a/nodedb/src/data/executor/handlers/vector_upsert.rs b/nodedb/src/data/executor/handlers/vector_upsert.rs index dd7d2bd04..222b4b72e 100644 --- a/nodedb/src/data/executor/handlers/vector_upsert.rs +++ b/nodedb/src/data/executor/handlers/vector_upsert.rs @@ -196,6 +196,7 @@ impl CoreLoop { "vector direct write" ); let database_id = task.request.database_id.as_u64(); + let identity_column = self.identity_column(database_id, tid, collection); let index_key = match self.vector_direct_index( task, @@ -251,7 +252,13 @@ impl CoreLoop { // `ON CONFLICT DO NOTHING`: nothing is written, so the count is // 0 and a `RETURNING` clause has no post-image to project. if let Some(spec) = returning { - return self.vector_stored_returning_response(task, spec, rls_filters, &[]); + return self.vector_stored_returning_response( + task, + spec, + rls_filters, + &identity_column, + &[], + ); } return self.response_affected(task, 0); } @@ -311,6 +318,7 @@ impl CoreLoop { task, spec, rls_filters, + &identity_column, &[(&storage_key, sidecar.as_slice())], ); } diff --git a/nodedb/src/data/executor/identity_column.rs b/nodedb/src/data/executor/identity_column.rs new file mode 100644 index 000000000..5e7005811 --- /dev/null +++ b/nodedb/src/data/executor/identity_column.rs @@ -0,0 +1,38 @@ +// SPDX-License-Identifier: BUSL-1.1 + +//! The column a collection's rows render their identity under. +//! +//! A row with a declared key carries its identity under that key. A row with +//! no declared key carries it under `id`. The Control Plane resolves the +//! declared key from the catalog and ships it on `DocumentOp::Register`, so +//! the Data Plane reads it from `doc_configs` and never from the catalog. +//! +//! Every reader that gives a sparse row its identity takes the column from +//! here: the scan row, the filter image, the `RETURNING` row, and the event +//! image. A row then never shows an `id` beside its declared key, and a +//! predicate naming the declared key finds it. + +use super::core_loop::CoreLoop; +use crate::types::{DatabaseId, TenantId}; + +impl CoreLoop { + /// The column `collection`'s rows render their identity under: the + /// registered declared key, else `id`. An unregistered collection has no + /// declared key. + pub(in crate::data::executor) fn identity_column( + &self, + database_id: u64, + tid: u64, + collection: &str, + ) -> String { + let key = ( + DatabaseId::new(database_id), + TenantId::new(tid), + collection.to_string(), + ); + self.doc_configs + .get(&key) + .and_then(|config| config.declared_key.clone()) + .unwrap_or_else(|| nodedb_types::DEFAULT_IDENTITY_COLUMN.to_string()) + } +} diff --git a/nodedb/src/data/executor/mod.rs b/nodedb/src/data/executor/mod.rs index 4cf98b258..648850612 100644 --- a/nodedb/src/data/executor/mod.rs +++ b/nodedb/src/data/executor/mod.rs @@ -17,6 +17,7 @@ pub mod enforcement; pub(crate) mod fts_text; mod graph_label_checkpoint; pub mod handlers; +mod identity_column; pub(crate) mod kv_checkpoint; pub(super) mod msgpack_utils; pub(crate) mod replay_abort; diff --git a/nodedb/src/data/executor/row_shape.rs b/nodedb/src/data/executor/row_shape.rs index 3b0e9af3e..885cff82f 100644 --- a/nodedb/src/data/executor/row_shape.rs +++ b/nodedb/src/data/executor/row_shape.rs @@ -60,118 +60,81 @@ pub(in crate::data::executor) fn sparse_body_to_msgpack<'a>( } } -/// Convert a single sparse/document row to a `(id, msgpack)` document. +/// Convert a single sparse/document row to a `(storage key, msgpack)` pair. /// -/// Normalizes the body per [`sparse_body_to_msgpack`], then injects the `id` -/// field. Injection is a no-op when the body already carries an `id` — a -/// vector-primary sidecar stores the user's declared primary key, and its -/// sparse key is the internal surrogate-hex, which must not displace it. -/// -/// `key` is the row's storage key. The client-visible identity is the row's -/// `_rowid` when its body carries one, else its surrogate's decimal string, -/// per [`StorageKey::to_identity`]. A first write sets `_rowid` to the -/// surrogate, and a copy under a new surrogate keeps it. Shared by the -/// materializing scan and the streaming scan so both paths produce -/// byte-identical output. +/// Normalizes the body per [`sparse_body_to_msgpack`], then gives it its +/// identity per [`inject_row_identity`] under `identity_column`, the column +/// `CoreLoop::identity_column` resolves for the collection. Shared by the +/// materializing scan, the streaming scan, and every per-engine read, so all +/// of them produce byte-identical rows. pub(in crate::data::executor) fn sparse_row_to_doc( key: &nodedb_types::StorageKey, raw: &[u8], format: SparseBodyFormatRef<'_>, + identity_column: &str, ) -> (String, Vec) { let mp = sparse_body_to_msgpack(raw, format); - let identity = msgpack_scan::extract_field(&mp, 0, nodedb_types::ROWID_COLUMN) - .and_then(|(start, _)| msgpack_scan::read_i64(&mp, start)) + ( + key.to_string(), + inject_row_identity(&mp, key, identity_column).into_owned(), + ) +} + +/// Give a msgpack row body its identity under `identity_column`. +/// +/// `identity_column` is the collection's declared key, else `id`. A body that +/// holds that column keeps it unchanged: a declared key is authoritative, and +/// a vector-primary sidecar stores the user's key, so no `id` appears beside +/// it. A body that lacks it gains the identity there: the body's `_rowid` +/// when it carries one, else the storage key's identity per +/// [`StorageKey::to_identity`](nodedb_types::StorageKey::to_identity). A first +/// write sets `_rowid` to the surrogate, and a copy under a new surrogate +/// keeps it. The Control Plane's envelope shaping applies the same rule. +/// The result borrows `body` when it already holds the column. +pub(in crate::data::executor) fn inject_row_identity<'a>( + body: &'a [u8], + key: &nodedb_types::StorageKey, + identity_column: &str, +) -> std::borrow::Cow<'a, [u8]> { + if msgpack_scan::extract_field(body, 0, identity_column).is_some() { + return std::borrow::Cow::Borrowed(body); + } + let identity = msgpack_scan::extract_field(body, 0, nodedb_types::ROWID_COLUMN) + .and_then(|(start, _)| msgpack_scan::read_integer(body, start)) .map(|rowid| rowid.to_string()) .unwrap_or_else(|| key.to_identity().into_string()); - let mp = msgpack_scan::inject_str_field(&mp, "id", &identity); - (key.to_string(), mp) + std::borrow::Cow::Owned(msgpack_scan::inject_str_field( + body, + identity_column, + &identity, + )) } /// Convert a single row from a `DecodedColumn` to a `nodedb_types::value::Value`. /// -/// `declared` is the column's declared type from the collection schema. It -/// decides how an eight-byte integer cell is typed, because the segment -/// reader infers a column's physical kind from its codec and decodes every -/// time column as `DecodedColumn::Int64`: a `Timestamp` column yields -/// `Value::NaiveDateTime`, a `Timestamptz` column `Value::DateTime`, both from -/// the epoch microseconds the segment stores. Every other declared type -/// yields the integer stored: a `SystemTimestamp` column is an engine-assigned -/// system-time count, and the bitemporal `_ts_system`, `_ts_valid_from`, -/// `_ts_valid_until` columns are declared `Int64` and hold epoch milliseconds -/// with `i64::MIN` / `i64::MAX` as the unbounded sentinels, which no instant -/// can carry. The live memtable applies the same rule, so a row reads -/// identically before and after a flush. +/// `declared` is the column's declared type from the collection schema. A +/// segment block records only its physical layout, so the declared type +/// decides the variant, by the live memtable's read rule: a row reads +/// identically before and after a flush. A `UUID` cell is UUID text, a `ULID` +/// cell ULID text, a `VECTOR` cell an array of floats, a `GEOMETRY` cell its +/// stored text, and a `JSON` cell its decoded value. A `Timestamp` column +/// yields `Value::NaiveDateTime` and a `Timestamptz` column `Value::DateTime`, +/// both from the epoch microseconds the segment stores. Every other declared +/// type backed by integer storage yields the integer stored: a +/// `SystemTimestamp` column is an engine-assigned system-time count, and the +/// bitemporal `_ts_system`, `_ts_valid_from`, `_ts_valid_until` columns are +/// declared `Int64` and hold epoch milliseconds with `i64::MIN` / `i64::MAX` +/// as the unbounded sentinels, which no instant can carry. /// -/// Returns `Value::Null` if the row index is out of range or the validity bit is false. +/// Returns `Value::Null` if the row index is out of range or the validity bit +/// is false. Returns `Err` when the cell's bytes do not hold a value of the +/// declared type: the segment is corrupt. pub(in crate::data::executor) fn decoded_col_to_value( col: &nodedb_columnar::reader::DecodedColumn, row_idx: usize, declared: &nodedb_types::columnar::ColumnType, -) -> nodedb_types::value::Value { - use nodedb_columnar::reader::DecodedColumn; - use nodedb_types::value::Value; - - match col { - DecodedColumn::Int64 { values, valid } | DecodedColumn::Timestamp { values, valid } => { - if row_idx < valid.len() && valid[row_idx] { - declared.time_cell(values[row_idx]) - } else { - Value::Null - } - } - DecodedColumn::Float64 { values, valid } => { - if row_idx < valid.len() && valid[row_idx] { - Value::Float(values[row_idx]) - } else { - Value::Null - } - } - DecodedColumn::Bool { values, valid } => { - if row_idx < valid.len() && valid[row_idx] { - Value::Bool(values[row_idx]) - } else { - Value::Null - } - } - DecodedColumn::Binary { - data, - offsets, - valid, - } => { - if row_idx < valid.len() && valid[row_idx] && row_idx + 1 < offsets.len() { - let start = offsets[row_idx] as usize; - let end = offsets[row_idx + 1] as usize; - if start <= end && end <= data.len() { - let bytes = &data[start..end]; - // Best-effort UTF-8 interpretation; fall back to bytes. - match std::str::from_utf8(bytes) { - Ok(s) => Value::String(s.to_string()), - Err(_) => Value::Bytes(bytes.to_vec()), - } - } else { - Value::Null - } - } else { - Value::Null - } - } - DecodedColumn::DictEncoded { - ids, - dictionary, - valid, - } => { - if row_idx < valid.len() && valid[row_idx] { - let id = ids[row_idx] as usize; - if id < dictionary.len() { - Value::String(dictionary[id].clone()) - } else { - Value::Null - } - } else { - Value::Null - } - } - } +) -> crate::Result { + nodedb_columnar::reader::decoded_cell_value(col, row_idx, declared).map_err(crate::Error::from) } #[cfg(test)] @@ -181,10 +144,14 @@ mod tests { use nodedb_types::value::Value; use nodedb_types::{InstantKind, NdbDateTime}; - use super::{decoded_col_to_value, kv_row_to_doc, msgpack_scan}; + use super::{SparseBodyFormatRef, kv_row_to_doc, msgpack_scan, sparse_row_to_doc}; const MICROS: i64 = 1_583_402_400_000_000; + fn decoded_col_to_value(col: &DecodedColumn, row: usize, ty: &ColumnType) -> Value { + super::decoded_col_to_value(col, row, ty).expect("decodable cell") + } + /// A time column as the segment reader decodes it: the reader infers the /// physical kind from the codec, so a time column arrives as `Int64`. fn time_column() -> DecodedColumn { @@ -260,6 +227,28 @@ mod tests { ); } + /// A 16-byte identifier cell reads as the text its declared type names, + /// as on the live memtable, never as raw bytes. + #[test] + fn an_identifier_cell_reads_back_as_its_declared_text() { + let uuid = uuid::Uuid::from_u128(0x67e5_5044_10b1_426f_9247_bb68_0e5f_e0c8); + let col = DecodedColumn::Binary { + data: uuid.as_bytes().to_vec(), + offsets: vec![0, 16], + valid: vec![true], + }; + assert_eq!( + decoded_col_to_value(&col, 0, &ColumnType::Uuid), + Value::Uuid(uuid.to_string()) + ); + // A ULID is 26 Crockford base32 characters. + assert!(matches!( + decoded_col_to_value(&col, 0, &ColumnType::Ulid), + Value::Ulid(ref text) if text.len() == 26 + )); + assert!(super::decoded_col_to_value(&col, 0, &ColumnType::Vector(2)).is_err()); + } + /// A raw (non-msgpack) KV value must be wrapped as a msgpack STRING, not /// appended verbatim. /// @@ -284,6 +273,64 @@ mod tests { assert_eq!(doc.get("key").and_then(|v| v.as_str()), Some("k1")); } + /// A schemaless body of `fields`, as msgpack. + fn body(fields: &[(&str, &str)]) -> Vec { + let mut mp = Vec::new(); + msgpack_scan::write_map_header(&mut mp, fields.len()); + for (name, value) in fields { + msgpack_scan::write_str(&mut mp, name); + msgpack_scan::write_str(&mut mp, value); + } + mp + } + + /// Shape one schemaless row stored under surrogate 7. + fn shaped( + fields: &[(&str, &str)], + identity_column: &str, + ) -> std::collections::HashMap { + let key = nodedb_types::StorageKey::for_surrogate(nodedb_types::Surrogate::new(7)); + let (_, mp) = sparse_row_to_doc( + &key, + &body(fields), + SparseBodyFormatRef::Document, + identity_column, + ); + let Ok(Value::Object(row)) = nodedb_types::value_from_msgpack(&mp) else { + panic!("a shaped row decodes as a map"); + }; + row + } + + /// A row holding its declared key keeps it and gains no `id`. + #[test] + fn a_row_holding_its_declared_key_gains_no_id() { + let row = shaped(&[("sku", "p1"), ("name", "pen")], "sku"); + assert_eq!(row.get("sku"), Some(&Value::String("p1".into()))); + assert!(!row.contains_key("id"), "{row:?}"); + assert_eq!(row.len(), 2, "{row:?}"); + } + + /// A row lacking its identity column gains the identity there, and only + /// there. + #[test] + fn a_row_lacking_its_identity_column_gains_it_there() { + let row = shaped(&[("name", "pen")], "sku"); + assert_eq!(row.get("sku"), Some(&Value::String("7".into()))); + assert!(!row.contains_key("id"), "{row:?}"); + + let row = shaped(&[("name", "pen")], "id"); + assert_eq!(row.get("id"), Some(&Value::String("7".into()))); + } + + /// A stored `id` stays authoritative under a collection with no declared + /// key. + #[test] + fn a_stored_id_stays_authoritative() { + let row = shaped(&[("id", "dup"), ("n", "1")], "id"); + assert_eq!(row.get("id"), Some(&Value::String("dup".into()))); + } + /// A msgpack-map value keeps its fields and gains `key`. #[test] fn a_msgpack_map_kv_value_keeps_its_fields() { diff --git a/nodedb/src/data/executor/scan_normalize.rs b/nodedb/src/data/executor/scan_normalize.rs index 61f397599..7f482400c 100644 --- a/nodedb/src/data/executor/scan_normalize.rs +++ b/nodedb/src/data/executor/scan_normalize.rs @@ -189,14 +189,14 @@ impl CoreLoop { // `scan_documents_for_each`. The body encoding is resolved once up // front through the same helper `scan_sparse` uses, so the streaming // and materializing scans cannot disagree about a row's shape. - let format = self.sparse_body_format( - crate::types::DatabaseId::new(did), - crate::types::TenantId::new(tid), - collection, - ); + let database_id = crate::types::DatabaseId::new(did); + let tenant_id = crate::types::TenantId::new(tid); + let format = self.sparse_body_format(database_id, tenant_id, collection); + let identity_column = self.identity_column(did, tid, collection); self.sparse .scan_documents_for_each(did, tid, collection, usize::MAX, |id, raw| { - let (id_s, mp) = sparse_row_to_doc(id, raw, format.as_format_ref()); + let (id_s, mp) = + sparse_row_to_doc(id, raw, format.as_format_ref(), &identity_column); f(&id_s, &mp) })?; Ok(()) @@ -299,67 +299,26 @@ impl CoreLoop { if results.len() >= limit { break; } - let seg_id = format!("{}", seg_idx as u64 + 1); - let reader = if let Some(ref reg) = self.quarantine_registry { - match crate::storage::quarantine::engines::open_segment_with_quarantine( - reg, seg_bytes, collection, &seg_id, - ) { - Ok(r) => r, - Err(e) => { - tracing::warn!(error = %e, segment_id = %seg_id, collection, "failed to open flushed columnar segment for scan"); - continue; - } - } - } else { - match nodedb_columnar::SegmentReader::open(seg_bytes) { - Ok(r) => r, - Err(e) => { - tracing::warn!(error = %e, "failed to open flushed columnar segment for scan"); - continue; - } - } - }; - let seg_row_count = reader.row_count() as usize; - let remaining = limit - results.len(); - let take = seg_row_count.min(remaining); - - // Decode all columns for this segment. - let col_count = schema.columns.len(); - let mut decoded_cols = Vec::with_capacity(col_count); - let mut decode_ok = true; - for col_idx in 0..col_count { - match reader.read_column(col_idx) { - Ok(dc) => decoded_cols.push(dc), - Err(e) => { - tracing::warn!(error = %e, col_idx, "failed to decode columnar segment column"); - decode_ok = false; - break; - } + let seg_id = seg_idx as u64 + 1; + // A segment that does not read refuses the scan. Skipping it + // would answer without the segment's rows. + let segment = self.decode_flushed_segment( + collection, + seg_id, + seg_bytes, + schema.columns.len(), + "columnar_normalized_scan", + )?; + let delete_bm = engine.delete_bitmap(seg_id); + for row_idx in 0..segment.row_count() { + if results.len() >= limit { + break; } - } - if !decode_ok { - continue; - } - - for row_idx in 0..take { - let mut map = std::collections::HashMap::new(); - let mut id = String::new(); - for (col_idx, col_def) in schema.columns.iter().enumerate() { - let val = decoded_col_to_value( - &decoded_cols[col_idx], - row_idx, - &col_def.column_type, - ); - if col_def.name == "id" - && let nodedb_types::value::Value::String(s) = &val - { - id.clone_from(s); - } - map.insert(col_def.name.clone(), val); + // A tombstoned row is not a row any more. + if delete_bm.is_some_and(|bm| bm.is_deleted(row_idx as u32)) { + continue; } - let ndb_val = nodedb_types::value::Value::Object(map); - let mp = nodedb_types::value_to_msgpack(&ndb_val).unwrap_or_default(); - results.push((id, mp)); + results.push(columnar_row_to_doc(schema, segment.row(schema, row_idx)?)?); } } } @@ -367,23 +326,8 @@ impl CoreLoop { // 2. Read from the live memtable (most-recent rows not yet flushed). if results.len() < limit { let remaining = limit - results.len(); - let rows: Vec<_> = engine.scan_memtable_rows().take(remaining).collect(); - for row in rows { - let mut map = std::collections::HashMap::new(); - let mut id = String::new(); - for (i, col_def) in schema.columns.iter().enumerate() { - if i < row.len() { - if col_def.name == "id" - && let nodedb_types::value::Value::String(s) = &row[i] - { - id.clone_from(s); - } - map.insert(col_def.name.clone(), row[i].clone()); - } - } - let ndb_val = nodedb_types::value::Value::Object(map); - let mp = nodedb_types::value_to_msgpack(&ndb_val).unwrap_or_default(); - results.push((id, mp)); + for row in engine.scan_memtable_rows().take(remaining) { + results.push(columnar_row_to_doc(schema, row?)?); } } @@ -400,20 +344,50 @@ impl CoreLoop { limit: usize, ) -> crate::Result)>> { let docs = self.sparse.scan_documents(did, tid, collection, limit)?; - let format = self.sparse_body_format( - crate::types::DatabaseId::new(did), - crate::types::TenantId::new(tid), - collection, - ); + let database_id = crate::types::DatabaseId::new(did); + let tenant_id = crate::types::TenantId::new(tid); + let format = self.sparse_body_format(database_id, tenant_id, collection); + let identity_column = self.identity_column(did, tid, collection); let mut normalized = Vec::with_capacity(docs.len()); for (id, raw) in docs { - normalized.push(sparse_row_to_doc(&id, &raw, format.as_format_ref())); + normalized.push(sparse_row_to_doc( + &id, + &raw, + format.as_format_ref(), + &identity_column, + )); } Ok(normalized) } } +/// A positional columnar row as `(id, msgpack object)`: the `id` column's +/// text, or empty when the row has none, and the row keyed by column name. +fn columnar_row_to_doc( + schema: &nodedb_types::columnar::ColumnarSchema, + row: Vec, +) -> crate::Result<(String, Vec)> { + let mut id = String::new(); + let mut map = std::collections::HashMap::with_capacity(schema.columns.len()); + for (col_def, val) in schema.columns.iter().zip(row) { + if col_def.name == "id" + && let nodedb_types::value::Value::String(s) = &val + { + id.clone_from(s); + } + map.insert(col_def.name.clone(), val); + } + let mp = + nodedb_types::value_to_msgpack(&nodedb_types::value::Value::Object(map)).map_err(|e| { + crate::Error::Serialization { + format: "msgpack".into(), + detail: format!("encode normalized columnar row: {e}"), + } + })?; + Ok((id, mp)) +} + // The per-row shape converters live in `row_shape`, beside each other and // away from the scan orchestration above. Re-exported here because every // caller in this directory reaches them through `scan_normalize::` and the diff --git a/nodedb/src/data/executor/scan_versioned.rs b/nodedb/src/data/executor/scan_versioned.rs index 7579e9635..72312ff60 100644 --- a/nodedb/src/data/executor/scan_versioned.rs +++ b/nodedb/src/data/executor/scan_versioned.rs @@ -17,7 +17,8 @@ impl CoreLoop { /// live version per `doc_id` and normalizes each body to standard msgpack /// via [`sparse_row_to_doc`], producing rows byte-identical in shape to /// [`CoreLoop::scan_collection`] (schemaless normalized from possibly-legacy - /// JSON, strict decoded from Binary Tuple, `id` injected). + /// JSON, strict decoded from Binary Tuple, identity injected under the + /// collection's identity column). pub(in crate::data::executor) fn scan_collection_versioned_current( &self, did: u64, @@ -39,15 +40,19 @@ impl CoreLoop { // own bound (an explicit `limit`), so no deadline cuts it short. &crate::engine::sparse::scan_stop::never_stop, )?; - let format = self.sparse_body_format( - crate::types::DatabaseId::new(did), - crate::types::TenantId::new(tid), - collection, - ); + let database_id = crate::types::DatabaseId::new(did); + let tenant_id = crate::types::TenantId::new(tid); + let format = self.sparse_body_format(database_id, tenant_id, collection); + let identity_column = self.identity_column(did, tid, collection); let mut normalized = Vec::with_capacity(docs.len()); for (key, raw) in docs { - normalized.push(sparse_row_to_doc(&key, &raw, format.as_format_ref())); + normalized.push(sparse_row_to_doc( + &key, + &raw, + format.as_format_ref(), + &identity_column, + )); } Ok(normalized) } diff --git a/nodedb/tests/native/cases/native_document_identity_update.rs b/nodedb/tests/native/cases/native_document_identity_update.rs new file mode 100644 index 000000000..6b93d4ea9 --- /dev/null +++ b/nodedb/tests/native/cases/native_document_identity_update.rs @@ -0,0 +1,104 @@ +// SPDX-License-Identifier: BUSL-1.1 + +//! A native document update keeps the row's identity column. +//! +//! The row stays stored under its document id, so an `id` assignment naming +//! another value splits SQL reads (which read the column) from key reads +//! (which read the key). The update is refused; assigning the current id is +//! accepted. + +use nodedb_test_support::native_harness::{NativeTestServer, do_handshake, send_request, send_sql}; +use nodedb_types::Value; +use nodedb_types::error::sqlstate; +use nodedb_types::protocol::HelloFrame; +use nodedb_types::protocol::opcodes::{OpCode, ResponseStatus}; +use nodedb_types::protocol::text_fields::TextFields; +use tokio::net::TcpStream; + +fn packed(value: &str) -> Vec { + nodedb_types::value_to_msgpack(&Value::String(value.into())).expect("encode") +} + +async fn seeded_session(server: &NativeTestServer, collection: &str) -> TcpStream { + let (mut stream, _ack) = do_handshake(server.addr, &HelloFrame::current()) + .await + .expect("native handshake"); + let create = send_sql(&mut stream, 1, &format!("CREATE COLLECTION {collection}")).await; + assert_ne!(create.status, ResponseStatus::Error, "create {collection}"); + let insert = send_sql( + &mut stream, + 2, + &format!("INSERT INTO {collection} (id, v) VALUES ('k1', 1)"), + ) + .await; + assert_ne!(insert.status, ResponseStatus::Error, "seed {collection}"); + stream +} + +async fn update_id(stream: &mut TcpStream, seq: u64, collection: &str, id: &str) -> ResponseStatus { + let resp = send_request( + stream, + seq, + OpCode::DocumentUpdate, + TextFields { + collection: Some(collection.to_string()), + document_id: Some("k1".to_string()), + updates: Some(vec![("id".to_string(), packed(id))]), + ..Default::default() + }, + ) + .await; + if resp.status == ResponseStatus::Error { + let err = resp.error.expect("error payload expected"); + assert_eq!( + err.code, + sqlstate::INTEGRITY_CONSTRAINT_VIOLATION, + "an identity change is an integrity violation, got {}: {}", + err.code, + err.message + ); + } + resp.status +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn native_update_refuses_an_id_naming_another_document() { + let server = NativeTestServer::start().await; + let mut stream = seeded_session(&server, "native_id_move").await; + + assert_eq!( + update_id(&mut stream, 3, "native_id_move", "k2").await, + ResponseStatus::Error, + "an id differing from the document id must be refused" + ); + + let read = send_sql(&mut stream, 4, "SELECT id FROM native_id_move").await; + assert_eq!( + read.rows, + Some(vec![vec![Value::String("k1".into())]]), + "the stored id must still name the key: {read:?}" + ); + let by_key = send_sql( + &mut stream, + 5, + "SELECT v FROM native_id_move WHERE id = 'k1'", + ) + .await; + assert_eq!( + by_key.rows, + Some(vec![vec![Value::Integer(1)]]), + "the row stays addressable by its key: {by_key:?}" + ); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn native_update_accepts_the_current_id() { + let server = NativeTestServer::start().await; + let mut stream = seeded_session(&server, "native_id_keep").await; + + assert_ne!( + update_id(&mut stream, 3, "native_id_keep", "k1").await, + ResponseStatus::Error, + "assigning the document's own id keeps its identity" + ); +} diff --git a/nodedb/tests/wire/cases/declared_key_scan_rows.rs b/nodedb/tests/wire/cases/declared_key_scan_rows.rs new file mode 100644 index 000000000..38890fb02 --- /dev/null +++ b/nodedb/tests/wire/cases/declared_key_scan_rows.rs @@ -0,0 +1,96 @@ +// SPDX-License-Identifier: BUSL-1.1 + +//! A row of a schemaless collection with a declared key names its identity by +//! that key. No read path adds a synthesized `id` column beside it: not a +//! full-text search, and not a join. + +use std::collections::HashMap; + +use crate::harness::TestServer; + +/// Create `items` keyed by `sku` and `orders` keyed by `oid`, both +/// schemaless, with one item and one order that names it. +async fn seed(srv: &TestServer) { + srv.exec("CREATE COLLECTION items (sku STRING NOT NULL PRIMARY KEY, name STRING)") + .await + .expect("create items"); + srv.exec("CREATE COLLECTION orders (oid STRING NOT NULL PRIMARY KEY, sku STRING, qty INT)") + .await + .expect("create orders"); + srv.exec("INSERT INTO items (sku, name) VALUES ('p1', 'blue pen')") + .await + .expect("insert item"); + srv.exec("INSERT INTO items (sku, name) VALUES ('p2', 'red mug')") + .await + .expect("insert item"); + srv.exec("INSERT INTO orders (oid, sku, qty) VALUES ('o1', 'p1', 3)") + .await + .expect("insert order"); +} + +/// Whether `column` names `field`, bare or qualified by a collection. +fn names(column: &str, field: &str) -> bool { + column == field || column.ends_with(&format!(".{field}")) +} + +/// Assert no column of `row` is an `id` column. +fn assert_no_id(row: &HashMap) { + assert!( + !row.keys().any(|column| names(column, "id")), + "a declared-key row shows no `id` column: {row:?}" + ); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn a_text_search_row_names_its_identity_by_the_declared_key() { + let srv = TestServer::start().await; + seed(&srv).await; + + let rows = srv + .query_named_rows("SELECT * FROM items WHERE text_match(name, 'pen')") + .await + .expect("text search"); + assert_eq!(rows.len(), 1, "one item matches 'pen': {rows:?}"); + let row = &rows[0]; + assert_no_id(row); + assert_eq!( + row.get("sku").map(String::as_str), + Some("p1"), + "the key column holds the key: {row:?}" + ); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn a_join_row_names_its_identity_by_the_declared_key() { + let srv = TestServer::start().await; + seed(&srv).await; + + let rows = srv + .query_named_rows("SELECT * FROM items i JOIN orders o ON o.sku = i.sku") + .await + .expect("join"); + assert_eq!(rows.len(), 1, "one order joins one item: {rows:?}"); + let row = &rows[0]; + assert_no_id(row); + assert!( + row.iter() + .any(|(column, value)| names(column, "sku") && value == "p1"), + "a key column holds the item key: {row:?}" + ); + assert!( + row.iter() + .any(|(column, value)| names(column, "oid") && value == "o1"), + "the order key column holds the order key: {row:?}" + ); + + // A join that names the declared keys finds them. + let rows = srv + .query_rows("SELECT i.sku, o.oid FROM items i JOIN orders o ON o.sku = i.sku") + .await + .expect("join on keys"); + assert_eq!( + rows, + vec![vec!["p1".to_string(), "o1".to_string()]], + "the projected keys hold the keys" + ); +} diff --git a/nodedb/tests/wire/cases/sql_on_conflict_primary_key_identity.rs b/nodedb/tests/wire/cases/sql_on_conflict_primary_key_identity.rs new file mode 100644 index 000000000..c9da56109 --- /dev/null +++ b/nodedb/tests/wire/cases/sql_on_conflict_primary_key_identity.rs @@ -0,0 +1,219 @@ +// SPDX-License-Identifier: BUSL-1.1 + +//! `INSERT ... ON CONFLICT DO UPDATE` keeps the conflicting row's key on the +//! key-value, vector-primary and columnar engines. +//! +//! The conflict branch merges its assignments into the row stored under the +//! conflicting key, and that row keeps its key and surrogate. An assignment +//! that moves the key column is refused with `23000`. Assigning the key the +//! row holds, or `EXCLUDED.`, is accepted. + +use crate::harness::TestServer; + +/// SQLSTATE of a statement that must fail, or `None` when it succeeded. +async fn sqlstate_of(server: &TestServer, sql: &str) -> Option { + match server.client.simple_query(sql).await { + Ok(_) => None, + Err(e) => Some( + e.as_db_error() + .unwrap_or_else(|| panic!("expected a DbError from {sql}, got: {e}")) + .code() + .code() + .to_string(), + ), + } +} + +async fn assert_key_change_refused(server: &TestServer, sql: &str) { + let state = sqlstate_of(server, sql) + .await + .unwrap_or_else(|| panic!("a primary-key change must be refused, but ran: {sql}")); + assert_eq!( + state, "23000", + "a primary-key change is an integrity violation, got SQLSTATE {state} for: {sql}" + ); +} + +async fn query(server: &TestServer, sql: &str) -> Vec { + server + .query_text(sql) + .await + .unwrap_or_else(|e| panic!("{sql}: {e}")) +} + +/// KV shapes: the built-in `key` column, and a named key column the body +/// also carries. +const KV_SHAPES: [(&str, &str); 2] = [ + ("key", "(key TEXT PRIMARY KEY, v TEXT)"), + ("k", "(k TEXT PRIMARY KEY, v TEXT)"), +]; + +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn kv_on_conflict_update_refuses_a_new_key() { + let server = TestServer::start().await; + for (i, (key, shape)) in KV_SHAPES.iter().enumerate() { + let name = format!("pk_conflict_kv_{i}"); + server + .exec(&format!( + "CREATE COLLECTION {name} {shape} WITH (engine='kv')" + )) + .await + .expect("create kv collection"); + server + .exec(&format!( + "INSERT INTO {name} ({key}, v) VALUES ('k1', 'one'), ('k2', 'two')" + )) + .await + .expect("seed rows"); + + for assignment in [format!("{key} = 'moved'"), format!("{key} = EXCLUDED.v")] { + assert_key_change_refused( + &server, + &format!( + "INSERT INTO {name} ({key}, v) VALUES ('k1', 'again') \ + ON CONFLICT ({key}) DO UPDATE SET {assignment}" + ), + ) + .await; + } + assert_eq!( + query(&server, &format!("SELECT v FROM {name} ORDER BY v")).await, + vec!["one".to_string(), "two".to_string()], + "a refused key change writes nothing" + ); + + server + .exec(&format!( + "INSERT INTO {name} ({key}, v) VALUES ('k1', 'again') \ + ON CONFLICT ({key}) DO UPDATE SET {key} = EXCLUDED.{key}, v = EXCLUDED.v" + )) + .await + .expect("assigning the conflicting key keeps the identity"); + server + .exec(&format!( + "INSERT INTO {name} ({key}, v) VALUES ('k2', 'ignored') \ + ON CONFLICT ({key}) DO UPDATE SET {key} = 'k2', v = 'same'" + )) + .await + .expect("assigning the current key keeps the identity"); + assert_eq!( + query(&server, &format!("SELECT v FROM {name} WHERE {key} = 'k1'")).await, + vec!["again".to_string()] + ); + assert_eq!( + query(&server, &format!("SELECT v FROM {name} WHERE {key} = 'k2'")).await, + vec!["same".to_string()] + ); + } +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn vector_primary_on_conflict_update_refuses_a_new_key() { + let server = TestServer::start().await; + server + .exec( + "CREATE COLLECTION pk_conflict_vp (id STRING PRIMARY KEY, vec VECTOR(3), \ + owner STRING) WITH (engine='vector', primary='vector', vector_field='vec', \ + dim=3, payload_indexes=['owner'])", + ) + .await + .expect("create vector-primary collection"); + server + .exec( + "INSERT INTO pk_conflict_vp (id, vec, owner) \ + VALUES ('r1', ARRAY[1.0, 0.0, 0.0], 'alice')", + ) + .await + .expect("seed row"); + + for assignment in ["id = 'moved'", "id = EXCLUDED.owner"] { + assert_key_change_refused( + &server, + &format!( + "INSERT INTO pk_conflict_vp (id, vec, owner) \ + VALUES ('r1', ARRAY[0.0, 1.0, 0.0], 'bob') \ + ON CONFLICT (id) DO UPDATE SET {assignment}" + ), + ) + .await; + } + assert_eq!( + query(&server, "SELECT owner FROM pk_conflict_vp WHERE id = 'r1'").await, + vec!["alice".to_string()], + "a refused key change writes nothing" + ); + assert!( + query( + &server, + "SELECT owner FROM pk_conflict_vp WHERE id = 'moved'" + ) + .await + .is_empty() + ); + + server + .exec( + "INSERT INTO pk_conflict_vp (id, vec, owner) \ + VALUES ('r1', ARRAY[0.0, 1.0, 0.0], 'bob') \ + ON CONFLICT (id) DO UPDATE SET id = EXCLUDED.id, owner = EXCLUDED.owner", + ) + .await + .expect("assigning the conflicting key keeps the identity"); + server + .exec( + "INSERT INTO pk_conflict_vp (id, vec, owner) \ + VALUES ('r1', ARRAY[0.0, 1.0, 0.0], 'ignored') \ + ON CONFLICT (id) DO UPDATE SET id = 'r1'", + ) + .await + .expect("assigning the current key keeps the identity"); + assert_eq!( + query(&server, "SELECT owner FROM pk_conflict_vp WHERE id = 'r1'").await, + vec!["bob".to_string()] + ); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn columnar_on_conflict_update_refuses_a_new_key() { + let server = TestServer::start().await; + server + .exec("CREATE COLLECTION pk_conflict_col (id TEXT PRIMARY KEY, v TEXT) WITH (engine='columnar')") + .await + .expect("create columnar collection"); + server + .exec("INSERT INTO pk_conflict_col (id, v) VALUES ('k1', 'one'), ('k2', 'two')") + .await + .expect("seed rows"); + + for assignment in ["id = 'moved'", "id = EXCLUDED.v"] { + assert_key_change_refused( + &server, + &format!( + "INSERT INTO pk_conflict_col (id, v) VALUES ('k1', 'again') \ + ON CONFLICT (id) DO UPDATE SET {assignment}" + ), + ) + .await; + } + assert_eq!( + query(&server, "SELECT id FROM pk_conflict_col ORDER BY id").await, + vec!["k1".to_string(), "k2".to_string()], + "a refused key change leaves one row per key" + ); + + server + .exec( + "INSERT INTO pk_conflict_col (id, v) VALUES ('k1', 'again') \ + ON CONFLICT (id) DO UPDATE SET id = EXCLUDED.id, v = EXCLUDED.v", + ) + .await + .expect("assigning the conflicting key keeps the identity"); + assert_eq!( + query(&server, "SELECT id FROM pk_conflict_col ORDER BY id").await, + vec!["k1".to_string(), "k2".to_string()] + ); + assert_eq!( + query(&server, "SELECT v FROM pk_conflict_col WHERE id = 'k1'").await, + vec!["again".to_string()] + ); +} diff --git a/nodedb/tests/wire/cases/sql_update_primary_key_identity.rs b/nodedb/tests/wire/cases/sql_update_primary_key_identity.rs new file mode 100644 index 000000000..999817c55 --- /dev/null +++ b/nodedb/tests/wire/cases/sql_update_primary_key_identity.rs @@ -0,0 +1,161 @@ +// SPDX-License-Identifier: BUSL-1.1 + +//! An UPDATE keeps a document's primary key. +//! +//! A document row stays stored under the key its primary key held at insert, +//! and an UPDATE rewrites the body in place. A changed key column would make +//! SQL reads and key reads name different documents, so a write that changes +//! it is refused with `23000`. Assigning the value the key already holds is +//! accepted. Changing a key is a DELETE plus an INSERT. + +use crate::harness::TestServer; + +/// SQLSTATE of a statement that must fail, or `None` when it succeeded. +async fn sqlstate_of(server: &TestServer, sql: &str) -> Option { + match server.client.simple_query(sql).await { + Ok(_) => None, + Err(e) => Some( + e.as_db_error() + .unwrap_or_else(|| panic!("expected a DbError from {sql}, got: {e}")) + .code() + .code() + .to_string(), + ), + } +} + +async fn assert_key_change_refused(server: &TestServer, sql: &str) { + let state = sqlstate_of(server, sql) + .await + .unwrap_or_else(|| panic!("a primary-key change must be refused, but ran: {sql}")); + assert_eq!( + state, "23000", + "a primary-key change is an integrity violation, got SQLSTATE {state} for: {sql}" + ); +} + +/// Create `name` with `CREATE COLLECTION {name}{shape}` and rows `k1` / `k2`. +async fn seeded(server: &TestServer, name: &str, shape: &str) { + server + .exec(&format!("CREATE COLLECTION {name}{shape}")) + .await + .expect("create collection"); + server + .exec(&format!( + "INSERT INTO {name} (id, v) VALUES ('k1', 'one'), ('k2', 'two')" + )) + .await + .expect("seed rows"); +} + +async fn assert_rows_unchanged(server: &TestServer, name: &str) { + let ids = server + .query_text(&format!("SELECT id FROM {name} ORDER BY id")) + .await + .expect("scan ids"); + assert_eq!(ids, vec!["k1".to_string(), "k2".to_string()]); + let by_key = server + .query_text(&format!("SELECT v FROM {name} WHERE id = 'k1'")) + .await + .expect("point read"); + assert_eq!( + by_key, + vec!["one".to_string()], + "k1 stays addressable by its key" + ); +} + +const SHAPES: [&str; 3] = [ + "", + " (id TEXT PRIMARY KEY, v TEXT)", + " (id TEXT PRIMARY KEY, v TEXT) WITH (engine='document_strict')", +]; + +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn point_update_refuses_a_new_primary_key() { + let server = TestServer::start().await; + for (i, shape) in SHAPES.iter().enumerate() { + let name = format!("pk_point_{i}"); + seeded(&server, &name, shape).await; + assert_key_change_refused( + &server, + &format!("UPDATE {name} SET id = 'moved' WHERE id = 'k1'"), + ) + .await; + assert_rows_unchanged(&server, &name).await; + } +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn predicate_and_multi_key_updates_refuse_a_new_primary_key() { + let server = TestServer::start().await; + for (i, shape) in SHAPES.iter().enumerate() { + let name = format!("pk_bulk_{i}"); + seeded(&server, &name, shape).await; + assert_key_change_refused( + &server, + &format!("UPDATE {name} SET id = 'moved' WHERE v = 'one'"), + ) + .await; + assert_key_change_refused( + &server, + &format!("UPDATE {name} SET id = v WHERE id IN ('k1', 'k2')"), + ) + .await; + assert_rows_unchanged(&server, &name).await; + } +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn transactional_update_refuses_a_new_primary_key() { + let server = TestServer::start().await; + seeded(&server, "pk_txn", "").await; + server.exec("BEGIN").await.expect("begin"); + assert_key_change_refused(&server, "UPDATE pk_txn SET id = 'moved' WHERE id = 'k1'").await; + server + .client + .simple_query("ROLLBACK") + .await + .expect("rollback"); + assert_rows_unchanged(&server, "pk_txn").await; +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn on_conflict_update_refuses_a_new_primary_key() { + let server = TestServer::start().await; + seeded(&server, "pk_conflict", " (id TEXT PRIMARY KEY, v TEXT)").await; + assert_key_change_refused( + &server, + "INSERT INTO pk_conflict (id, v) VALUES ('k1', 'again') \ + ON CONFLICT (id) DO UPDATE SET id = 'moved'", + ) + .await; + assert_rows_unchanged(&server, "pk_conflict").await; + server + .exec( + "INSERT INTO pk_conflict (id, v) VALUES ('k1', 'one') \ + ON CONFLICT (id) DO UPDATE SET id = EXCLUDED.id, v = EXCLUDED.v", + ) + .await + .expect("assigning the conflicting key keeps the identity"); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn assigning_the_current_primary_key_is_accepted() { + let server = TestServer::start().await; + for (i, shape) in SHAPES.iter().enumerate() { + let name = format!("pk_same_{i}"); + seeded(&server, &name, shape).await; + server + .exec(&format!( + "UPDATE {name} SET id = 'k1', v = 'one' WHERE id = 'k1'" + )) + .await + .expect("assigning the current key keeps the identity"); + server + .exec(&format!("UPDATE {name} SET id = id WHERE v = 'two'")) + .await + .expect("a self-assignment keeps the identity"); + assert_rows_unchanged(&server, &name).await; + } +} From 94bba7d3f3d328edfaa9db36bb715ab7313f05f5 Mon Sep 17 00:00:00 2001 From: Farhan Syah Date: Thu, 8 Oct 2026 16:51:47 +0800 Subject: [PATCH 10/24] fix(executor): reverse in-memory index writes of an abandoned write Point puts and deletes return undo entries for every R-tree, vector and sparse-vector mutation they make. A write abandoned after it applied (an enforcement refusal, a chain settle error, a commit error or a dropped batch) reverses them together with its cache entries and target row writes, and a failed reversal fail-stops the core. Undo failures carry a typed UndoError with their cause. An edge version the CSR refuses is taken back out of the edge store. Document collections no longer ingest geometry rows into a columnar memtable. --- .../src/collection/lifecycle_insert_ops.rs | 34 ++ .../src/data/executor/core_loop/fail_stop.rs | 24 +- .../executor/core_loop/graph_partition.rs | 5 +- .../core_loop/vector_index_rebuild.rs | 27 +- .../data/executor/enforcement/chain_guard.rs | 121 +++++- .../enforcement/materialized_sum/apply.rs | 20 +- .../handlers/columnar_mutation_apply.rs | 373 +++++++++++++----- .../handlers/columnar_write/spatial.rs | 89 ----- .../handlers/control/crdt_materialize.rs | 34 +- .../handlers/document/resolve/apply_row.rs | 136 +++++-- .../handlers/document/write/batch_insert.rs | 180 +++++++-- .../handlers/graph_edge_write/delete.rs | 75 ++-- .../handlers/graph_edge_write/delete_batch.rs | 86 ++-- .../executor/handlers/graph_edge_write/put.rs | 100 +++-- .../handlers/graph_edge_write/put_batch.rs | 54 ++- .../handlers/graph_edge_write/shared.rs | 27 +- .../handlers/merge_orchestrated/abort.rs | 5 +- .../merge_orchestrated/apply/update_rows.rs | 11 +- .../merge_orchestrated/apply_support.rs | 40 +- .../merge_orchestrated/delete_arms.rs | 31 +- .../executor/handlers/point/apply_delete.rs | 91 ++--- .../executor/handlers/point/apply_put/core.rs | 77 +++- .../handlers/point/apply_put/index.rs | 340 ++++++++++++---- .../handlers/point/apply_put/sparse.rs | 94 ++++- .../handlers/point/apply_put/types.rs | 28 +- .../handlers/point/apply_put/vector/put.rs | 344 +++++++++++----- .../handlers/point/apply_put/vector/remove.rs | 8 +- .../handlers/point/apply_put/vector/types.rs | 19 + .../data/executor/handlers/point/delete.rs | 291 +++++++++++++- .../data/executor/handlers/point/insert.rs | 66 ++-- .../src/data/executor/handlers/point/put.rs | 129 ++++-- .../handlers/point/update_reindex_sparse.rs | 21 +- .../handlers/point/update_reindex_vector.rs | 22 +- .../transaction/redo_apply/document.rs | 120 +++--- .../handlers/transaction/redo_apply/passes.rs | 14 +- .../handlers/transaction/undo/apply.rs | 266 ++++++++++--- .../transaction/undo/crdt_collection.rs | 11 +- .../handlers/transaction/undo/crdt_row.rs | 80 ++++ .../handlers/transaction/undo/document.rs | 113 +++--- .../handlers/transaction/undo/document_fts.rs | 35 +- .../transaction/undo/document_outcome.rs | 66 +--- .../handlers/transaction/undo/edge_cut.rs | 32 +- .../handlers/transaction/undo/edge_write.rs | 20 +- .../handlers/transaction/undo/entry.rs | 22 ++ .../handlers/transaction/undo/error.rs | 116 ++++++ .../handlers/transaction/undo/fts_doc.rs | 9 +- .../handlers/transaction/undo/graph_node.rs | 27 +- .../executor/handlers/transaction/undo/kv.rs | 20 +- .../handlers/transaction/undo/memory.rs | 75 ++++ .../executor/handlers/transaction/undo/mod.rs | 4 + .../handlers/transaction/undo/rollback.rs | 43 +- .../handlers/transaction/undo/spatial.rs | 8 +- .../handlers/transaction/undo/spatial_row.rs | 9 +- .../handlers/transaction/undo/stats.rs | 14 +- .../handlers/transaction/undo/timeseries.rs | 17 +- .../transaction/undo/truncate_columnar.rs | 6 +- .../transaction/undo/vector_truncate.rs | 9 +- .../handlers/transaction/undo/vector_write.rs | 35 +- .../executor/handlers/upsert/exec/insert.rs | 62 ++- .../handlers/upsert/exec/overwrite.rs | 78 ++-- .../src/data/executor/handlers/write_batch.rs | 58 ++- .../executor/wal_replay_columnar_image.rs | 8 +- .../wal_replay_redo_document_apply.rs | 44 ++- 63 files changed, 3200 insertions(+), 1223 deletions(-) delete mode 100644 nodedb/src/data/executor/handlers/columnar_write/spatial.rs create mode 100644 nodedb/src/data/executor/handlers/transaction/undo/crdt_row.rs create mode 100644 nodedb/src/data/executor/handlers/transaction/undo/error.rs create mode 100644 nodedb/src/data/executor/handlers/transaction/undo/memory.rs diff --git a/nodedb-vector/src/collection/lifecycle_insert_ops.rs b/nodedb-vector/src/collection/lifecycle_insert_ops.rs index 2c3713ba7..06699918e 100644 --- a/nodedb-vector/src/collection/lifecycle_insert_ops.rs +++ b/nodedb-vector/src/collection/lifecycle_insert_ops.rs @@ -268,6 +268,22 @@ impl VectorCollection { } false } + + /// Un-delete `id` and bind it to `surrogate` again: the reverse of + /// [`Self::delete`] on a bound node, which drops the binding. + /// + /// Returns `false` and changes nothing when `id` carries no tombstone. + /// [`Surrogate::ZERO`] binds nothing, as in [`Self::insert_with_surrogate`]. + pub fn undelete_bound(&mut self, id: u32, surrogate: Surrogate) -> bool { + if !self.undelete(id) { + return false; + } + if surrogate != Surrogate::ZERO { + self.surrogate_map.insert(id, surrogate); + self.surrogate_to_local.insert(surrogate, id); + } + true + } } /// The FP32 vector at `local` in a sealed segment: the mmap tier when the @@ -325,6 +341,24 @@ mod tests { assert!(!coll.delete(first), "the first node is already gone"); } + #[test] + fn undelete_bound_restores_the_node_and_its_binding() { + let mut coll = collection(); + let s = Surrogate::new(5); + let id = coll.insert_with_surrogate(vec![1.0, 0.0], s).unwrap(); + assert!(coll.delete(id)); + assert_eq!(coll.local_for_surrogate(s), None); + + assert!(coll.undelete_bound(id, s)); + assert_eq!(coll.live_count(), 1); + assert_eq!(coll.local_for_surrogate(s), Some(id)); + assert_eq!(coll.get_surrogate(id), Some(s)); + assert!( + !coll.undelete_bound(id, s), + "a live node has no tombstone to clear" + ); + } + #[test] fn delete_by_surrogate_is_idempotent() { let mut coll = collection(); diff --git a/nodedb/src/data/executor/core_loop/fail_stop.rs b/nodedb/src/data/executor/core_loop/fail_stop.rs index 28e48bf81..5f81ddcff 100644 --- a/nodedb/src/data/executor/core_loop/fail_stop.rs +++ b/nodedb/src/data/executor/core_loop/fail_stop.rs @@ -112,12 +112,24 @@ impl CoreLoop { /// Fail-stop the core when `response` reports a failed rollback. pub(in crate::data::executor) fn fail_stop_on_rollback_failure(&mut self, response: &Response) { - if let Some(ErrorCode::RollbackFailed { + if let Some(code) = response.error_code.as_deref() { + self.fail_stop_on_rollback_code(code); + } + } + + /// Fail-stop the core when `code` is `RollbackFailed`. The logged detail + /// names the undo entry, the reverse write, and its typed cause. + pub(in crate::data::executor) fn fail_stop_on_rollback_code(&mut self, code: &ErrorCode) { + if let ErrorCode::RollbackFailed { entry_index, detail, - }) = response.error_code.as_deref() + cause, + } = code { - let detail = format!("undo entry {entry_index}: {detail}"); + let detail = match cause { + Some(cause) => format!("undo entry {entry_index}: {detail}; cause: {cause:?}"), + None => format!("undo entry {entry_index}: {detail}"), + }; self.fail_stop_core(FailStopCause::RollbackFailed, &detail); } } @@ -186,12 +198,18 @@ mod tests { ErrorCode::RollbackFailed { entry_index: 2, detail: "restore failed".into(), + cause: Some(Box::new(ErrorCode::Internal { + detail: "storage error (sparse): commit".into(), + })), }, ); core.fail_stop_on_rollback_failure(&response); assert!(core.fail_stop.is_stopped()); + let (cause, detail) = core.fail_stop.cause().expect("stopped"); + assert_eq!(cause, FailStopCause::RollbackFailed); + assert!(detail.starts_with("undo entry 2: restore failed; cause: Internal")); assert_eq!(metrics.core_fail_stops.stopped_cores(), 1); assert_eq!( metrics.core_fail_stops.report().map(|r| r.cause), diff --git a/nodedb/src/data/executor/core_loop/graph_partition.rs b/nodedb/src/data/executor/core_loop/graph_partition.rs index 379c523d4..776b55eb5 100644 --- a/nodedb/src/data/executor/core_loop/graph_partition.rs +++ b/nodedb/src/data/executor/core_loop/graph_partition.rs @@ -68,8 +68,9 @@ impl CoreLoop { /// Reverse a [`CoreLoop::mark_node_deleted`] by removing `node_id` of /// `collection` from the caller's `(database, tenant)` deleted-nodes /// set. Called only on transaction rollback, and only for a node THIS - /// transaction newly marked (see the `was_newly_marked` capture in - /// `apply_point_delete`), so it never removes a pre-existing tombstone. + /// transaction newly marked (see the `MarkNodeDeleted` undo entry + /// `apply_point_delete` pushes), so it never removes a pre-existing + /// tombstone. #[inline] pub(in crate::data::executor) fn unmark_node_deleted( &mut self, diff --git a/nodedb/src/data/executor/core_loop/vector_index_rebuild.rs b/nodedb/src/data/executor/core_loop/vector_index_rebuild.rs index 129c9d99d..5fb63b3c8 100644 --- a/nodedb/src/data/executor/core_loop/vector_index_rebuild.rs +++ b/nodedb/src/data/executor/core_loop/vector_index_rebuild.rs @@ -103,16 +103,21 @@ impl CoreLoop { crate::engine::document::store::StorageKey::for_surrogate(surrogate); // Same as WAL replay: the document is already durable, so a // width mismatch from before the forward-path check existed is - // reported and skipped rather than aborting the rebuild. - let deltas = match self.apply_point_put_vector_indexes(VectorIndexPutParams { - database_id: db, - tid: tenant_id, - collection: &collection, - storage_key, - value: &value, - wal_lsn: 0, - }) { - Ok(deltas) => deltas, + // reported and skipped rather than aborting the rebuild. The + // rebuild re-derives durable rows and has no write to abandon, + // so the undo entries go unused. + let inserted = match self.apply_point_put_vector_indexes( + VectorIndexPutParams { + database_id: db, + tid: tenant_id, + collection: &collection, + storage_key, + value: &value, + wal_lsn: 0, + }, + &mut Vec::new(), + ) { + Ok(inserted) => inserted, Err(e) => { tracing::warn!( core = self.core_id, @@ -124,7 +129,7 @@ impl CoreLoop { continue; } }; - if !deltas.is_empty() { + if inserted > 0 { rebuilt += 1; } } diff --git a/nodedb/src/data/executor/enforcement/chain_guard.rs b/nodedb/src/data/executor/enforcement/chain_guard.rs index b15cbfbe5..542cfc027 100644 --- a/nodedb/src/data/executor/enforcement/chain_guard.rs +++ b/nodedb/src/data/executor/enforcement/chain_guard.rs @@ -21,6 +21,10 @@ use redb::WriteTransaction; use crate::data::executor::core_loop::CoreLoop; use crate::data::executor::doc_format; use crate::data::executor::enforcement::hash_chain::{self, ChainHead}; +use crate::data::executor::enforcement::materialized_sum::apply::TargetWrite; +use crate::data::executor::handlers::transaction::undo::UndoEntry; +use crate::data::executor::handlers::transaction::undo::memory::abort_error; +use crate::engine::document::store::StorageKey; use crate::types::{DatabaseId, TenantId}; /// The chain link a pending write asks `build_stored_body` to write. @@ -285,27 +289,114 @@ impl ChainGuard { } } +/// A document write abandoned after `apply_point_put` ran, before its +/// transaction commits. One write can store several rows of one collection. +pub(in crate::data::executor) struct AbandonedWrite<'a> { + pub database_id: u64, + pub tid: u64, + pub collection: &'a str, + /// Every row `apply_point_put` was called for, including a row whose + /// call failed: it cached the row before it failed. + pub row_keys: &'a [StorageKey], + /// The rows' in-memory index entries to reverse, from + /// `PointPutOutcome::memory_undo` or `PointDeleteOutcome::memory_undo`, + /// in the order the rows were written. + /// A row whose `apply_point_put` failed has none: the call reverses its + /// own entries before it returns. + pub memory_undo: Vec, + /// The target rows the write's enforcement updated. Empty before + /// enforcement ran. + pub target_writes: Vec, +} + +impl<'a> AbandonedWrite<'a> { + /// A write of one row, with no in-memory entries to reverse yet. + pub(in crate::data::executor) fn row( + database_id: u64, + tid: u64, + collection: &'a str, + row_key: &'a StorageKey, + ) -> Self { + Self::rows(database_id, tid, collection, std::slice::from_ref(row_key)) + } + + /// A write of several rows, with no in-memory entries to reverse yet. + pub(in crate::data::executor) fn rows( + database_id: u64, + tid: u64, + collection: &'a str, + row_keys: &'a [StorageKey], + ) -> Self { + Self { + database_id, + tid, + collection, + row_keys, + memory_undo: Vec::new(), + target_writes: Vec::new(), + } + } + + /// The rows' `PointPutOutcome::memory_undo` or + /// `PointDeleteOutcome::memory_undo`. + pub(in crate::data::executor) fn undo(mut self, memory_undo: Vec) -> Self { + self.memory_undo = memory_undo; + self + } + + /// The target rows the write's enforcement updated. + pub(in crate::data::executor) fn targets(mut self, target_writes: Vec) -> Self { + self.target_writes = target_writes; + self + } +} + /// Undo the in-memory side effects an abort AFTER `apply_point_put` leaves -/// behind, before the caller drops its transaction uncommitted. +/// behind, before the caller drops its transaction uncommitted. Returns the +/// error the caller reports: `error`, or `RollbackFailed` when an in-memory +/// entry did not reverse. /// -/// `apply_point_put` populates the read-through document cache with the body it -/// wrote. Dropping the redb transaction reverses the durable write but not that -/// cache entry, so every subsequent read of the row would be served the -/// post-image of a write that never landed — a row visible to readers and -/// absent from storage. Restoring the hash-chain head is the same class of -/// in-memory reversal, so both happen here rather than one being remembered at -/// each abort site and the other forgotten. +/// Dropping the redb transaction reverses the durable writes and nothing in +/// memory. Four things stay behind, and all four are reversed here, so no +/// abort site remembers one and forgets another: +/// - the document-cache entries `apply_point_put` wrote: reads would serve +/// rows absent from storage +/// - the R-tree, vector and sparse entries of the rows: spatial predicates +/// and vector search would answer with them +/// - the same cache and index entries of every materialized-sum target row +/// - the advanced hash-chain head pub(in crate::data::executor) fn abort_after_apply( core: &mut CoreLoop, guard: &mut ChainGuard, - database_id: u64, - tid: u64, - collection: &str, - row_key: &crate::engine::document::store::StorageKey, -) { + write: AbandonedWrite<'_>, + error: crate::Error, +) -> crate::Error { guard.restore(core); - core.doc_cache - .invalidate(database_id, tid, collection, row_key); + abandon_write(core, write, error) +} + +/// [`abort_after_apply`] for a write that set no hash-chain intent. +pub(in crate::data::executor) fn abandon_write( + core: &mut CoreLoop, + write: AbandonedWrite<'_>, + error: crate::Error, +) -> crate::Error { + let AbandonedWrite { + database_id, + tid, + collection, + row_keys, + memory_undo, + target_writes, + } = write; + for row_key in row_keys { + core.doc_cache + .invalidate(database_id, tid, collection, row_key); + } + // A target is written after its row, so targets are reversed first. + let targets = core.abandon_target_writes(database_id, tid, target_writes); + let row = core.undo_memory_effects(database_id, tid, memory_undo); + abort_error(error, targets.and(row)) } impl CoreLoop { diff --git a/nodedb/src/data/executor/enforcement/materialized_sum/apply.rs b/nodedb/src/data/executor/enforcement/materialized_sum/apply.rs index e4e259456..699ed597d 100644 --- a/nodedb/src/data/executor/enforcement/materialized_sum/apply.rs +++ b/nodedb/src/data/executor/enforcement/materialized_sum/apply.rs @@ -78,6 +78,7 @@ use super::rmw::BalanceRmw; use crate::data::executor::core_loop::CoreLoop; use crate::data::executor::enforcement::images::{EnforcementCtx, RowImages}; use crate::data::executor::handlers::point::apply_put::PointPutOutcome; +use crate::data::executor::handlers::transaction::undo::memory::abort_error; use crate::types::DatabaseId; /// A target row this write updated, captured so a transactional caller can @@ -188,20 +189,11 @@ impl CoreLoop { Err(e) => { // The caller drops `txn`, which reverses every target // row this pass already wrote — but not the read-through - // cache entries those writes populated. Left behind, they - // serve balances that no longer exist in storage. - for write in &writes { - let key = crate::engine::document::store::StorageKey::for_surrogate( - write.surrogate, - ); - self.doc_cache.invalidate( - ctx.database_id, - ctx.tid, - &write.collection, - &key, - ); - } - return Err(e); + // cache entries or the in-memory index entries those + // writes made. Left behind, they serve balances that no + // longer exist in storage. + let undone = self.abandon_target_writes(ctx.database_id, ctx.tid, writes); + return Err(abort_error(e, undone)); } } } diff --git a/nodedb/src/data/executor/handlers/columnar_mutation_apply.rs b/nodedb/src/data/executor/handlers/columnar_mutation_apply.rs index 6d9998661..d5d3c43b7 100644 --- a/nodedb/src/data/executor/handlers/columnar_mutation_apply.rs +++ b/nodedb/src/data/executor/handlers/columnar_mutation_apply.rs @@ -12,13 +12,13 @@ use nodedb_columnar::pk_index::{RowLocation, encode_pk}; use nodedb_types::Value; use nodedb_types::columnar::ColumnarSchema; -use tracing::warn; use crate::data::executor::core_loop::CoreLoop; use crate::data::executor::handlers::columnar_write::{ GeometryIndexDelta, RemovedSpatialEntry, row_values_to_object, schema_has_geometry, }; use crate::data::executor::handlers::transaction::undo::UndoEntry; +use crate::data::executor::handlers::transaction::undo::memory::abort_error; use crate::data::executor::task::ExecutionTask; /// Engine map key of one columnar collection. @@ -66,6 +66,12 @@ impl CoreLoop { /// columns the old row's R-tree entries go and the post-image is /// indexed. When `undo_log` is `Some`, every in-memory change this /// makes is recorded on it. + /// + /// The statement applies whole or not at all. Every pre-image is read + /// before the first row is mutated, so a pre-image that does not read + /// changes nothing. A row the engine refuses reverses the rows this + /// statement already changed and returns the engine's error. A reversal + /// that fails returns `RollbackFailed`, which fail-stops the core. pub(in crate::data::executor) fn apply_columnar_update_rows( &mut self, task: &ExecutionTask, @@ -73,10 +79,14 @@ impl CoreLoop { schema: &ColumnarSchema, rows: &[(Value, Vec)], undo_log: Option<&mut Vec>, - ) -> ColumnarUpdateOutcome { - let track = undo_log.is_some(); + ) -> crate::Result { let has_geometry = schema_has_geometry(schema); let collection = key.2.as_str(); + let row_count_before = self + .columnar_engines + .get(key) + .map(|engine| engine.memtable().row_count()) + .ok_or_else(|| engine_missing(collection))?; let mut outcome = ColumnarUpdateOutcome { affected: 0, inserted_pks: Vec::new(), @@ -85,67 +95,71 @@ impl CoreLoop { }; let mut spatial_removed: Vec = Vec::new(); let mut geometry_delta = GeometryIndexDelta::default(); + // Every pre-image is read before the first update unbinds a PK: it + // names the R-tree entries the old row owns. + let old_pks = rows.iter().map(|r| &r.0); + let pre_images = self.read_geometry_pre_images(key, has_geometry, old_pks)?; - for (old_pk, new_row) in rows { + for ((old_pk, new_row), pre_image) in rows.iter().zip(pre_images) { let old_pk_bytes = encode_pk(old_pk); - // The pre-image is read BEFORE the update unbinds the PK: it is - // what names the R-tree entries the old row owns. - let pre_image = if has_geometry { - self.read_columnar_row_by_pk(key, &old_pk_bytes) - } else { - None - }; - let Some(engine) = self.columnar_engines.get_mut(key) else { - break; - }; - // Capture the pre-image BEFORE mutating (the update removes the - // old PK binding and appends a new row): the tombstoned - // original's location, the appended replacement's PK, and — for - // a PK-changing update — the memtable row its insert half - // displaces. - let capture = if track { - let old_location = engine.pk_index().get(&old_pk_bytes).copied(); - let new_pk_bytes = engine.encode_pk_from_row(new_row).ok(); - let displaced_entry = match &new_pk_bytes { - Some(nb) if *nb != old_pk_bytes => engine - .pk_index() - .get(nb) - .copied() - .filter(|loc| loc.segment_id == engine.memtable_segment_id()) - .map(|loc| (nb.clone(), loc)), - _ => None, - }; - Some((old_location, new_pk_bytes, displaced_entry)) - } else { - None + let applied = match self.columnar_engines.get_mut(key) { + Some(engine) => { + // Capture before mutating (the update removes the old PK + // binding and appends a new row): the tombstoned + // original's location, the appended replacement's PK, + // and, for a PK-changing update, the row its insert half + // displaces. That row is in the memtable or a flushed + // segment, and the undo puts back either one. + let old_location = engine.pk_index().get(&old_pk_bytes).copied(); + let new_pk_bytes = engine.encode_pk_from_row(new_row).ok(); + let displaced_entry = match &new_pk_bytes { + Some(nb) if *nb != old_pk_bytes => engine + .pk_index() + .get(nb) + .copied() + .map(|loc| (nb.clone(), loc)), + _ => None, + }; + // A flushed row's surrogate lives in its segment's + // sidecar, which the engine does not hold. The + // replacement row keeps it. + let flushed_surrogate = old_location + .filter(|loc| loc.segment_id != engine.memtable_segment_id()) + .and_then(|loc| { + flushed_row_surrogate(&self.columnar_flushed_surrogates, key, loc) + }); + engine + .update(old_pk, new_row, flushed_surrogate) + .map(|_| (old_location, new_pk_bytes, displaced_entry)) + .map_err(crate::Error::from) + } + None => Err(engine_missing(collection)), }; - // A flushed row's surrogate lives in its segment's sidecar, which - // the engine does not hold; the replacement row keeps it. - let flushed_surrogate = engine - .pk_index() - .get(&old_pk_bytes) - .filter(|loc| loc.segment_id != engine.memtable_segment_id()) - .and_then(|loc| { - flushed_row_surrogate(&self.columnar_flushed_surrogates, key, *loc) - }); - match engine.update(old_pk, new_row, flushed_surrogate) { - Ok(_) => {} - Err(e) => { - warn!(core = self.core_id, %collection, error = %e, "columnar update row failed"); - continue; + let (old_location, new_pk_bytes, displaced_entry) = match applied { + Ok(capture) => capture, + Err(error) => { + let mut undo = Vec::new(); + push_removed_spatial_undo(&mut undo, spatial_removed); + Self::push_geometry_index_undo(&mut undo, geometry_delta); + undo.push(UndoEntry::ColumnarUpdate { + collection_key: key.clone(), + row_count_before, + inserted_pks: outcome.inserted_pks, + displaced: outcome.displaced, + restored: outcome.restored, + }); + return Err(self.reverse_columnar_statement(key, undo, error)); } - } + }; outcome.affected += 1; - if let Some((old_location, new_pk_bytes, displaced_entry)) = capture { - if let Some(nb) = new_pk_bytes { - outcome.inserted_pks.push(nb); - } - if let Some(loc) = old_location { - outcome.restored.push((old_pk_bytes, loc)); - } - if let Some(d) = displaced_entry { - outcome.displaced.push(d); - } + if let Some(nb) = new_pk_bytes { + outcome.inserted_pks.push(nb); + } + if let Some(loc) = old_location { + outcome.restored.push((old_pk_bytes, loc)); + } + if let Some(d) = displaced_entry { + outcome.displaced.push(d); } if has_geometry { if let Some(row) = &pre_image { @@ -168,21 +182,26 @@ impl CoreLoop { push_removed_spatial_undo(log, spatial_removed); Self::push_geometry_index_undo(log, geometry_delta); } - outcome + Ok(outcome) } /// Remove the rows bound to `pks` through `MutationEngine::delete`. For /// a collection with geometry columns each row's R-tree entries go too. /// When `undo_log` is `Some`, every in-memory change this makes is /// recorded on it. + /// + /// The statement applies whole or not at all. Every pre-image is read + /// before the first row is removed, so a pre-image that does not read + /// changes nothing. A row the engine refuses restores the rows this + /// statement already removed and returns the engine's error. A reversal + /// that fails returns `RollbackFailed`, which fail-stops the core. pub(in crate::data::executor) fn apply_columnar_delete_pks( &mut self, key: &ColumnarEngineKey, schema: &ColumnarSchema, pks: &[Value], undo_log: Option<&mut Vec>, - ) -> ColumnarDeleteOutcome { - let track = undo_log.is_some(); + ) -> crate::Result { let has_geometry = schema_has_geometry(schema); let collection = key.2.as_str(); let mut outcome = ColumnarDeleteOutcome { @@ -190,37 +209,37 @@ impl CoreLoop { restored: Vec::new(), }; let mut spatial_removed: Vec = Vec::new(); + let pre_images = self.read_geometry_pre_images(key, has_geometry, pks.iter())?; - for pk in pks { + for (pk, pre_image) in pks.iter().zip(pre_images) { let pk_bytes = encode_pk(pk); - let pre_image = if has_geometry { - self.read_columnar_row_by_pk(key, &pk_bytes) - } else { - None - }; - let Some(engine) = self.columnar_engines.get_mut(key) else { - break; - }; - // Read the location BEFORE the delete removes the PK binding. - let captured = if track { - engine - .pk_index() - .get(&pk_bytes) - .copied() - .map(|loc| (pk_bytes.clone(), loc)) - } else { - None + let applied = match self.columnar_engines.get_mut(key) { + Some(engine) => { + // Read the location before the delete removes the PK + // binding. + let location = engine.pk_index().get(&pk_bytes).copied(); + engine + .delete(pk) + .map(|_| location) + .map_err(crate::Error::from) + } + None => Err(engine_missing(collection)), }; - match engine.delete(pk) { - Ok(_) => {} - Err(e) => { - warn!(core = self.core_id, %collection, error = %e, "columnar delete row failed"); - continue; + let location = match applied { + Ok(location) => location, + Err(error) => { + let mut undo = Vec::new(); + push_removed_spatial_undo(&mut undo, spatial_removed); + undo.push(UndoEntry::ColumnarDelete { + collection_key: key.clone(), + restored: outcome.restored, + }); + return Err(self.reverse_columnar_statement(key, undo, error)); } - } + }; outcome.affected += 1; - if let Some(entry) = captured { - outcome.restored.push(entry); + if let Some(loc) = location { + outcome.restored.push((pk_bytes, loc)); } if let Some(row) = &pre_image { spatial_removed.extend( @@ -232,7 +251,50 @@ impl CoreLoop { if let Some(log) = undo_log { push_removed_spatial_undo(log, spatial_removed); } - outcome + Ok(outcome) + } + + /// The pre-image of each row `pks` binds, in order, for a collection + /// with geometry columns. One `None` per PK when `has_geometry` is false, + /// so the caller reads nothing it does not need. + /// + /// `Err` on the first pre-image that does not read: a corrupt row would + /// otherwise keep its R-tree entries after it is gone. + fn read_geometry_pre_images<'v>( + &self, + key: &ColumnarEngineKey, + has_geometry: bool, + pks: impl Iterator, + ) -> crate::Result>>> { + pks.map(|pk| { + if has_geometry { + self.read_columnar_row_by_pk(key, &encode_pk(pk)) + } else { + Ok(None) + } + }) + .collect() + } + + /// Reverse the in-memory changes `undo` records for a statement that + /// stopped at `error`, last change first. The answer is `error`, or + /// `RollbackFailed` when a change does not reverse: the core's state is + /// then unknown, and the core fail-stops when that error leaves it. + fn reverse_columnar_statement( + &mut self, + key: &ColumnarEngineKey, + undo: Vec, + error: crate::Error, + ) -> crate::Error { + let reversal = self.undo_memory_effects(key.0.as_u64(), key.1.as_u64(), undo); + abort_error(error, reversal) + } +} + +/// The error for a columnar engine that is absent mid-statement. +fn engine_missing(collection: &str) -> crate::Error { + crate::Error::Internal { + detail: format!("columnar engine not found for collection '{collection}'"), } } @@ -287,16 +349,18 @@ mod tests { core.columnar_flushed_surrogates .insert(key.clone(), vec![sidecar]); - let outcome = core.apply_columnar_update_rows( - &make_default_task(), - &key, - &schema(), - &[( - Value::Integer(1), - vec![Value::Integer(1), Value::Integer(99)], - )], - None, - ); + let outcome = core + .apply_columnar_update_rows( + &make_default_task(), + &key, + &schema(), + &[( + Value::Integer(1), + vec![Value::Integer(1), Value::Integer(99)], + )], + None, + ) + .expect("apply"); assert_eq!(outcome.affected, 1); let live: Vec<(Option, Vec)> = core @@ -304,7 +368,8 @@ mod tests { .get(&key) .expect("engine") .scan_memtable_rows_with_surrogates() - .collect(); + .collect::>() + .expect("read"); assert_eq!( live, vec![( @@ -313,4 +378,112 @@ mod tests { )] ); } + + fn two_row_key(core: &mut CoreLoop) -> ColumnarEngineKey { + let key: ColumnarEngineKey = ( + DatabaseId::DEFAULT, + crate::types::TenantId::new(1), + "m".to_string(), + ); + let mut engine = nodedb_columnar::MutationEngine::new("m".to_string(), schema()); + engine + .insert(&[Value::Integer(1), Value::Integer(10)]) + .expect("insert"); + engine + .insert(&[Value::Integer(2), Value::Integer(20)]) + .expect("insert"); + core.columnar_engines.insert(key.clone(), engine); + key + } + + fn live_rows(core: &CoreLoop, key: &ColumnarEngineKey) -> Vec> { + core.columnar_engines + .get(key) + .expect("engine") + .scan_memtable_rows() + .collect::>() + .expect("read") + } + + fn bound_row(core: &CoreLoop, key: &ColumnarEngineKey, pk: i64) -> Option { + core.columnar_engines + .get(key) + .expect("engine") + .pk_index() + .get(&encode_pk(&Value::Integer(pk))) + .copied() + } + + #[test] + fn an_update_refused_partway_reverses_the_rows_it_already_changed() { + let dir = tempfile::tempdir().expect("tempdir"); + let (mut core, _tx, _rx) = make_core_with_dir(dir.path()); + let key = two_row_key(&mut core); + + // Row 2's post-image breaks NOT NULL, after row 1 is already updated. + let result = core.apply_columnar_update_rows( + &make_default_task(), + &key, + &schema(), + &[ + ( + Value::Integer(1), + vec![Value::Integer(1), Value::Integer(11)], + ), + (Value::Integer(2), vec![Value::Integer(2), Value::Null]), + ], + None, + ); + + match result { + Err(crate::Error::RejectedConstraint { .. }) => {} + Err(other) => panic!("expected the NOT NULL refusal, got {other:?}"), + Ok(outcome) => panic!("expected a refusal, {} rows applied", outcome.affected), + } + assert_eq!( + live_rows(&core, &key), + vec![ + vec![Value::Integer(1), Value::Integer(10)], + vec![Value::Integer(2), Value::Integer(20)], + ] + ); + assert_eq!( + core.columnar_engines + .get(&key) + .expect("engine") + .memtable() + .row_count(), + 2 + ); + assert_eq!(bound_row(&core, &key, 1).map(|loc| loc.row_index), Some(0)); + } + + #[test] + fn a_delete_refused_partway_restores_the_rows_it_already_removed() { + let dir = tempfile::tempdir().expect("tempdir"); + let (mut core, _tx, _rx) = make_core_with_dir(dir.path()); + let key = two_row_key(&mut core); + + // PK 99 is unbound, after row 1 is already removed. + let result = core.apply_columnar_delete_pks( + &key, + &schema(), + &[Value::Integer(1), Value::Integer(99)], + None, + ); + + match result { + Err(crate::Error::Storage { .. }) => {} + Err(other) => panic!("expected the unbound-key refusal, got {other:?}"), + Ok(outcome) => panic!("expected a refusal, {} rows applied", outcome.affected), + } + assert_eq!( + live_rows(&core, &key), + vec![ + vec![Value::Integer(1), Value::Integer(10)], + vec![Value::Integer(2), Value::Integer(20)], + ] + ); + assert_eq!(bound_row(&core, &key, 1).map(|loc| loc.row_index), Some(0)); + } } diff --git a/nodedb/src/data/executor/handlers/columnar_write/spatial.rs b/nodedb/src/data/executor/handlers/columnar_write/spatial.rs deleted file mode 100644 index 63720d0ee..000000000 --- a/nodedb/src/data/executor/handlers/columnar_write/spatial.rs +++ /dev/null @@ -1,89 +0,0 @@ -// SPDX-License-Identifier: BUSL-1.1 - -//! Spatial-side columnar ingest: build a row from a JSON object (first-row -//! schema inference on engine creation). - -use nodedb_columnar::MutationEngine; -use nodedb_types::columnar::schema::{TS_SYSTEM, TS_VALID_FROM, TS_VALID_UNTIL}; -use nodedb_types::value::Value; - -use crate::data::executor::core_loop::CoreLoop; - -use super::schema::{infer_schema_from_json, ndb_field_to_value, prepend_bitemporal_columns}; - -impl CoreLoop { - /// Ingest a single JSON document into the columnar engine for a collection. - /// - /// Creates the engine on first call. Used by the spatial insert path. - pub(in crate::data::executor) fn ingest_doc_to_columnar( - &mut self, - database_id: u64, - tid: u64, - collection: &str, - obj: &serde_json::Map, - ) { - let engine_key = ( - nodedb_types::DatabaseId::new(database_id), - crate::types::TenantId::new(tid), - collection.to_string(), - ); - let bitemporal = self.is_bitemporal(database_id, tid, collection); - let sys_now = if bitemporal { - self.bitemporal_now_ms() - } else { - 0 - }; - if !self.columnar_engines.contains_key(&engine_key) { - let base_schema = infer_schema_from_json(&serde_json::Value::Object(obj.clone())); - let schema = if bitemporal { - prepend_bitemporal_columns(base_schema) - } else { - base_schema - }; - let engine = MutationEngine::with_flush_threshold( - collection.to_string(), - schema, - self.query_tuning.columnar_flush_threshold, - ); - self.columnar_engines.insert(engine_key.clone(), engine); - } - - let Some(engine) = self.columnar_engines.get_mut(&engine_key) else { - return; - }; - let schema = engine.schema().clone(); - - let ndb_obj: std::collections::HashMap = obj - .iter() - .map(|(k, v)| (k.clone(), Value::from(v.clone()))) - .collect(); - let values: Vec = match schema - .columns - .iter() - .map(|col| match col.name.as_str() { - TS_SYSTEM if bitemporal => Ok(Value::Integer(sys_now)), - TS_VALID_FROM if bitemporal => Ok(match ndb_obj.get(TS_VALID_FROM) { - Some(Value::Integer(i)) => Value::Integer(*i), - _ => Value::Integer(i64::MIN), - }), - TS_VALID_UNTIL if bitemporal => Ok(match ndb_obj.get(TS_VALID_UNTIL) { - Some(Value::Integer(i)) => Value::Integer(*i), - _ => Value::Integer(i64::MAX), - }), - _ => ndb_field_to_value(ndb_obj.get(&col.name), &col.column_type), - }) - .collect::, crate::Error>>() - { - Ok(v) => v, - Err(e) => { - tracing::warn!( - collection, - "spatial columnar ingest skipped: timestamp coercion overflow: {e}" - ); - return; - } - }; - - let _ = engine.insert(&values); - } -} diff --git a/nodedb/src/data/executor/handlers/control/crdt_materialize.rs b/nodedb/src/data/executor/handlers/control/crdt_materialize.rs index b23ce921a..796c1d433 100644 --- a/nodedb/src/data/executor/handlers/control/crdt_materialize.rs +++ b/nodedb/src/data/executor/handlers/control/crdt_materialize.rs @@ -36,6 +36,7 @@ use crate::data::executor::handlers::transaction::undo::document_outcome::{ use nodedb_types::{RowIdentity, Surrogate}; use crate::data::executor::core_loop::CoreLoop; +use crate::data::executor::enforcement::chain_guard::{AbandonedWrite, abandon_write}; use crate::data::executor::handlers::point::apply_put::PointPutParams; use crate::data::executor::task::ExecutionTask; use crate::engine::crdt::tenant_state::TenantCrdtEngine; @@ -152,8 +153,11 @@ impl CoreLoop { // No chain guard: DDL refuses HASH_CHAIN on a CRDT collection, so this // write never reaches a chained row. + // An abort drops the txn uncommitted, which reverses the durable + // writes. `abandon_write` reverses the cache and in-memory index + // entries. let txn = self.sparse.begin_write()?; - let outcome = self.apply_point_put( + let outcome = match self.apply_point_put( &txn, PointPutParams { database_id, @@ -169,11 +173,27 @@ impl CoreLoop { wal_lsn: task.wal_lsn(), resolved_targets: &[], }, - )?; - txn.commit().map_err(|e| crate::Error::Storage { - engine: "sparse".into(), - detail: format!("crdt materialize commit: {e}"), - })?; + ) { + Ok(outcome) => outcome, + Err(e) => { + return Err(abandon_write( + self, + AbandonedWrite::row(database_id, tid, collection, &storage_key), + e, + )); + } + }; + if let Err(e) = txn.commit() { + return Err(abandon_write( + self, + AbandonedWrite::row(database_id, tid, collection, &storage_key) + .undo(outcome.memory_undo), + crate::Error::Storage { + engine: "sparse".into(), + detail: format!("crdt materialize commit: {e}"), + }, + )); + } self.checkpoint_coordinator.mark_dirty("sparse", 1); @@ -193,8 +213,6 @@ impl CoreLoop { push_put_undo( &mut undo, DocumentRow { - database_id, - tid, collection, storage_key, }, diff --git a/nodedb/src/data/executor/handlers/document/resolve/apply_row.rs b/nodedb/src/data/executor/handlers/document/resolve/apply_row.rs index 50bd2a4fe..102f17d4d 100644 --- a/nodedb/src/data/executor/handlers/document/resolve/apply_row.rs +++ b/nodedb/src/data/executor/handlers/document/resolve/apply_row.rs @@ -12,10 +12,13 @@ use nodedb_types::Surrogate; use crate::bridge::envelope::{ErrorCode, WriteSetEntry}; use crate::data::executor::core_loop::CoreLoop; use crate::data::executor::core_loop::redo_image::{removed_row_image, submitted_row_image}; -use crate::data::executor::enforcement::chain_guard::{ChainGuard, abort_after_apply}; +use crate::data::executor::enforcement::chain_guard::{ + AbandonedWrite, ChainGuard, abandon_write, abort_after_apply, +}; use crate::data::executor::enforcement::write_hook::{self, HookCtx, ImageBody, WriteImages}; use crate::data::executor::handlers::point::apply_delete::PointDeleteParams; -use crate::data::executor::handlers::point::apply_put::PointPutParams; +use crate::data::executor::handlers::point::apply_put::{PointPutParams, VectorIndexDelta}; +use crate::data::executor::handlers::transaction::undo::UndoEntry; use crate::data::executor::task::ExecutionTask; use crate::engine::document::store::{RowIdentity, StorageKey}; @@ -67,12 +70,6 @@ impl CoreLoop { let row_identity = RowIdentity::from_user_key(document_id); let has_vectors = self.collection_has_vectors(database_id, tid, collection); - // HNSW insert appends rather than replaces, so the prior embedding - // must come out first or KNN keeps scoring both. - if has_vectors && precondition.is_some() { - self.remove_document_vector_indexes(database_id, tid, collection, storage_key); - } - // A row absent before this write is a link of a HASH_CHAIN target. let mut chain = ChainGuard::begin(self, database_id, tid, collection); if precondition.is_none() { @@ -87,6 +84,18 @@ impl CoreLoop { return Err(ErrorCode::from(e)); } }; + + // HNSW insert appends rather than replaces, so the prior embedding + // must come out first or KNN keeps scoring both. The removal is in + // memory, so an abort below puts the prior nodes back. + let mut memory_undo: Vec = Vec::new(); + if has_vectors && precondition.is_some() { + memory_undo.extend( + self.remove_document_vector_indexes(database_id, tid, collection, storage_key) + .into_iter() + .map(VectorIndexDelta::into_delete_undo), + ); + } let mut outcome = match self.apply_point_put( &txn, PointPutParams { @@ -106,16 +115,29 @@ impl CoreLoop { ) { Ok(outcome) => outcome, Err(e) => { - // Dropping `txn` reverses the write but not the cache entry. - abort_after_apply(self, &mut chain, database_id, tid, collection, &storage_key); + // Dropping `txn` reverses the write but not the cache entry + // or the prior vector removal. + let e = abort_after_apply( + self, + &mut chain, + AbandonedWrite::row(database_id, tid, collection, &storage_key) + .undo(memory_undo), + e, + ); return Err(ErrorCode::from(e)); } }; + memory_undo.append(&mut outcome.memory_undo); if let Err(e) = chain .settle(self, surrogate, &outcome.stored_value) .and_then(|()| chain.persist_head(self, &txn)) { - abort_after_apply(self, &mut chain, database_id, tid, collection, &storage_key); + let e = abort_after_apply( + self, + &mut chain, + AbandonedWrite::row(database_id, tid, collection, &storage_key).undo(memory_undo), + e, + ); return Err(ErrorCode::from(e)); } @@ -139,24 +161,46 @@ impl CoreLoop { let enforcement = match write_hook::run(self, &txn, &hook_ctx, images) { Ok(enforcement) => enforcement, Err(e) => { - abort_after_apply(self, &mut chain, database_id, tid, collection, &storage_key); + let e = abort_after_apply( + self, + &mut chain, + AbandonedWrite::row(database_id, tid, collection, &storage_key) + .undo(memory_undo), + e, + ); return Err(ErrorCode::from(e)); } }; let target_write_set = write_hook::target_write_set(&enforcement.target_writes); + let target_writes = enforcement.target_writes; if let Err(e) = self.settle_balanced_entries(database_id, tid, collection, enforcement.balanced_entries) { - abort_after_apply(self, &mut chain, database_id, tid, collection, &storage_key); + let e = abort_after_apply( + self, + &mut chain, + AbandonedWrite::row(database_id, tid, collection, &storage_key) + .undo(memory_undo) + .targets(target_writes), + e, + ); return Err(ErrorCode::from(e)); } if let Err(e) = txn.commit() { - abort_after_apply(self, &mut chain, database_id, tid, collection, &storage_key); - return Err(ErrorCode::Internal { - detail: format!("commit: {e}"), - }); + let e = abort_after_apply( + self, + &mut chain, + AbandonedWrite::row(database_id, tid, collection, &storage_key) + .undo(memory_undo) + .targets(target_writes), + crate::Error::Storage { + engine: "sparse".into(), + detail: format!("commit: {e}"), + }, + ); + return Err(ErrorCode::from(e)); } self.checkpoint_coordinator.mark_dirty("sparse", 1); @@ -219,7 +263,7 @@ impl CoreLoop { let database_id = task.request.database_id.as_u64(); let txn = self.sparse.begin_write().map_err(ErrorCode::from)?; - let outcome = self + let mut outcome = self .apply_point_delete( &txn, PointDeleteParams { @@ -234,6 +278,11 @@ impl CoreLoop { }, ) .map_err(ErrorCode::from)?; + // Every abort below drops `txn` uncommitted, which reverses the + // durable writes only. `abandon_write` reverses the in-memory + // cascades and the target rows' cache and index entries. + let storage_key = StorageKey::for_surrogate(surrogate); + let memory_undo = std::mem::take(&mut outcome.memory_undo); let hook_ctx = HookCtx { database_id, @@ -246,25 +295,55 @@ impl CoreLoop { // The pre-image is the ONLY image a delete has, and it is what tells the // fold to take the removed row's contribution off the total. let enforcement = match outcome.prior_value { - Some(ref old) => write_hook::run( + Some(ref old) => match write_hook::run( self, &txn, &hook_ctx, WriteImages::Delete { old: ImageBody::Stored(old), }, - ) - .map_err(ErrorCode::from)?, + ) { + Ok(enforcement) => enforcement, + Err(e) => { + let e = abandon_write( + self, + AbandonedWrite::row(database_id, tid, collection, &storage_key) + .undo(memory_undo), + e, + ); + return Err(ErrorCode::from(e)); + } + }, None => Default::default(), }; let target_write_set = write_hook::target_write_set(&enforcement.target_writes); + let target_writes = enforcement.target_writes; - self.settle_balanced_entries(database_id, tid, collection, enforcement.balanced_entries) - .map_err(ErrorCode::from)?; + if let Err(e) = + self.settle_balanced_entries(database_id, tid, collection, enforcement.balanced_entries) + { + let e = abandon_write( + self, + AbandonedWrite::row(database_id, tid, collection, &storage_key) + .undo(memory_undo) + .targets(target_writes), + e, + ); + return Err(ErrorCode::from(e)); + } - txn.commit().map_err(|e| ErrorCode::Internal { - detail: format!("commit: {e}"), - })?; + if let Err(e) = txn.commit() { + let e = abandon_write( + self, + AbandonedWrite::row(database_id, tid, collection, &storage_key) + .undo(memory_undo) + .targets(target_writes), + crate::Error::DataPlane(ErrorCode::Internal { + detail: format!("commit: {e}"), + }), + ); + return Err(ErrorCode::from(e)); + } self.checkpoint_coordinator.mark_dirty("sparse", 1); if let Some(prior_bytes) = outcome.prior_value.as_deref() { @@ -280,13 +359,12 @@ impl CoreLoop { lsn, ); } - let old_converted = - self.resolve_event_payload(database_id, tid, collection, prior_bytes); self.emit_document_delete_event( task, + tid, collection, RowIdentity::from_user_key(document_id), - Some(old_converted.as_deref().unwrap_or(prior_bytes)), + Some(prior_bytes), ); } // The removal, journalled after apply, naming its own collection, diff --git a/nodedb/src/data/executor/handlers/document/write/batch_insert.rs b/nodedb/src/data/executor/handlers/document/write/batch_insert.rs index 72ecdc65c..873e56b93 100644 --- a/nodedb/src/data/executor/handlers/document/write/batch_insert.rs +++ b/nodedb/src/data/executor/handlers/document/write/batch_insert.rs @@ -7,10 +7,14 @@ use tracing::{debug, warn}; use crate::bridge::envelope::{ErrorCode, Response, WriteSetEntry}; use crate::data::executor::core_loop::CoreLoop; use crate::data::executor::core_loop::redo_image::submitted_row_image; -use crate::data::executor::enforcement::chain_guard::ChainGuard; +use crate::data::executor::enforcement::chain_guard::{ + AbandonedWrite, ChainGuard, abort_after_apply, +}; +use crate::data::executor::enforcement::materialized_sum::apply::TargetWrite; use crate::data::executor::enforcement::unique::SubmittedWrite; use crate::data::executor::enforcement::write_hook::{self, HookCtx, ImageBody, WriteImages}; use crate::data::executor::handlers::point::apply_put::PointPutParams; +use crate::data::executor::handlers::transaction::undo::UndoEntry; use crate::data::executor::task::ExecutionTask; use nodedb_physical::physical_plan::{ResolvedSumTarget, ReturningSpec}; @@ -190,11 +194,16 @@ impl CoreLoop { // judged together and a journal written as several rows of one // statement balances. let mut balanced_entries = Vec::new(); + // The page's in-memory side effects, in write order: every applied + // row's R-tree, vector and sparse entries, and the target rows its + // enforcement wrote. Dropping `txn` reverses the durable writes and + // none of these. + let mut page_undo: Vec = Vec::new(); + let mut page_targets: Vec = Vec::new(); // The row that failed, plus why. Collected rather than returned inline - // so the page's in-memory side effects — the advanced chain head and the - // document-cache entries `apply_point_put` populated — are reversed in - // one place. Dropping `txn` reverses the durable writes; it does not - // reverse either of those. + // so the page's in-memory side effects — these, the advanced chain + // head and the document-cache entries `apply_point_put` populated — + // are reversed in one place. let mut failure: Option<(nodedb_types::StorageKey, crate::Error)> = None; for (i, (document_id, value)) in documents.iter().enumerate() { let surrogate = surrogates[i]; @@ -212,7 +221,7 @@ impl CoreLoop { failure = Some((key, e)); break; } - let outcome = match self.apply_point_put( + let mut outcome = match self.apply_point_put( &txn, PointPutParams { database_id, @@ -238,6 +247,7 @@ impl CoreLoop { break; } }; + page_undo.append(&mut outcome.memory_undo); if let Err(e) = chain.settle(self, surrogate, &outcome.stored_value) { failure = Some((key, e)); break; @@ -262,6 +272,7 @@ impl CoreLoop { } }; target_write_set.extend(write_hook::target_write_set(&enforcement.target_writes)); + page_targets.extend(enforcement.target_writes); balanced_entries.extend(enforcement.balanced_entries); if returning.is_some() { stored_bodies.push(outcome.stored_value); @@ -281,19 +292,24 @@ impl CoreLoop { applied.push((row_identity, key)); } + // Every row the page called `apply_point_put` for. An abort below + // drops their cache entries: a cached body for a row that never + // committed is served to readers as though it had. + let mut written_keys: Vec = + applied.iter().map(|(_, key)| *key).collect(); + if let Some((failed_key, error)) = failure { - // The whole page rolls back, so put the chain head back where it - // started and drop every cache entry the abandoned rows populated — - // a cached body for a row that never committed is served to readers - // as though it had. - chain.restore(self); - for key in applied - .iter() - .map(|(_, key)| key) - .chain(std::iter::once(&failed_key)) - { - self.doc_cache.invalidate(database_id, tid, collection, key); - } + // The whole page rolls back: the chain head, the cache entries + // and every in-memory index entry of the page go back. + written_keys.push(failed_key); + let error = abort_after_apply( + self, + &mut chain, + AbandonedWrite::rows(database_id, tid, collection, &written_keys) + .undo(page_undo) + .targets(page_targets), + error, + ); return self.response_error(task, error); } @@ -302,34 +318,44 @@ impl CoreLoop { // no rows at all. if let Err(e) = self.settle_balanced_entries(database_id, tid, collection, balanced_entries) { - chain.restore(self); - for (_, key) in &applied { - self.doc_cache.invalidate(database_id, tid, collection, key); - } + let e = abort_after_apply( + self, + &mut chain, + AbandonedWrite::rows(database_id, tid, collection, &written_keys) + .undo(page_undo) + .targets(page_targets), + e, + ); return self.response_error(task, e); } // The advanced head lands in the SAME transaction as the rows whose // hashes it covers. if let Err(e) = chain.persist_head(self, &txn) { - chain.restore(self); - for (_, key) in &applied { - self.doc_cache.invalidate(database_id, tid, collection, key); - } + let e = abort_after_apply( + self, + &mut chain, + AbandonedWrite::rows(database_id, tid, collection, &written_keys) + .undo(page_undo) + .targets(page_targets), + e, + ); return self.response_error(task, e); } if let Err(e) = txn.commit() { - chain.restore(self); - for (_, key) in &applied { - self.doc_cache.invalidate(database_id, tid, collection, key); - } - return self.response_error( - task, - ErrorCode::Internal { + let e = abort_after_apply( + self, + &mut chain, + AbandonedWrite::rows(database_id, tid, collection, &written_keys) + .undo(page_undo) + .targets(page_targets), + crate::Error::Storage { + engine: "sparse".into(), detail: format!("batch insert commit: {e}"), }, ); + return self.response_error(task, e); } // Record each committed row's touched secondary-index values into the @@ -377,17 +403,20 @@ impl CoreLoop { .zip(stored_bodies.iter()) .map(|(identity, stored)| (identity, stored.as_slice())) .collect(); - self.stored_returning_response(task, spec, rls_filters, strict_schema.as_ref(), &rows) + let identity_column = self.identity_column(database_id, tid, collection); + self.stored_returning_response( + task, + spec, + rls_filters, + strict_schema.as_ref(), + &identity_column, + &rows, + ) } else { match crate::data::executor::response_codec::encode_count("inserted", documents.len()) { Ok(bytes) => self.response_with_payload(task, bytes), Err(e) => { - return self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ); + return self.response_error(task, ErrorCode::from(e)); } } }; @@ -570,6 +599,77 @@ mod tests { ); } + /// A geometry + vector row body, in the MessagePack every handler gets. + fn geo_vector_body(x: f64, embedding: &[f64]) -> Vec { + doc_format::encode_to_msgpack(&serde_json::json!({ + "loc": format!(r#"{{"type":"Point","coordinates":[{x},1.0]}}"#), + "embedding": embedding, + })) + } + + /// A batch whose last row is refused leaves no in-memory trace of the + /// rows before it. Dropping the transaction reverses the stored rows + /// only, so without the undo rows 1-2 keep answering spatial predicates + /// and vector search. Row 3 is refused in its vector step, after its own + /// R-tree entry landed, so it also checks the refused row's own undo. + #[test] + fn a_refused_last_row_leaves_no_spatial_or_vector_trace() { + let dir = tempfile::tempdir().unwrap(); + let (mut core, _tx, _rx) = make_core_with_dir(dir.path()); + let db = DatabaseId::DEFAULT; + core.vector_params.insert( + (db, TenantId::new(TID), COLL.to_string()), + crate::engine::vector::hnsw::HnswParams::default(), + ); + let documents = vec![ + ("g1".to_string(), geo_vector_body(1.0, &[1.0, 0.0, 0.0])), + ("g2".to_string(), geo_vector_body(2.0, &[0.0, 1.0, 0.0])), + // The index is three wide, so this row is refused. + ("g3".to_string(), geo_vector_body(3.0, &[1.0, 1.0])), + ]; + let surrogates = vec![Surrogate(31), Surrogate(32), Surrogate(33)]; + + let task = batch_task(&documents, &surrogates); + let resp = core.execute_document_batch_insert( + &task, + DocumentBatchInsertParams { + tid: TID, + collection: COLL, + documents: &documents, + surrogates: &surrogates, + returning: None, + rls_filters: &[], + resolved_sum_targets: &[], + deferred_sum_targets: &[], + }, + ); + + assert_eq!(resp.status, Status::Error, "row 3 must refuse the batch"); + for surrogate in &surrogates { + assert!(stored(&core, *surrogate).is_none(), "no row may be stored"); + } + let spatial_key = (db, TenantId::new(TID), COLL.to_string(), "loc".to_string()); + assert!( + core.spatial_indexes + .get(&spatial_key) + .is_none_or(|rtree| rtree.entries().is_empty()), + "the R-tree must hold no entry of the abandoned rows" + ); + assert!( + core.spatial_doc_map.is_empty(), + "the reverse spatial map must hold no entry of the abandoned rows" + ); + let vector_key = CoreLoop::vector_index_key(db.as_u64(), TID, COLL, "embedding"); + assert!( + !core.vector_collections.contains_key(&vector_key), + "row 1 created the vector index, so the abandoned page must remove it" + ); + assert!( + core.vector_doc_map.is_empty(), + "the vector reverse map must hold no entry of the abandoned rows" + ); + } + /// A batch carrying fewer surrogates than documents has no cross-engine /// identity for its rows, so every index would silently omit them. It is /// refused outright rather than stored-and-reported-successful. diff --git a/nodedb/src/data/executor/handlers/graph_edge_write/delete.rs b/nodedb/src/data/executor/handlers/graph_edge_write/delete.rs index e489e25dd..36330f675 100644 --- a/nodedb/src/data/executor/handlers/graph_edge_write/delete.rs +++ b/nodedb/src/data/executor/handlers/graph_edge_write/delete.rs @@ -25,9 +25,9 @@ impl CoreLoop { /// Edge delete with optional transactional compensation. /// - /// The `UndoEntry::EdgeWrite` is recorded once the tombstone is written, - /// never before: it names the tombstone version, and a rollback removes - /// exactly that version. + /// The `UndoEntry::EdgeWrite` is recorded once the tombstone and its CSR + /// change both stand, never before: it names the tombstone version, and a + /// rollback removes exactly that version. /// /// The RLS write policy is decided against that same pre-image and BEFORE /// the tombstone: the row a policy governs is the edge that exists now, and @@ -69,18 +69,19 @@ impl CoreLoop { // The pre-image is always read: the RLS write gate needs it for any // non-admit-all policy, and the response needs it to report a // truthful affected count. - let old_properties = self - .edge_store - .get_edge( - database_id, - TenantId::new(tid), - collection, - src_id, - label, - dst_id, - ) - .ok() - .flatten(); + // A pre-image read error refuses the delete: read as absent, it would + // admit the delete without the policy and report nothing removed. + let old_properties = match self.edge_store.get_edge( + database_id, + TenantId::new(tid), + collection, + src_id, + label, + dst_id, + ) { + Ok(properties) => properties, + Err(e) => return self.response_error(task, ErrorCode::from(e)), + }; let existed = old_properties.is_some(); if let Err(error) = crate::data::executor::handlers::rls_write_gate::admit_edge_properties( @@ -105,8 +106,9 @@ impl CoreLoop { label, dst_id, }; - // The CSR state the undo puts back, read only when an undo is kept. - let csr_prior = undo.is_some().then(|| self.capture_edge_csr(&target)); + // The CSR state a reversal puts back: a transaction rollback, or the + // reversal of a tombstone the CSR refuses below. + let csr_prior = self.capture_edge_csr(&target); use crate::engine::graph::edge_store::EdgeRef; match self.edge_store.soft_delete_edge_recorded( EdgeRef::new( @@ -122,14 +124,11 @@ impl CoreLoop { ) { Ok(tombstone) => { let current = tombstone.current.clone(); - // The tombstone is written whether or not the edge was live, - // so the undo that removes it is recorded either way. - if let (Some(undo), Some(csr)) = (undo, csr_prior) { - undo.push(UndoEntry::EdgeWrite(Box::new(target.undo(tombstone, csr)))); - } + let edge_undo = target.undo(tombstone, csr_prior); // The CSR follows what the edge resolves to: a tombstone a // TRUNCATE hides, or one below a newer version, leaves the - // edge as it was. + // edge as it was. A CSR refusal takes the tombstone back out, + // so the edge store and the CSR never disagree. if let Err(e) = self.mirror_edge_csr( database_id, tid, @@ -137,12 +136,13 @@ impl CoreLoop { collection, current.as_deref(), ) { - return self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ); + let code = self.reverse_edge_write(edge_undo, e); + return self.response_error(task, code); + } + // The tombstone is written whether or not the edge was live, + // so the undo that removes it is recorded either way. + if let Some(undo) = undo { + undo.push(UndoEntry::EdgeWrite(Box::new(edge_undo))); } self.checkpoint_coordinator.mark_dirty("sparse", 1); self.note_edge_write_lsn(task, tid, collection, src_id, label, dst_id); @@ -178,12 +178,7 @@ impl CoreLoop { ))]; response } - Err(e) => self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ), + Err(e) => self.response_error(task, ErrorCode::from(e)), } } } @@ -211,6 +206,12 @@ mod tests { zerompk::to_msgpack_vec(&vec![filter]).expect("encode policy filter") } + /// The stored property map `{"owner": owner}`, as plain MessagePack. + fn owner_properties(owner: &str) -> Vec { + nodedb_types::json_msgpack::json_to_msgpack(&serde_json::json!({ "owner": owner })) + .expect("encode properties") + } + /// The delete is decided against the edge's STORED property object before /// the tombstone is written, so a rejected delete leaves the edge in place. #[test] @@ -227,7 +228,7 @@ mod tests { src_id: "a", label: "KNOWS", dst_id: "b", - properties: br#"{"owner":"alice"}"#, + properties: &owner_properties("alice"), src_surrogate: Surrogate::new(1), dst_surrogate: Surrogate::new(2), }, @@ -285,7 +286,7 @@ mod tests { src_id: "a", label: "KNOWS", dst_id: "b", - properties: br#"{"owner":"alice"}"#, + properties: &owner_properties("alice"), src_surrogate: Surrogate::new(1), dst_surrogate: Surrogate::new(2), }, diff --git a/nodedb/src/data/executor/handlers/graph_edge_write/delete_batch.rs b/nodedb/src/data/executor/handlers/graph_edge_write/delete_batch.rs index 67ba228c8..3a18c67ae 100644 --- a/nodedb/src/data/executor/handlers/graph_edge_write/delete_batch.rs +++ b/nodedb/src/data/executor/handlers/graph_edge_write/delete_batch.rs @@ -4,9 +4,10 @@ use tracing::debug; -use crate::bridge::envelope::{EdgeImage, ErrorCode, Response, WriteSetEntry}; +use crate::bridge::envelope::{EdgeImage, Response, WriteSetEntry}; use crate::data::executor::core_loop::CoreLoop; -use crate::data::executor::handlers::partial_refusal::refusal_after_partial_apply; +use crate::data::executor::handlers::partial_refusal::refusal_after_rows; +use crate::data::executor::handlers::transaction::undo::edge_write::EdgeTarget; use crate::data::executor::task::ExecutionTask; use crate::types::TenantId; @@ -52,25 +53,37 @@ impl CoreLoop { // batch with the landed tombstones, which stay. let mut write_set: Vec = Vec::with_capacity(edges.len()); for edge in edges { - let existed = self - .edge_store - .get_edge( - database_id, - TenantId::new(tid), - edge.collection.as_str(), - &edge.src_id, - &edge.label, - &edge.dst_id, - ) - .ok() - .flatten() - .is_some(); - if existed { - removed += 1; - } + // A pre-image read error refuses the batch: read as absent, the + // count would leave out an edge the tombstone removes. + let existed = match self.edge_store.get_edge( + database_id, + TenantId::new(tid), + edge.collection.as_str(), + &edge.src_id, + &edge.label, + &edge.dst_id, + ) { + Ok(properties) => properties.is_some(), + Err(e) => { + let code = refusal_after_rows(write_set.len() as u64, e); + return self.refusal_with_landed_rows(task, code, write_set); + } + }; + let target = EdgeTarget { + database_id, + tid, + collection: edge.collection.as_str(), + src_id: &edge.src_id, + label: &edge.label, + dst_id: &edge.dst_id, + }; + let csr_prior = self.capture_edge_csr(&target); let stamp = match self.graph_write_stamp() { Ok(stamp) => stamp, - Err(e) => return self.refusal_with_landed_rows(task, e.into(), write_set), + Err(e) => { + let code = refusal_after_rows(write_set.len() as u64, e); + return self.refusal_with_landed_rows(task, code, write_set); + } }; let ord = stamp.system_from; use crate::engine::graph::edge_store::EdgeRef; @@ -91,12 +104,28 @@ impl CoreLoop { ) { Ok(tombstone) => tombstone, Err(e) => { - let code = ErrorCode::Internal { - detail: e.to_string(), - }; + let code = refusal_after_rows(write_set.len() as u64, e); return self.refusal_with_landed_rows(task, code, write_set); } }; + let current = tombstone.current.clone(); + // The CSR follows what the edge resolves to after the tombstone. + // A CSR refusal takes this tombstone back out; the edges before it + // stand and are journalled. + if let Err(e) = self.mirror_edge_csr( + database_id, + tid, + (&edge.src_id, &edge.label, &edge.dst_id), + edge.collection.as_str(), + current.as_deref(), + ) { + let reversed = self.reverse_edge_write(target.undo(tombstone, csr_prior), e); + let code = refusal_after_rows(write_set.len() as u64, reversed); + return self.refusal_with_landed_rows(task, code, write_set); + } + if existed { + removed += 1; + } write_set.push(WriteSetEntry::edge(EdgeImage::Delete( crate::wal::EdgeDeleteRedo { collection: edge.collection.to_string(), @@ -109,19 +138,6 @@ impl CoreLoop { applied: (stamp.applied != ord).then_some(stamp.applied), }, ))); - // The CSR follows what the edge resolves to after the tombstone. - if let Err(e) = self.mirror_edge_csr( - database_id, - tid, - (&edge.src_id, &edge.label, &edge.dst_id), - edge.collection.as_str(), - tombstone.current.as_deref(), - ) { - let code = refusal_after_partial_apply(ErrorCode::Internal { - detail: format!("edge CSR update: {e}"), - }); - return self.refusal_with_landed_rows(task, code, write_set); - } } if !edges.is_empty() { self.checkpoint_coordinator diff --git a/nodedb/src/data/executor/handlers/graph_edge_write/put.rs b/nodedb/src/data/executor/handlers/graph_edge_write/put.rs index f7b55e327..acc04cf07 100644 --- a/nodedb/src/data/executor/handlers/graph_edge_write/put.rs +++ b/nodedb/src/data/executor/handlers/graph_edge_write/put.rs @@ -25,11 +25,12 @@ impl CoreLoop { /// Edge upsert with optional transactional compensation. /// - /// When `undo` is `Some`, the `UndoEntry::EdgeWrite` is recorded after - /// the edge-store version is written and before the fallible CSR - /// mutation. It names the version the put added, so a rollback removes - /// exactly that version. An entry recorded before the store write would - /// name a version that does not exist. + /// When `undo` is `Some`, the `UndoEntry::EdgeWrite` is recorded once the + /// edge-store version and its CSR edge both stand. It names the version + /// the put added, so a rollback removes exactly that version. + /// + /// A CSR refusal after the version is stored takes the version back out, + /// so the edge store never holds an edge the CSR misses. /// /// A put writes a new edge-store version and makes the CSR edge live with /// the weight in `properties`, so a successful put always reports exactly @@ -105,8 +106,9 @@ impl CoreLoop { label, dst_id, }; - // The CSR state the undo puts back, read only when an undo is kept. - let csr_prior = undo.is_some().then(|| self.capture_edge_csr(&target)); + // The CSR state a reversal puts back: a transaction rollback, or the + // reversal of a version the CSR refuses below. + let csr_prior = self.capture_edge_csr(&target); use crate::engine::graph::edge_store::EdgeRef; match self.edge_store.put_edge_version_recorded( EdgeRef::new( @@ -126,11 +128,7 @@ impl CoreLoop { ) { Ok(version) => { let current = version.current.clone(); - // Edge-store version is now durable; the compensation entry is - // valid from here on even if the CSR mutation below fails. - if let (Some(undo), Some(csr)) = (undo, csr_prior) { - undo.push(UndoEntry::EdgeWrite(Box::new(target.undo(version, csr)))); - } + let edge_undo = target.undo(version, csr_prior); // The CSR follows what the edge resolves to: a version a // TRUNCATE hides, or one below a newer version, leaves the // edge as it was. @@ -143,6 +141,11 @@ impl CoreLoop { ); match csr_result { Ok(()) => { + // The version and its CSR edge both stand, so a + // transaction rollback reverses them together. + if let Some(undo) = undo { + undo.push(UndoEntry::EdgeWrite(Box::new(edge_undo))); + } let partition = self.csr_partition_mut(database_id, tid); // Populate the per-node surrogates so future bitmap-gated // traversals can check membership without a separate lookup. @@ -184,20 +187,15 @@ impl CoreLoop { ))]; response } - Err(e) => self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ), + // The CSR never misses a stored edge: the version the CSR + // refused is taken back out of the edge store. + Err(e) => { + let code = self.reverse_edge_write(edge_undo, e); + self.response_error(task, code) + } } } - Err(e) => self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ), + Err(e) => self.response_error(task, ErrorCode::from(e)), } } } @@ -254,6 +252,60 @@ mod tests { assert_eq!(event.new_value.as_deref(), Some(b"w=1".as_slice())); } + /// A CSR refusal after the edge store took the version takes the version + /// back out, and the answer keeps the CSR error's typed code. + #[test] + fn a_reversed_edge_write_leaves_no_version_and_keeps_the_typed_code() { + use crate::data::executor::handlers::transaction::undo::edge_write::EdgeTarget; + use crate::engine::graph::edge_store::{EdgeRef, VersionStamp}; + + let mut h = make_core(); + let task = make_task_with_lsn(79); + let database_id = task.request.database_id; + let target = EdgeTarget { + database_id: database_id.as_u64(), + tid: 1, + collection: "knows", + src_id: "a", + label: "KNOWS", + dst_id: "b", + }; + let csr_prior = h.core.capture_edge_csr(&target); + let version = h + .core + .edge_store + .put_edge_version_recorded( + EdgeRef::new(database_id, TenantId::new(1), "knows", "a", "KNOWS", "b") + .with_surrogates(Surrogate::new(1), Surrogate::new(2)), + b"w=1", + VersionStamp::at(100), + 100, + i64::MAX, + true, + ) + .expect("put edge version"); + + let code = h.core.reverse_edge_write( + target.undo(version, csr_prior), + nodedb_graph::GraphError::RebuildInProgress, + ); + + assert!(matches!(code, ErrorCode::BadRequest { .. }), "{code:?}"); + let stored = h + .core + .edge_store + .get_edge( + database_id.as_u64(), + TenantId::new(1), + "knows", + "a", + "KNOWS", + "b", + ) + .expect("edge lookup"); + assert!(stored.is_none(), "the reversed version is gone"); + } + /// An endpoint under `Surrogate::ZERO` names no node: the put is refused /// and no edge version is written. #[test] diff --git a/nodedb/src/data/executor/handlers/graph_edge_write/put_batch.rs b/nodedb/src/data/executor/handlers/graph_edge_write/put_batch.rs index 500bfe481..3e1fae0f4 100644 --- a/nodedb/src/data/executor/handlers/graph_edge_write/put_batch.rs +++ b/nodedb/src/data/executor/handlers/graph_edge_write/put_batch.rs @@ -6,7 +6,8 @@ use tracing::debug; use crate::bridge::envelope::{EdgeImage, ErrorCode, Response, WriteSetEntry}; use crate::data::executor::core_loop::CoreLoop; -use crate::data::executor::handlers::partial_refusal::refusal_after_partial_apply; +use crate::data::executor::handlers::partial_refusal::refusal_after_rows; +use crate::data::executor::handlers::transaction::undo::edge_write::EdgeTarget; use crate::data::executor::task::ExecutionTask; use crate::types::TenantId; @@ -60,10 +61,22 @@ impl CoreLoop { // apply. An edge that failed after earlier ones landed refuses the // batch with the landed versions, which stay. let mut write_set: Vec = Vec::with_capacity(edges.len()); - for (idx, edge) in edges.iter().enumerate() { + for edge in edges { + let target = EdgeTarget { + database_id, + tid, + collection: edge.collection.as_str(), + src_id: &edge.src_id, + label: &edge.label, + dst_id: &edge.dst_id, + }; + let csr_prior = self.capture_edge_csr(&target); let stamp = match self.graph_write_stamp() { Ok(stamp) => stamp, - Err(e) => return self.refusal_with_landed_rows(task, e.into(), write_set), + Err(e) => { + let code = refusal_after_rows(write_set.len() as u64, e); + return self.refusal_with_landed_rows(task, code, write_set); + } }; let ord = stamp.system_from; let valid_from_ms = nodedb_types::ordinal_to_ms(ord); @@ -85,7 +98,22 @@ impl CoreLoop { owns_logical_edge_stats(task, &edge.src_id), ) { Ok(version) => { - // The version is in the edge store from here on. + let current = version.current.clone(); + // The CSR follows what the edge resolves to: a version a + // TRUNCATE hides leaves the edge as it was. A CSR refusal + // takes this version back out; the edges before it stand + // and are journalled. + if let Err(e) = self.mirror_edge_csr( + database_id, + tid, + (&edge.src_id, &edge.label, &edge.dst_id), + edge.collection.as_str(), + current.as_deref(), + ) { + let reversed = self.reverse_edge_write(target.undo(version, csr_prior), e); + let code = refusal_after_rows(write_set.len() as u64, reversed); + return self.refusal_with_landed_rows(task, code, write_set); + } write_set.push(WriteSetEntry::edge(EdgeImage::Put( crate::wal::EdgePutRedo { collection: edge.collection.to_string(), @@ -99,28 +127,12 @@ impl CoreLoop { applied: (stamp.applied != ord).then_some(stamp.applied), }, ))); - // The CSR follows what the edge resolves to: a version a - // TRUNCATE hides leaves the edge as it was. - if let Err(e) = self.mirror_edge_csr( - database_id, - tid, - (&edge.src_id, &edge.label, &edge.dst_id), - edge.collection.as_str(), - version.current.as_deref(), - ) { - let code = refusal_after_partial_apply(ErrorCode::Internal { - detail: format!("edge {idx} (label interning): {e}"), - }); - return self.refusal_with_landed_rows(task, code, write_set); - } let partition = self.csr_partition_mut(database_id, tid); partition.set_node_surrogate(&edge.src_id, edge.src_surrogate); partition.set_node_surrogate(&edge.dst_id, edge.dst_surrogate); } Err(e) => { - let code = ErrorCode::Internal { - detail: format!("edge {idx}: {e}"), - }; + let code = refusal_after_rows(write_set.len() as u64, e); return self.refusal_with_landed_rows(task, code, write_set); } } diff --git a/nodedb/src/data/executor/handlers/graph_edge_write/shared.rs b/nodedb/src/data/executor/handlers/graph_edge_write/shared.rs index dbf2cc9f1..de22a0ca1 100644 --- a/nodedb/src/data/executor/handlers/graph_edge_write/shared.rs +++ b/nodedb/src/data/executor/handlers/graph_edge_write/shared.rs @@ -2,7 +2,9 @@ //! Shared param structs and helpers for the edge write handlers. +use crate::bridge::envelope::ErrorCode; use crate::data::executor::core_loop::CoreLoop; +use crate::data::executor::handlers::transaction::undo::edge_write::EdgeWriteUndo; use crate::data::executor::task::ExecutionTask; use crate::types::{RecordHomes, TenantId}; @@ -87,6 +89,29 @@ impl CoreLoop { .restore_edge_in_collection(src_id, label, dst_id, collection, weight) } + /// Take back an edge write the CSR refused. The edge store already holds + /// the write's version, so the version is removed and the CSR is put back + /// to the state the write found. The answer is the CSR error. A failed + /// reversal leaves the edge store and the CSR apart: the answer is then + /// `RollbackFailed`, which fail-stops the core. + pub(in crate::data::executor) fn reverse_edge_write( + &mut self, + undo: EdgeWriteUndo, + csr_error: nodedb_graph::GraphError, + ) -> ErrorCode { + let csr_detail = csr_error.to_string(); + match self.apply_undo_edge_write(0, undo) { + Ok(()) => crate::Error::from(csr_error).into(), + Err(mut undo_error) => { + undo_error.action = format!( + "{}; while reversing an edge write the CSR refused: {csr_detail}", + undo_error.action + ); + ErrorCode::from(undo_error) + } + } + } + /// Record a committed edge write's version, keyed by the edge's /// `(src, label, dst)` identity, if a WAL LSN was threaded onto the task. pub(in crate::data::executor) fn note_edge_write_lsn( @@ -167,7 +192,7 @@ pub(super) mod test_support { vshard_id: VShardId::new(0), plan: PhysicalPlan::Graph(GraphOp::Neighbors { node_id: "x".to_string(), - edge_label: None, + edge_labels: Vec::new(), direction: nodedb_graph::Direction::Out, rls_filters: Vec::new(), collection: None, diff --git a/nodedb/src/data/executor/handlers/merge_orchestrated/abort.rs b/nodedb/src/data/executor/handlers/merge_orchestrated/abort.rs index e12fad985..6ce8c74c9 100644 --- a/nodedb/src/data/executor/handlers/merge_orchestrated/abort.rs +++ b/nodedb/src/data/executor/handlers/merge_orchestrated/abort.rs @@ -73,10 +73,7 @@ impl CoreLoop { } = p; let final_err = match self.rollback_undo_log(database_id, tid, undo_log) { Ok(()) => err, - Err((entry_index, detail)) => ErrorCode::RollbackFailed { - entry_index, - detail, - }, + Err(undo_error) => ErrorCode::from(undo_error), }; self.rollback_merge_cache(database_id, tid, collection, applied_keys); self.response_error(task, final_err) diff --git a/nodedb/src/data/executor/handlers/merge_orchestrated/apply/update_rows.rs b/nodedb/src/data/executor/handlers/merge_orchestrated/apply/update_rows.rs index 3522a56d9..85cc3126b 100644 --- a/nodedb/src/data/executor/handlers/merge_orchestrated/apply/update_rows.rs +++ b/nodedb/src/data/executor/handlers/merge_orchestrated/apply/update_rows.rs @@ -79,6 +79,7 @@ impl CoreLoop { balanced_entries, returned_docs, } = tally; + let identity_column = self.identity_column(database_id, tid, collection); for upd in updates { let surrogate = upd.key.surrogate(); @@ -95,13 +96,7 @@ impl CoreLoop { if has_vectors { for d in self.remove_document_vector_indexes(database_id, tid, collection, upd.key) { - undo_log.push(UndoEntry::DeleteVector { - index_key: d.index_key, - vector_id: d.vector_id, - collection: d.collection, - field: d.field, - doc_id: Some(d.doc_id), - }); + undo_log.push(d.into_delete_undo()); } } match self.apply_point_put( @@ -170,7 +165,7 @@ impl CoreLoop { outcome.bitemporal_sys_from_ms, )); if returning { - match returning_doc(&upd.body, &upd.key) { + match returning_doc(&upd.body, &upd.key, &identity_column) { Ok(doc) => returned_docs.push(doc), Err(e) => { return Err(self.abort_merge_apply(MergeAbort { diff --git a/nodedb/src/data/executor/handlers/merge_orchestrated/apply_support.rs b/nodedb/src/data/executor/handlers/merge_orchestrated/apply_support.rs index 883d6c037..c40c8f52a 100644 --- a/nodedb/src/data/executor/handlers/merge_orchestrated/apply_support.rs +++ b/nodedb/src/data/executor/handlers/merge_orchestrated/apply_support.rs @@ -19,23 +19,12 @@ pub(super) type MergePutEvent<'a> = (RowIdentity, &'a [u8], Option>); /// Record the in-memory index mutations a successful /// [`crate::data::executor::core_loop::CoreLoop::apply_point_put`] performed as -/// undo entries. The HNSW vector index and the spatial R-tree live OUTSIDE the +/// undo entries. The vector, sparse and spatial indexes live OUTSIDE the /// shared redb transaction, so dropping that transaction on abort does not -/// reverse them — they must be undone explicitly. Drains the outcome's insert -/// deltas (leaving `prior_value` for the caller's event emission). +/// reverse them. Drains the outcome's undo +/// entries and leaves `prior_value` for the caller's event emission. pub(super) fn record_put_index_undo(undo_log: &mut Vec, outcome: &mut PointPutOutcome) { - for d in std::mem::take(&mut outcome.vector_inserts) { - undo_log.push(UndoEntry::InsertVector { - index_key: d.index_key, - vector_id: d.vector_id, - collection: d.collection, - field: d.field, - doc_id: Some(d.doc_id), - }); - } - for (key, entry_id) in std::mem::take(&mut outcome.spatial_inserts) { - undo_log.push(UndoEntry::SpatialInsert { key, entry_id }); - } + undo_log.append(&mut outcome.memory_undo); } /// Decide every resolved arm of a MERGE against the target's compiled write @@ -54,6 +43,7 @@ pub(super) fn record_put_index_undo(undo_log: &mut Vec, outcome: &mut pub(super) fn gate_merge_arms( plan: &MergePlanActions, rls_write_check: &nodedb_types::RlsWriteCheck, + identity_column: &str, tid: u64, collection: &str, ) -> crate::Result<()> { @@ -73,7 +63,15 @@ pub(super) fn gate_merge_arms( .chain(plan.deletes.iter().map(|d| (d.body.as_slice(), d.key))); for (body, key) in doc_arms { let identity = key.to_identity(); - rls_write_gate::admit_stored_row(rls_write_check, body, &identity, None, tid, collection)?; + rls_write_gate::admit_stored_row( + rls_write_check, + body, + &identity, + None, + identity_column, + tid, + collection, + )?; } for insert in &plan.inserts { let identity = @@ -83,6 +81,7 @@ pub(super) fn gate_merge_arms( &insert.body, &identity, None, + identity_column, tid, collection, )?; @@ -102,7 +101,12 @@ pub(super) fn gate_merge_arms( /// bodies are MessagePack for BOTH storage modes (`collect_merge_plan` decodes /// a strict target's Binary Tuple and re-encodes the resolved row before the /// apply pass ever sees it), so the strict decoder would have nothing to read. -pub(super) fn returning_doc(body: &[u8], key: &StorageKey) -> crate::Result { +/// `identity_column` is the column the identity renders under. +pub(super) fn returning_doc( + body: &[u8], + key: &StorageKey, + identity_column: &str, +) -> crate::Result { let identity = key.to_identity(); - super::super::returning_doc::from_stored(body, &identity, None) + super::super::returning_doc::from_stored(body, &identity, None, identity_column) } diff --git a/nodedb/src/data/executor/handlers/merge_orchestrated/delete_arms.rs b/nodedb/src/data/executor/handlers/merge_orchestrated/delete_arms.rs index fd511a8c6..9d43b4abf 100644 --- a/nodedb/src/data/executor/handlers/merge_orchestrated/delete_arms.rs +++ b/nodedb/src/data/executor/handlers/merge_orchestrated/delete_arms.rs @@ -15,6 +15,7 @@ use crate::bridge::envelope::{Response, WriteSetEntry}; use crate::data::executor::core_loop::CoreLoop; use crate::data::executor::core_loop::redo_image::removed_row_image; +use crate::data::executor::enforcement::chain_guard::{AbandonedWrite, abandon_write}; use crate::data::executor::enforcement::write_hook; use crate::data::executor::handlers::point::apply_delete::PointDeleteParams; use crate::data::executor::task::ExecutionTask; @@ -71,6 +72,7 @@ impl CoreLoop { write_set, returned_docs, } = tally; + let identity_column = self.identity_column(database_id, tid, collection); for del in deletes { let surrogate = del.key.surrogate(); @@ -98,7 +100,11 @@ impl CoreLoop { resolved_targets, }, ) { - Ok(outcome) => { + Ok(mut outcome) => { + // An abort below drops `txn` uncommitted, which reverses + // the durable writes only. `abandon_write` reverses the + // arm's in-memory cascades. + let memory_undo = std::mem::take(&mut outcome.memory_undo); // A DELETE arm takes the removed row's contribution // back off its target, folded inside THIS arm's // transaction so the debit and the removal commit @@ -129,16 +135,28 @@ impl CoreLoop { Ok(enforcement) => enforcement.target_writes, // Dropping `txn` un-committed reverses the // removal and every target it had debited. - Err(e) => return Err(self.response_error(task, e)), + Err(e) => { + let e = abandon_write( + self, + AbandonedWrite::row(database_id, tid, collection, &del.key) + .undo(memory_undo), + e, + ); + return Err(self.response_error(task, e)); + } }; if let Err(e) = txn.commit() { - return Err(self.response_error( - task, + let e = abandon_write( + self, + AbandonedWrite::row(database_id, tid, collection, &del.key) + .undo(memory_undo) + .targets(target_writes), crate::Error::Storage { engine: "sparse".into(), detail: format!("merge delete commit: {e}"), }, - )); + ); + return Err(self.response_error(task, e)); } // Journalled only once the arm committed: the removal // and the target rows its debit rewrote. @@ -159,7 +177,7 @@ impl CoreLoop { // is the raw stored form (Binary Tuple on a // strict target) and would need re-decoding. if returning { - match returning_doc(&del.body, &del.key) { + match returning_doc(&del.body, &del.key, &identity_column) { Ok(doc) => returned_docs.push(doc), Err(e) => return Err(self.response_error(task, e)), } @@ -167,6 +185,7 @@ impl CoreLoop { } self.emit_document_delete_event( task, + tid, collection, row_identity, outcome.prior_value.as_deref(), diff --git a/nodedb/src/data/executor/handlers/point/apply_delete.rs b/nodedb/src/data/executor/handlers/point/apply_delete.rs index ac85cdcf0..fd9860c42 100644 --- a/nodedb/src/data/executor/handlers/point/apply_delete.rs +++ b/nodedb/src/data/executor/handlers/point/apply_delete.rs @@ -19,7 +19,7 @@ use nodedb_types::Surrogate; use crate::data::executor::handlers::point::apply_put::SpatialEntryId; use crate::data::executor::handlers::point::apply_put::VectorIndexDelta; use crate::data::executor::handlers::point::apply_put::map_enforcement_error; -use crate::data::executor::spatial_key::SpatialIndexKey; +use crate::data::executor::handlers::transaction::undo::UndoEntry; /// Parameters for [`CoreLoop::apply_point_delete`]. pub(in crate::data::executor) struct PointDeleteParams<'a> { @@ -68,28 +68,18 @@ pub(in crate::data::executor) struct PointDeleteOutcome { /// where a rolled-back DELETE never restored its secondary-index entries. /// Empty on the bitemporal path (which has no plain INDEXES entries). pub secondary_index_tuples: Vec<(String, String)>, - /// Vector index mutations this delete soft-deleted from HNSW vector - /// indexes. Populated unconditionally (autocommit and transactional) so the - /// owning document's vectors never orphan; a transactional caller pushes an - /// `UndoEntry::DeleteVector` per entry so a rolled-back delete restores - /// them, including the paired `vector_doc_map` entry this cascade removed. - pub vector_deletes: Vec, - /// `(spatial_index_key, entry_id, bbox, document_id)` tuples this delete - /// removed from per-field spatial R-trees (and the reverse - /// `spatial_doc_map`). The bbox is captured BEFORE the R-tree `delete` - /// (which does not return it) so a transactional caller can push - /// `UndoEntry::SpatialDelete` re-insert reversals. Empty when the document - /// had no spatial fields. Autocommit callers ignore it (an aborted redb txn - /// does not reverse in-memory spatial writes). - pub spatial_deletes: Vec<(SpatialIndexKey, u64, nodedb_types::BoundingBox, String)>, - /// The node id this delete NEWLY marked deleted in the in-memory - /// `deleted_nodes` edge referential-integrity tracker, if any. `Some(id)` - /// only when `mark_node_deleted` newly inserted the node (it was not already - /// tombstoned by a prior committed op). A transactional caller pushes an - /// `UndoEntry::MarkNodeDeleted` so a rolled-back delete un-marks exactly the - /// node it added — never resurrecting a pre-existing tombstone. `None` when - /// the node was already marked. Autocommit callers ignore it. - pub mark_node_deleted: Option, + /// Undo entries for the in-memory cascades, in the order they ran: + /// - `SpatialDelete` for each R-tree entry and `spatial_doc_map` record + /// removed + /// - `MarkNodeDeleted` when this delete newly marked the row's node + /// deleted. A node a prior write marked is never un-marked. + /// - `DeleteVector` for each vector node soft-deleted and each + /// `vector_doc_map` entry removed + /// - `SparseDoc` for each sparse-vector document removed + /// + /// Dropping the caller's transaction does not reverse these. A caller + /// that abandons the delete reverses them with `undo_memory_effects`. + pub memory_undo: Vec, } impl CoreLoop { @@ -113,7 +103,8 @@ impl CoreLoop { /// the spatial / vector / sparse-vector removals are in-memory, so those /// cascades are unaffected by the caller's transaction. /// - /// On `Err` the caller MUST drop `txn` without committing. + /// On `Err` the caller MUST drop `txn` without committing. An `Err` + /// leaves no in-memory change behind. /// /// Does NOT emit WriteEvents, mark checkpoints dirty, or build /// RETURNING payloads — those stay with the caller. @@ -121,8 +112,9 @@ impl CoreLoop { /// Returns a [`PointDeleteOutcome`] capturing the prior stored bytes /// (present when a row was actually removed) plus the bitemporal system /// time and versioned index tombstone tuples written, so a transactional - /// caller can build a fully-reversible undo entry. Autocommit callers - /// read only `prior_value`. + /// caller can build a fully-reversible undo entry. A caller that drops + /// `txn` uncommitted after `Ok` reverses `memory_undo` with + /// `abandon_write`. pub(in crate::data::executor) fn apply_point_delete( &mut self, txn: &WriteTransaction, @@ -343,26 +335,23 @@ impl CoreLoop { // Cascade 3: Remove from spatial R-tree indexes + reverse map, and // record the node deletion for edge referential integrity. Both are - // fully captured (`spatial_deletes` + `mark_node_deleted` in the - // outcome) and reversed on rollback, so they run unconditionally for - // both the autocommit and transactional delete paths. + // captured in `memory_undo` and reversed on rollback or abort, so + // they run unconditionally for both the autocommit and transactional + // delete paths. // // `apply_point_put` hashes the substrate row key as the R-tree entry // id, so delete must hash the same key to find the entry. Hashing the // user PK would leak ghost bbox entries that survive the row's removal. - // `(spatial_index_key, entry_id, bbox, document_id)` tuples removed by - // the spatial cascade below are captured so a transactional caller can - // reverse them. - let mut mark_node_deleted_capture: Option = None; - // The put path hashes the same storage key via `SpatialEntryId`, so - // the shared removal hashes the same key to find and drop every - // per-field entry + reverse-map pair for this document. Captures - // each removed `(skey, entry_id, bbox, doc)` for reversible undo. - let spatial_deletes = self.remove_document_spatial_indexes( + // + // No step from here on fails, so an `Err` from this function leaves + // no in-memory change behind. + let mut memory_undo: Vec = Vec::new(); + self.remove_document_spatial_indexes_with_undo( database_id, tid, collection, SpatialEntryId::from_storage_key(storage_key), + &mut memory_undo, ); // Record deletion for edge referential integrity. Capture the id @@ -370,7 +359,12 @@ impl CoreLoop { // a prior committed op already tombstoned would wrongly resurrect // it as a valid edge target. if self.mark_node_deleted(database_id, tid, collection, document_id) { - mark_node_deleted_capture = Some(document_id.to_string()); + memory_undo.push(UndoEntry::MarkNodeDeleted { + database_id, + tid, + collection: collection.to_string(), + node_id: document_id.to_string(), + }); } // Cascade 4 (CORE, UNCONDITIONAL): soft-delete any HNSW vector entries @@ -385,14 +379,23 @@ impl CoreLoop { // enumeration the put path uses, so each `vector_doc_map` entry is // looked up by its exact key rather than scanning the whole map on // every delete. Shared with the PointUpdate re-index path. - let vector_deletes = - self.remove_document_vector_indexes(database_id, tid, collection, storage_key); + memory_undo.extend( + self.remove_document_vector_indexes(database_id, tid, collection, storage_key) + .into_iter() + .map(VectorIndexDelta::into_delete_undo), + ); // Sparse inverted-index cleanup, mirroring the dense-vector cascade // above: drop this document's sparse posting entries under the same hex // surrogate row key the put path indexed them by. A no-op unless the // strict schema declares a `SparseVector` column. - self.remove_document_sparse_indexes(database_id, tid, collection, storage_key); + self.remove_document_sparse_indexes( + database_id, + tid, + collection, + storage_key, + &mut memory_undo, + ); // Invalidate document cache. self.doc_cache @@ -409,9 +412,7 @@ impl CoreLoop { bitemporal_sys_from_ms, bitemporal_index_tuples, secondary_index_tuples, - vector_deletes, - spatial_deletes, - mark_node_deleted: mark_node_deleted_capture, + memory_undo, }) } } diff --git a/nodedb/src/data/executor/handlers/point/apply_put/core.rs b/nodedb/src/data/executor/handlers/point/apply_put/core.rs index 848f1dd10..8802fdb63 100644 --- a/nodedb/src/data/executor/handlers/point/apply_put/core.rs +++ b/nodedb/src/data/executor/handlers/point/apply_put/core.rs @@ -14,6 +14,8 @@ use crate::data::executor::doc_format; use crate::data::executor::enforcement::unique::{ PostImage, UniqueJudge, UniqueScope, check_unique_post_state, }; +use crate::data::executor::handlers::transaction::undo::UndoEntry; +use crate::data::executor::handlers::transaction::undo::memory::abort_error; use super::enforce::PutEnforcement; use super::types::{PointPutOutcome, PointPutParams}; @@ -24,6 +26,12 @@ impl CoreLoop { /// cache. Does NOT commit — on `Err` the caller MUST drop `txn` /// uncommitted, or a row publishes with indexes nothing re-derives. /// + /// The R-tree, vector and sparse writes are in memory, so dropping `txn` + /// does not reverse them. On `Err` this function has + /// already reversed its own. On `Ok` they come back as + /// `PointPutOutcome::memory_undo`: a caller that abandons the write after + /// this returns reverses them with `undo_memory_effects`. + /// /// `value` always arrives WITH the write (never a row read back to /// reconcile), so a `decode_document(value)` guard below that fails /// quietly is an intentional "no fields to derive from", not a swallowed @@ -172,7 +180,7 @@ impl CoreLoop { // inverted-index write failure is real and rejects the write. if let Ok(doc) = doc_format::decode_document(value) { // Shared with the DELETE-rollback re-index path. - let text_content = crate::data::executor::fts_text::extract_fts_text(&doc); + let text_content = crate::data::executor::fts_text::extract_fts_fields(&doc); // Empty text is NOT skipped — stripping every indexable word // must still remove the document from the index. if index_text { @@ -311,10 +319,12 @@ impl CoreLoop { } } - let spatial_inserts = - self.apply_point_put_spatial(database_id, tid, collection, storage_key, value); - let vector_inserts = self.apply_point_put_vector_indexes( - crate::data::executor::handlers::point::apply_put::VectorIndexPutParams { + // Every write below is in memory, where the caller's transaction + // cannot reach it. A failure part-way reverses what already ran, so + // an `Err` from this function leaves no in-memory index entry behind. + let mut memory_undo: Vec = Vec::new(); + if let Err(e) = self.apply_point_put_memory_indexes( + MemoryIndexPut { database_id, tid, collection, @@ -322,9 +332,11 @@ impl CoreLoop { value, wal_lsn: wal_lsn.map(|l| l.as_u64()).unwrap_or(0), }, - )?; - // No-op unless the strict schema declares a `SparseVector` column. - self.apply_point_put_sparse_indexes(database_id, tid, collection, storage_key, value); + &mut memory_undo, + ) { + let undone = self.undo_memory_effects(database_id, tid, memory_undo); + return Err(abort_error(e, undone)); + } Ok(PointPutOutcome { prior_value: prior, @@ -333,11 +345,54 @@ impl CoreLoop { bitemporal_index_tuples, secondary_index_added, secondary_index_removed, - vector_inserts, - spatial_inserts, + memory_undo, stats_prior, }) } + + /// The in-memory index writes of one put: the R-tree, the vector + /// indexes, and the sparse-vector indexes. Each step pushes an undo entry + /// onto `undo` for every mutation it makes, also when it fails part-way. + fn apply_point_put_memory_indexes( + &mut self, + put: MemoryIndexPut<'_>, + undo: &mut Vec, + ) -> crate::Result<()> { + let MemoryIndexPut { + database_id, + tid, + collection, + storage_key, + value, + wal_lsn, + } = put; + self.apply_point_put_spatial(database_id, tid, collection, storage_key, value, undo); + self.apply_point_put_vector_indexes( + crate::data::executor::handlers::point::apply_put::VectorIndexPutParams { + database_id, + tid, + collection, + storage_key, + value, + wal_lsn, + }, + undo, + )?; + // No-op unless the strict schema declares a `SparseVector` column. + self.apply_point_put_sparse_indexes(database_id, tid, collection, storage_key, value, undo); + Ok(()) + } +} + +/// Inputs to `apply_point_put_memory_indexes`. `wal_lsn` is `0` when the +/// write carries no LSN. +struct MemoryIndexPut<'a> { + database_id: u64, + tid: u64, + collection: &'a str, + storage_key: crate::engine::document::store::StorageKey, + value: &'a [u8], + wal_lsn: u64, } #[cfg(test)] @@ -362,7 +417,7 @@ mod tests { const COLL: &str = "articles"; const SURROGATE: Surrogate = Surrogate(7); /// Raw JSON body — `doc_format::decode_document`'s JSON fallback accepts - /// it, and its single string field is what `extract_fts_text` feeds the + /// it, and its single string field is what `extract_fts_fields` feeds the /// inverted index, so this document has real text to index. const BODY: &[u8] = br#"{"title":"alpha bravo charlie"}"#; diff --git a/nodedb/src/data/executor/handlers/point/apply_put/index.rs b/nodedb/src/data/executor/handlers/point/apply_put/index.rs index f3f7b3b94..0521c89c7 100644 --- a/nodedb/src/data/executor/handlers/point/apply_put/index.rs +++ b/nodedb/src/data/executor/handlers/point/apply_put/index.rs @@ -1,13 +1,17 @@ // SPDX-License-Identifier: BUSL-1.1 -//! Spatial R-tree + columnar ingest side-effect for `apply_point_put`: -//! geometry-field detection, per-field R-tree insert, reverse entry→doc map, -//! and columnar-memtable ingest. HNSW vector indexing lives in the sibling -//! `vector` module. Split out of `apply_put.rs` to keep that file focused on -//! the core document-write transaction. +//! Spatial R-tree side-effect for `apply_point_put`: geometry-field +//! detection, per-field R-tree insert, and the reverse entry→doc map. HNSW +//! vector indexing lives in the sibling `vector` module. +//! +//! A document collection keeps its rows in the sparse store only. Its +//! scans, aggregates and joins read them there, so every document is visible +//! and every delete is honored. The R-tree is the one derived structure a +//! geometry field adds. use crate::data::executor::core_loop::CoreLoop; use crate::data::executor::doc_format; +use crate::data::executor::handlers::transaction::undo::UndoEntry; use crate::data::executor::spatial_key::SpatialIndexKey; use crate::engine::document::store::StorageKey; @@ -43,15 +47,14 @@ impl SpatialEntryId { } impl CoreLoop { - /// Spatial R-tree + columnar ingest side-effect: parse geometry fields, - /// insert into the per-field R-tree, maintain the reverse entry→doc map, - /// and (when geometry present) ingest into the columnar memtable so bare - /// scans/aggregates over spatial collections work. + /// Spatial R-tree side-effect: parse geometry fields, insert into the + /// per-field R-tree, and maintain the reverse entry→doc map. /// - /// Returns the `(spatial_index_key, entry_id)` pairs inserted so a - /// transactional caller can push `UndoEntry::SpatialInsert` reversals. The - /// spatial writes are in-memory (an aborted redb txn does not reverse them), - /// so explicit undo is required. Empty when no geometry fields are present. + /// The spatial writes are in memory, so an aborted redb txn does not + /// reverse them. Pushes onto `undo` a `SpatialDelete` for each prior entry + /// of the document it removes, then a `SpatialInsert` for each entry it + /// inserts. Rollback runs the log in reverse, so it removes the new + /// entries before it puts the prior ones back. pub(in crate::data::executor) fn apply_point_put_spatial( &mut self, database_id: u64, @@ -59,54 +62,33 @@ impl CoreLoop { collection: &str, storage_key: StorageKey, value: &[u8], - ) -> Vec<( - ( - nodedb_types::DatabaseId, - crate::types::TenantId, - String, - String, - ), - u64, - )> { + undo: &mut Vec, + ) { // Rendered once here; every reverse-map / hash use below shares it. let document_id = storage_key.to_string(); let document_id = document_id.as_str(); let spatial_entry_id = SpatialEntryId::from_rendered(document_id); - let mut inserts = Vec::new(); - // Re-indexing a document must REPLACE, not append: `RTree::insert` - // blindly pushes a fresh entry even when one with this `entry_id` - // already exists, so a live geometry UPDATE, a WAL replay, or the - // crash-recovery rebuild would otherwise leave stale duplicate bbox - // entries scoring alongside the new one. Clear any prior geometry for - // this document first (idempotent — a no-op on a genuine first insert). - // The removed tuples are discarded here, mirroring the vector put path: - // only the new inserts are captured for transactional undo. - let _ = - self.remove_document_spatial_indexes(database_id, tid, collection, spatial_entry_id); - // Spatial index: detect geometry fields and insert into R-tree. - // Tries to parse each field as a GeoJSON Geometry — either a native - // JSON object (schemaless document writes, e.g. - // `{"type":"Point","coordinates":[...]}`) or a JSON string containing - // GeoJSON (SQL `ST_Point(...)` inserts, which serialize geometry to a - // string). See `nodedb_types::geometry::from_geojson_str` — shared - // with the read path (`extract_geometry` in spatial.rs) and the - // columnar index path (`geometry_index.rs`); keep all three in sync. - // If successful, computes bbox and inserts into the per-field R-tree. - // Also writes the document to columnar_memtables so that bare table scans - // and aggregates on spatial collections read from columnar (spatial extends columnar). + // Spatial index: detect geometry fields. Tries to parse each field as + // a GeoJSON Geometry — either a native JSON object (schemaless + // document writes, e.g. `{"type":"Point","coordinates":[...]}`) or a + // JSON string containing GeoJSON (SQL `ST_Point(...)` inserts, which + // serialize geometry to a string). See + // `nodedb_types::geometry::from_geojson_str` — shared with the read + // path (`extract_geometry` in spatial.rs) and the columnar index path + // (`geometry_index.rs`); keep all three in sync. // // `value` is `apply_point_put`'s incoming body, and the invariant on // that function applies unchanged here: geometry is detected by walking // a decoded document's fields, so a body that is not one carries no - // geometry to index and nothing to ingest. This would be wrong if + // geometry to index. This would be wrong if // `value` were ever the STORED row instead — a stored geometry that // failed to decode would silently drop out of the R-tree while the row - // stayed queryable, which is the desync the delete-then-insert above + // stayed queryable, which is the desync the delete-then-insert below // exists to prevent. + let mut geometries = Vec::new(); if let Ok(doc) = doc_format::decode_document(value) && let Some(obj) = doc.as_object() { - let mut has_geometry = false; for (field_name, field_value) in obj { let parsed_geom = match field_value { serde_json::Value::String(s) => nodedb_types::geometry::from_geojson_str(s), @@ -116,46 +98,72 @@ impl CoreLoop { .ok(), }; if let Some(geom) = parsed_geom { - has_geometry = true; - let bbox = nodedb_types::bbox::geometry_bbox(&geom); - let db_id = nodedb_types::DatabaseId::new(database_id); - let tid_id = crate::types::TenantId::new(tid); - let spatial_key = (db_id, tid_id, collection.to_string(), field_name.clone()); - let entry_id = spatial_entry_id.as_u64(); - let memory = nodedb_mem::ScopedMemory::new( - self.governor.clone(), - db_id, - tid_id, - nodedb_mem::EngineId::Spatial, - ); - let rtree = self - .spatial_indexes - .entry(spatial_key.clone()) - .or_insert_with(|| crate::engine::spatial::RTree::new(memory)); - rtree.insert(crate::engine::spatial::RTreeEntry { id: entry_id, bbox }); - // Maintain reverse map: entry_id → document_id. - self.spatial_doc_map.insert( - ( - db_id, - tid_id, - collection.to_string(), - field_name.clone(), - entry_id, - ), - document_id.to_string(), - ); - inserts.push((spatial_key, entry_id)); + geometries.push((field_name.clone(), geom)); } } + } - // If document has geometry, also write to columnar memtable. - // This ensures bare scans + aggregates work via columnar path. - if has_geometry { - self.ingest_doc_to_columnar(database_id, tid, collection, obj); - } + // Re-indexing a document must REPLACE, not append: `RTree::insert` + // blindly pushes a fresh entry even when one with this `entry_id` + // already exists, so a live geometry UPDATE, a WAL replay, or the + // crash-recovery rebuild would otherwise leave stale duplicate bbox + // entries scoring alongside the new one. Clear any prior geometry for + // this document first (idempotent — a no-op on a genuine first insert). + self.remove_document_spatial_indexes_with_undo( + database_id, + tid, + collection, + spatial_entry_id, + undo, + ); + let db_id = nodedb_types::DatabaseId::new(database_id); + let tid_id = crate::types::TenantId::new(tid); + let entry_id = spatial_entry_id.as_u64(); + for (field_name, geom) in geometries { + let bbox = nodedb_types::bbox::geometry_bbox(&geom); + let spatial_key = (db_id, tid_id, collection.to_string(), field_name.clone()); + let memory = nodedb_mem::ScopedMemory::new( + self.governor.clone(), + db_id, + tid_id, + nodedb_mem::EngineId::Spatial, + ); + let rtree = self + .spatial_indexes + .entry(spatial_key.clone()) + .or_insert_with(|| crate::engine::spatial::RTree::new(memory)); + rtree.insert(crate::engine::spatial::RTreeEntry { id: entry_id, bbox }); + // Maintain reverse map: entry_id → document_id. + self.spatial_doc_map.insert( + (db_id, tid_id, collection.to_string(), field_name, entry_id), + document_id.to_string(), + ); + undo.push(UndoEntry::SpatialInsert { + key: spatial_key, + entry_id, + }); } + } - inserts + /// [`Self::remove_document_spatial_indexes`], pushing a `SpatialDelete` + /// undo entry onto `undo` for each entry it removes. + pub(in crate::data::executor) fn remove_document_spatial_indexes_with_undo( + &mut self, + database_id: u64, + tid: u64, + collection: &str, + entry_id: SpatialEntryId, + undo: &mut Vec, + ) { + let removed = self.remove_document_spatial_indexes(database_id, tid, collection, entry_id); + for (key, entry_id, bbox, document_id) in removed { + undo.push(UndoEntry::SpatialDelete { + key, + entry_id, + bbox, + document_id, + }); + } } /// Remove every R-tree entry (and its paired `spatial_doc_map` reverse @@ -346,4 +354,172 @@ mod tests { assert!(core.spatial_indexes.contains_key(&key)); assert_eq!(core.spatial_indexes.get(&key).unwrap().entries().len(), 1); } + + /// A schemaless collection whose geometry documents carry different value + /// types for the same field, or no `id` field, accepts every write. A + /// collection scan returns every document, including one with no + /// geometry, from the sparse store. + #[test] + fn schemaless_geometry_docs_with_mixed_field_types_are_all_scanned() { + let dir = tempfile::tempdir().unwrap(); + let (mut core, _tx, _rx) = make_core_with_dir(dir.path()); + let coll = "geo_mixed"; + let docs: [(&str, u32, &[u8]); 4] = [ + ( + "a", + 1, + br#"{"id":"a","loc":{"type":"Point","coordinates":[1.0,2.0]},"n":1}"#, + ), + ( + "b", + 2, + br#"{"id":"b","loc":{"type":"Point","coordinates":[3.0,4.0]},"n":1.5}"#, + ), + ( + "c", + 3, + br#"{"loc":{"type":"Point","coordinates":[5.0,6.0]},"n":"text"}"#, + ), + ("d", 4, br#"{"id":"d","n":2}"#), + ]; + for (document_id, surrogate, doc) in docs { + let task = point_put_task(coll, document_id, doc); + let resp = core.execute_point_put( + &task, + PointPutExec { + tid: 1, + collection: coll, + document_id, + surrogate: Surrogate::new(surrogate), + value: doc, + returning: None, + rls_filters: &[], + resolved_sum_targets: &[], + }, + ); + assert_eq!(resp.status, Status::Ok, "put of '{document_id}' refused"); + } + + let key = ( + DatabaseId::DEFAULT, + TenantId::new(1), + coll.to_string(), + "loc".to_string(), + ); + assert_eq!(core.spatial_indexes.get(&key).unwrap().entries().len(), 3); + assert!( + !core.columnar_engines.contains_key(&( + DatabaseId::DEFAULT, + TenantId::new(1), + coll.to_string() + )), + "a document collection keeps no columnar copy of its rows" + ); + + let scanned = core + .scan_collection(DatabaseId::DEFAULT.as_u64(), 1, coll, usize::MAX) + .unwrap(); + assert_eq!(scanned.len(), 4, "every document must be scanned"); + } + + /// A put that overwrites a geometry row and then aborts puts the row's + /// prior R-tree entry and its `spatial_doc_map` record back. + #[test] + fn an_aborted_geometry_overwrite_restores_the_prior_entry() { + use crate::data::executor::enforcement::chain_guard::{AbandonedWrite, abandon_write}; + use crate::data::executor::enforcement::unique::UniqueJudge; + use crate::data::executor::handlers::point::apply_put::{PointPutParams, SpatialEntryId}; + use crate::engine::document::store::StorageKey; + + let dir = tempfile::tempdir().unwrap(); + let (mut core, _tx, _rx) = make_core_with_dir(dir.path()); + let coll = "geo_abort"; + let db = DatabaseId::DEFAULT.as_u64(); + let surrogate = Surrogate::new(1); + let storage_key = StorageKey::for_surrogate(surrogate); + let entry_id = SpatialEntryId::from_storage_key(storage_key).as_u64(); + let key = ( + DatabaseId::DEFAULT, + TenantId::new(1), + coll.to_string(), + "loc".to_string(), + ); + let map_key = ( + DatabaseId::DEFAULT, + TenantId::new(1), + coll.to_string(), + "loc".to_string(), + entry_id, + ); + + let old = br#"{"loc":{"type":"Point","coordinates":[1.0,2.0]}}"#; + let task = point_put_task(coll, "d1", old); + let resp = core.execute_point_put( + &task, + PointPutExec { + tid: 1, + collection: coll, + document_id: "d1", + surrogate, + value: old, + returning: None, + rls_filters: &[], + resolved_sum_targets: &[], + }, + ); + assert_eq!(resp.status, Status::Ok); + let old_entries = core.spatial_indexes.get(&key).unwrap().entries(); + assert_eq!(old_entries.len(), 1); + let old_bbox = old_entries[0].bbox; + let old_doc = core.spatial_doc_map.get(&map_key).cloned(); + assert!(old_doc.is_some()); + + let new = br#"{"loc":{"type":"Point","coordinates":[30.0,40.0]}}"#; + let txn = core.sparse.begin_write().unwrap(); + let outcome = core + .apply_point_put( + &txn, + PointPutParams { + database_id: db, + tid: 1, + collection: coll, + storage_key, + surrogate, + value: new, + index_text: true, + user_roles: &[], + enforce: true, + unique: UniqueJudge::Row, + wal_lsn: None, + resolved_targets: &[], + }, + ) + .unwrap(); + let new_entries = core.spatial_indexes.get(&key).unwrap().entries(); + assert_eq!(new_entries.len(), 1); + assert_ne!( + new_entries[0].bbox, old_bbox, + "the overwrite replaced the entry" + ); + + drop(txn); + let error = abandon_write( + &mut core, + AbandonedWrite::row(db, 1, coll, &storage_key).undo(outcome.memory_undo), + crate::Error::Storage { + engine: "sparse".into(), + detail: "commit refused".into(), + }, + ); + assert!( + matches!(error, crate::Error::Storage { .. }), + "a clean undo reports the original error, got {error:?}" + ); + + let restored = core.spatial_indexes.get(&key).unwrap().entries(); + assert_eq!(restored.len(), 1, "exactly the prior entry is back"); + assert_eq!(restored[0].id, entry_id); + assert_eq!(restored[0].bbox, old_bbox); + assert_eq!(core.spatial_doc_map.get(&map_key).cloned(), old_doc); + } } diff --git a/nodedb/src/data/executor/handlers/point/apply_put/sparse.rs b/nodedb/src/data/executor/handlers/point/apply_put/sparse.rs index d20576690..a5dc1bb33 100644 --- a/nodedb/src/data/executor/handlers/point/apply_put/sparse.rs +++ b/nodedb/src/data/executor/handlers/point/apply_put/sparse.rs @@ -10,6 +10,7 @@ //! the dense-vector path needs. use crate::data::executor::core_loop::CoreLoop; +use crate::data::executor::handlers::transaction::undo::UndoEntry; use crate::engine::document::store::StorageKey; impl CoreLoop { @@ -91,6 +92,10 @@ impl CoreLoop { /// /// No-op (byte-identical to a collection without sparse columns) when /// `strict_sparse_fields` is empty. + /// + /// Pushes a `SparseDoc` undo entry onto `undo` before each upsert. The + /// entry holds the document's prior postings in that index and the + /// index's id counter, or marks the index as created by this write. pub(in crate::data::executor) fn apply_point_put_sparse_indexes( &mut self, database_id: u64, @@ -98,6 +103,7 @@ impl CoreLoop { collection: &str, storage_key: StorageKey, value: &[u8], + undo: &mut Vec, ) { let sparse_fields = self.strict_sparse_fields(database_id, tid, collection); if sparse_fields.is_empty() { @@ -119,6 +125,20 @@ impl CoreLoop { let Ok(sv) = nodedb_types::SparseVector::parse_literal(literal) else { continue; }; + let key = Self::sparse_index_key(database_id, tid, collection, field); + let (prior, next_id) = match self.sparse_vector_indexes.get(&key) { + Some(index) => ( + index.doc_image(&document_id), + Some(index.next_internal_id()), + ), + None => (None, None), + }; + undo.push(UndoEntry::SparseDoc { + key, + doc_id: document_id.clone(), + prior, + next_id, + }); self.get_or_create_sparse_index(database_id, tid, collection, field) .insert(&document_id, &sv); // Sparse indexes are in-memory with no redb store behind them; the @@ -133,12 +153,16 @@ impl CoreLoop { /// and the PointUpdate re-index (which clears the old literal before /// inserting the new one). Mirrors `remove_document_vector_indexes`. /// No-op when the collection declares no sparse columns. + /// + /// Pushes a `SparseDoc` undo entry onto `undo` for each document entry it + /// removes, carrying the entry's internal id and postings. pub(in crate::data::executor) fn remove_document_sparse_indexes( &mut self, database_id: u64, tid: u64, collection: &str, storage_key: StorageKey, + undo: &mut Vec, ) { let sparse_fields = self.strict_sparse_fields(database_id, tid, collection); if sparse_fields.is_empty() { @@ -147,12 +171,23 @@ impl CoreLoop { // The sparse index keys postings by the rendered storage key. let row_key = storage_key.to_string(); for field in &sparse_fields { - if self - .get_or_create_sparse_index(database_id, tid, collection, field) - .delete(&row_key) - { - self.checkpoint_coordinator.mark_dirty("vector", 1); - } + let key = Self::sparse_index_key(database_id, tid, collection, field); + // A field with no index holds no entry to remove. + let Some(index) = self.sparse_vector_indexes.get_mut(&key) else { + continue; + }; + let Some(prior) = index.doc_image(&row_key) else { + continue; + }; + let next_id = index.next_internal_id(); + index.delete(&row_key); + undo.push(UndoEntry::SparseDoc { + key, + doc_id: row_key.clone(), + prior: Some(prior), + next_id: Some(next_id), + }); + self.checkpoint_coordinator.mark_dirty("vector", 1); } } } @@ -255,6 +290,7 @@ mod tests { collection, StorageKey::for_surrogate(Surrogate::new(1)), &doc, + &mut Vec::new(), ); assert_eq!( @@ -284,6 +320,7 @@ mod tests { collection, StorageKey::for_surrogate(Surrogate::new(1)), &doc_with_sparse(field, "{3:0.5, 7:1.5}"), + &mut Vec::new(), ); core.apply_point_put_sparse_indexes( db, @@ -291,6 +328,7 @@ mod tests { collection, StorageKey::for_surrogate(Surrogate::new(1)), &doc_with_sparse(field, "{1:0.9}"), + &mut Vec::new(), ); assert_eq!( @@ -320,20 +358,63 @@ mod tests { collection, StorageKey::for_surrogate(Surrogate::new(1)), &doc_with_sparse(field, "{3:0.5, 7:1.5}"), + &mut Vec::new(), ); assert_eq!(doc_count(core, db, tid, collection, field), 1); + let mut undo = Vec::new(); core.remove_document_sparse_indexes( db, tid, collection, StorageKey::for_surrogate(Surrogate::new(1)), + &mut undo, ); assert_eq!( doc_count(core, db, tid, collection, field), 0, "remove must drop the document from the sparse field's index" ); + assert_eq!(undo.len(), 1, "the removal pushes one undo entry"); + } + + /// Reversing a removal's undo entries puts the document's postings back + /// under the same internal id. + #[test] + fn undoing_a_remove_restores_the_sparse_entry() { + let mut harness = make_core(); + let core = &mut harness.core; + let (db, tid, collection, field) = (0u64, 1u64, "docs", "terms"); + register_strict_sparse(core, tid, collection, field); + let row = StorageKey::for_surrogate(Surrogate::new(1)); + + core.apply_point_put_sparse_indexes( + db, + tid, + collection, + row, + &doc_with_sparse(field, "{3:0.5, 7:1.5}"), + &mut Vec::new(), + ); + let key = CoreLoop::sparse_index_key(db, tid, collection, field); + let before = core + .sparse_vector_indexes + .get(&key) + .and_then(|index| index.doc_image(&row.to_string())) + .expect("the put indexed the row"); + + let mut undo = Vec::new(); + core.remove_document_sparse_indexes(db, tid, collection, row, &mut undo); + assert_eq!(doc_count(core, db, tid, collection, field), 0); + + core.undo_memory_effects(db, tid, undo) + .expect("the removal reverses"); + let after = core + .sparse_vector_indexes + .get(&key) + .and_then(|index| index.doc_image(&row.to_string())) + .expect("the row is indexed again"); + assert_eq!(after, before, "same internal id and postings"); } /// A collection with no `SparseVector` column is untouched: no index is @@ -360,6 +441,7 @@ mod tests { collection, StorageKey::for_surrogate(Surrogate::new(1)), &doc_with_sparse("terms", "{3:0.5}"), + &mut Vec::new(), ); assert!( diff --git a/nodedb/src/data/executor/handlers/point/apply_put/types.rs b/nodedb/src/data/executor/handlers/point/apply_put/types.rs index 247d3d2cc..866edc4ab 100644 --- a/nodedb/src/data/executor/handlers/point/apply_put/types.rs +++ b/nodedb/src/data/executor/handlers/point/apply_put/types.rs @@ -6,7 +6,7 @@ use nodedb_types::Surrogate; use crate::bridge::envelope::ErrorCode; -use crate::data::executor::spatial_key::SpatialIndexKey; +use crate::data::executor::handlers::transaction::undo::UndoEntry; use crate::engine::document::store::StorageKey; use nodedb_physical::physical_plan::ResolvedSumTarget; @@ -87,16 +87,13 @@ pub(in crate::data::executor) struct PointPutOutcome { /// bitemporal path. A transactional caller re-inserts these on rollback. /// Autocommit callers ignore it. pub secondary_index_removed: Vec<(String, String)>, - /// Vector index mutations this put performed, so a transactional caller - /// can push `UndoEntry::InsertVector` reversals (which also undo the - /// paired `vector_doc_map` entry). Empty when the document had no vector - /// fields. Autocommit callers ignore it. - pub vector_inserts: Vec, - /// `(spatial_index_key, entry_id)` pairs this put inserted into per-field - /// spatial R-trees, so a transactional caller can push - /// `UndoEntry::SpatialInsert` reversals. Empty when the document had no - /// spatial fields. Autocommit callers ignore it. - pub spatial_inserts: Vec<(SpatialIndexKey, u64)>, + /// Undo entries for every in-memory mutation this put made, in the order + /// it made them: R-tree and `spatial_doc_map` entries, vector nodes and + /// indexes, and sparse-vector postings. A dropped redb transaction does not reverse + /// them. A transactional caller pushes them onto its undo log. An + /// autocommit caller that abandons the write reverses them with + /// `undo_memory_effects`. + pub memory_undo: Vec, /// Pre-images of the column-stats read-modify-write this put performed, so a /// transactional caller can push `UndoEntry::StatsRestore` reversals. Each /// element is `(stats_key, prior_bytes)`: `prior_bytes = Some(b)` restores @@ -170,6 +167,7 @@ pub(in crate::data::executor) fn map_enforcement_error(e: ErrorCode) -> crate::E | ErrorCode::CollectionDraining { .. } | ErrorCode::RecursionDepthExceeded { .. } | ErrorCode::UndefinedColumn { .. } + | ErrorCode::TextColumn { .. } | ErrorCode::Internal { .. } | ErrorCode::Unsupported { .. } | ErrorCode::RollbackFailed { .. } @@ -178,12 +176,18 @@ pub(in crate::data::executor) fn map_enforcement_error(e: ErrorCode) -> crate::E | ErrorCode::DivisionByZero | ErrorCode::UndefinedFunction { .. } | ErrorCode::DataException { .. } + | ErrorCode::NumericValueOutOfRange { .. } + | ErrorCode::InvalidTextRepresentation { .. } + | ErrorCode::DatatypeMismatch { .. } + | ErrorCode::InvalidDatetimeFormat { .. } + | ErrorCode::DatetimeFieldOverflow { .. } | ErrorCode::DispatchCapacity { .. } | ErrorCode::ExpiredBeforeExecution | ErrorCode::BadRequest { .. } | ErrorCode::TransactionRollback { .. } | ErrorCode::ActiveSqlTransaction { .. } - | ErrorCode::DependentObjectsExist { .. }) => crate::Error::DataPlane(other), + | ErrorCode::DependentObjectsExist { .. } + | ErrorCode::NodeLabelLimit { .. }) => crate::Error::DataPlane(other), } } diff --git a/nodedb/src/data/executor/handlers/point/apply_put/vector/put.rs b/nodedb/src/data/executor/handlers/point/apply_put/vector/put.rs index 72dc9ccfc..6ed6e0d1c 100644 --- a/nodedb/src/data/executor/handlers/point/apply_put/vector/put.rs +++ b/nodedb/src/data/executor/handlers/point/apply_put/vector/put.rs @@ -3,9 +3,12 @@ //! Index a document's vectors into their HNSW collections on a point put. use crate::data::executor::core_loop::CoreLoop; +use crate::data::executor::handlers::transaction::undo::UndoEntry; +use crate::data::executor::handlers::vector_direct_row::VectorIndexKey; use crate::data::executor::vector_string::floats_from_value; -use super::types::{VectorFieldInsert, VectorIndexDelta, VectorIndexPutParams}; +use super::fields::DEFAULT_VECTOR_FIELD; +use super::types::{VectorFieldInsert, VectorIndexPutParams}; /// A document vector field whose width differs from its index: the /// caller's data error, SQLSTATE `22000`, naming the field. @@ -21,10 +24,14 @@ fn field_dimension_mismatch(field_name: &str, expected: usize, got: usize) -> cr impl CoreLoop { /// HNSW vector indexing side-effect: index declared strict-schema /// `Vector(dim)` columns, or (schemaless) fields matched by registered - /// `vector_params`, into the corresponding `VectorCollection`. + /// `vector_params` and declared `VECTOR(n)` columns, into the + /// corresponding `VectorCollection`. /// - /// Returns the `(index_key, vector_id)` pairs inserted so a transactional - /// caller can push `UndoEntry::InsertVector` reversals. Each inserted + /// Pushes an undo entry onto `undo` for each in-memory mutation as it + /// makes it: a vector index it creates, the prior node it soft-deletes, + /// and the node it inserts. The entries cover every mutation that ran, + /// also when a later field fails. Returns how many nodes it inserted. + /// Each inserted /// vector is also recorded in `vector_doc_map` keyed by the hex surrogate /// row key, so `apply_point_delete` can soft-delete it when the owning /// document is removed (closing the vector-orphan leak). @@ -36,15 +43,17 @@ impl CoreLoop { /// this LSN is skipped rather than re-appended as a duplicate HNSW node. /// /// Fails when a vector's width disagrees with the index it would land in — - /// either the width declared by `CREATE VECTOR INDEX ... DIM ` or the - /// width an already-materialized index carries. The write is refused - /// rather than the field skipped: a document that silently loses its + /// the width declared by `CREATE VECTOR INDEX ... DIM `, the width of a + /// declared `VECTOR(n)` column, or the width an already-materialized index + /// carries. The write is refused rather than the field skipped: a + /// document that silently loses its /// embedding is indistinguishable, at query time, from one that was never /// similar to anything. pub(in crate::data::executor) fn apply_point_put_vector_indexes( &mut self, params: VectorIndexPutParams<'_>, - ) -> crate::Result> { + undo: &mut Vec, + ) -> crate::Result { let VectorIndexPutParams { database_id, tid, @@ -53,7 +62,7 @@ impl CoreLoop { value, wal_lsn, } = params; - let mut inserts: Vec = Vec::new(); + let mut inserted: Vec = Vec::new(); // Vector index: if the strict schema declares Vector(dim) columns, // extract float arrays and insert into HNSW so KNN search works. @@ -87,22 +96,28 @@ impl CoreLoop { // all of them replayed before the core serves a request, // so a live write is never named. let skip = wal_lsn != 0 && self.vector_replay_skips(wal_lsn); - self.ensure_vector_collection(&index_key, &index_key, *dim as usize)?; + self.ensure_vector_collection_with_undo( + &index_key, + &index_key, + *dim as usize, + undo, + )?; if skip { continue; } - if let Some(delta) = - self.remove_then_insert_vector_field(VectorFieldInsert { + if self.remove_then_insert_vector_field( + VectorFieldInsert { database_id, tid, - index_key, + index_key: index_key.clone(), collection, field_name, storage_key, floats, - })? - { - inserts.push(delta); + }, + undo, + )? { + inserted.push(index_key); } } } @@ -119,9 +134,10 @@ impl CoreLoop { let bare_key = (db_key, tid_key, collection.to_string()); let field_names = self.schemaless_vector_field_names(database_id, tid, collection); - // Each field name maps back to its `vector_params` map key: either - // the field-qualified key (if one was registered) or the bare key - // (single default-"embedding" field, no per-field registration). + // Each field name maps back to its `vector_params` map key: the + // field-qualified key (if one was registered), the bare key for + // the default "embedding" field, or the field-qualified key of a + // declared column with no registration (default parameters). let schemaless_keys: Vec<( (nodedb_types::DatabaseId, crate::types::TenantId, String), String, @@ -129,7 +145,9 @@ impl CoreLoop { .into_iter() .map(|field| { let qualified = (db_key, tid_key, format!("{field_prefix}{field}")); - let params_key = if self.vector_params.contains_key(&qualified) { + let params_key = if self.vector_params.contains_key(&qualified) + || field != DEFAULT_VECTOR_FIELD + { qualified } else { bare_key.clone() @@ -154,25 +172,37 @@ impl CoreLoop { let store_key = Self::vector_index_key(database_id, tid, collection, field_name); self.check_vector_width(&store_key, field_name, floats.len())?; + // A declared `VECTOR(n)` column refuses another width, as + // a strict schema's vector column does. + if let Some(declared) = self.declared_schemaless_vector_dim( + database_id, + tid, + collection, + field_name, + ) && declared != floats.len() + { + return Err(field_dimension_mismatch(field_name, declared, floats.len())); + } let dim = floats.len(); // Same stamp gate as the strict arm above. let skip = wal_lsn != 0 && self.vector_replay_skips(wal_lsn); - self.ensure_vector_collection(&store_key, params_key, dim)?; + self.ensure_vector_collection_with_undo(&store_key, params_key, dim, undo)?; if skip { continue; } - if let Some(delta) = - self.remove_then_insert_vector_field(VectorFieldInsert { + if self.remove_then_insert_vector_field( + VectorFieldInsert { database_id, tid, - index_key: store_key, + index_key: store_key.clone(), collection, field_name, storage_key, floats, - })? - { - inserts.push(delta); + }, + undo, + )? { + inserted.push(store_key); } } } @@ -183,11 +213,30 @@ impl CoreLoop { // once the whole record landed, so a rollback finds its inserts in the // growing segment. if !self.recording_redo_undo() { - for delta in &inserts { - self.settle_vector_collection(&delta.index_key); + for index_key in &inserted { + self.settle_vector_collection(index_key); } } - Ok(inserts) + Ok(inserted.len()) + } + + /// `ensure_vector_collection`, plus an undo entry when the call creates + /// the index. + fn ensure_vector_collection_with_undo( + &mut self, + index_key: &VectorIndexKey, + config_key: &VectorIndexKey, + dim: usize, + undo: &mut Vec, + ) -> crate::Result<()> { + let created = !self.vector_collections.contains_key(index_key); + self.ensure_vector_collection(index_key, config_key, dim)?; + if created { + undo.push(UndoEntry::VectorCollectionCreated { + index_key: index_key.clone(), + }); + } + Ok(()) } /// Reject a vector whose width disagrees with the index it targets. @@ -230,14 +279,19 @@ impl CoreLoop { /// Binds the vector node to the document's global surrogate so /// cross-engine identity holds: a search hit resolves back to this row's /// surrogate (and thus its user PK at the response boundary) instead of - /// leaking a headless local node id. Returns `Ok(None)` if `index_key`'s + /// leaking a headless local node id. Returns `Ok(false)` if `index_key`'s /// `VectorCollection` was somehow absent (defensive — it was just /// populated via `entry().or_insert_with()` by the caller), and the /// collection's typed error when the vector does not fit it. + /// + /// Pushes a `DeleteVector` undo entry for the prior node it removes, a + /// `DisplacedVector` entry for an unrecorded node bound to the surrogate, + /// and an `InsertVector` entry for the node it inserts. fn remove_then_insert_vector_field( &mut self, params: VectorFieldInsert<'_>, - ) -> crate::Result> { + undo: &mut Vec, + ) -> crate::Result { let VectorFieldInsert { database_id, tid, @@ -247,17 +301,32 @@ impl CoreLoop { storage_key, floats, } = params; - let _ = self.remove_document_vector_index_field( + if let Some(prior) = self.remove_document_vector_index_field( database_id, tid, collection, field_name, storage_key, - ); + ) { + undo.push(prior.into_delete_undo()); + } let Some(coll) = self.vector_collections.get_mut(&index_key) else { - return Ok(None); + return Ok(false); }; - let vector_id = coll.insert_with_surrogate(floats, storage_key.surrogate())?; + // A node still bound to the surrogate has no `vector_doc_map` entry: + // a direct vector write bound it. `insert_with_surrogate` displaces + // such a node, so it is deleted here first, where undo can see it. + let surrogate = storage_key.surrogate(); + if let Some(displaced) = coll.local_for_surrogate(surrogate) + && coll.delete(displaced) + { + undo.push(UndoEntry::DisplacedVector { + index_key: index_key.clone(), + vector_id: displaced, + surrogate, + }); + } + let vector_id = coll.insert_with_surrogate(floats, surrogate)?; self.vector_doc_map.insert( ( index_key.0, @@ -268,13 +337,14 @@ impl CoreLoop { ), vector_id, ); - Ok(Some(VectorIndexDelta { + undo.push(UndoEntry::InsertVector { index_key, vector_id, collection: collection.to_string(), field: field_name.to_string(), - doc_id: storage_key, - })) + doc_id: Some(storage_key), + }); + Ok(true) } } @@ -395,25 +465,31 @@ mod tests { register_bare_field(core, db_id, tid, collection); let first = doc_with_vectors(&[("embedding", &[1.0, 0.0, 0.0])]); - core.apply_point_put_vector_indexes(VectorIndexPutParams { - database_id: db_id, - tid, - collection, - storage_key, - value: &first, - wal_lsn: 0, - }) + core.apply_point_put_vector_indexes( + VectorIndexPutParams { + database_id: db_id, + tid, + collection, + storage_key, + value: &first, + wal_lsn: 0, + }, + &mut Vec::new(), + ) .expect("vector indexing must accept this fixture"); let second = doc_with_vectors(&[("embedding", &[0.0, 1.0, 0.0])]); - core.apply_point_put_vector_indexes(VectorIndexPutParams { - database_id: db_id, - tid, - collection, - storage_key, - value: &second, - wal_lsn: 0, - }) + core.apply_point_put_vector_indexes( + VectorIndexPutParams { + database_id: db_id, + tid, + collection, + storage_key, + value: &second, + wal_lsn: 0, + }, + &mut Vec::new(), + ) .expect("vector indexing must accept this fixture"); assert_eq!( @@ -452,14 +528,17 @@ mod tests { ("embedding", &[1.0, 0.0, 0.0]), ("title_vec", &[0.0, 1.0, 0.0, 0.0]), ]); - core.apply_point_put_vector_indexes(VectorIndexPutParams { - database_id: db_id, - tid, - collection, - storage_key, - value: &doc, - wal_lsn: 0, - }) + core.apply_point_put_vector_indexes( + VectorIndexPutParams { + database_id: db_id, + tid, + collection, + storage_key, + value: &doc, + wal_lsn: 0, + }, + &mut Vec::new(), + ) .expect("vector indexing must accept this fixture"); assert_eq!( @@ -487,14 +566,17 @@ mod tests { register_bare_field(core, db_id, tid, collection); let doc = doc_with_vectors(&[("embedding", &[1.0, 0.0, 0.0])]); - core.apply_point_put_vector_indexes(VectorIndexPutParams { - database_id: db_id, - tid, - collection, - storage_key, - value: &doc, - wal_lsn: 0, - }) + core.apply_point_put_vector_indexes( + VectorIndexPutParams { + database_id: db_id, + tid, + collection, + storage_key, + value: &doc, + wal_lsn: 0, + }, + &mut Vec::new(), + ) .expect("vector indexing must accept this fixture"); let key = CoreLoop::vector_index_key(db_id, tid, collection, "embedding"); core.vector_collections @@ -513,6 +595,69 @@ mod tests { ); } + /// A put over a row whose surrogate binds a node no `vector_doc_map` + /// entry records displaces that node. Reversing the put's undo entries + /// brings the node and its binding back. + #[test] + fn undoing_a_put_restores_a_displaced_unrecorded_node() { + let mut harness = make_core(); + let core = &mut harness.core; + let (db_id, tid, collection) = (0u64, 1u64, "docs"); + let surrogate = Surrogate::new(1); + let storage_key = crate::engine::document::store::StorageKey::for_surrogate(surrogate); + register_bare_field(core, db_id, tid, collection); + + // A direct vector write binds a node to the surrogate. It records no + // `vector_doc_map` entry. + let key = CoreLoop::vector_index_key(db_id, tid, collection, "embedding"); + core.ensure_vector_collection(&key, &key, 3) + .expect("vector index"); + let direct = core + .vector_collections + .get_mut(&key) + .expect("collection") + .insert_with_surrogate(vec![0.0, 1.0, 0.0], surrogate) + .expect("vector write"); + + let doc = doc_with_vectors(&[("embedding", &[1.0, 0.0, 0.0])]); + let mut undo = Vec::new(); + core.apply_point_put_vector_indexes( + VectorIndexPutParams { + database_id: db_id, + tid, + collection, + storage_key, + value: &doc, + wal_lsn: 0, + }, + &mut undo, + ) + .expect("vector indexing must accept this fixture"); + assert_eq!(live_count(core, db_id, tid, collection, "embedding"), 1); + assert_ne!( + core.vector_collections + .get(&key) + .and_then(|c| c.local_for_surrogate(surrogate)), + Some(direct), + "the put must bind its own node" + ); + + core.undo_memory_effects(db_id, tid, undo) + .expect("the put's entries reverse"); + + let coll = core.vector_collections.get(&key).expect("collection stays"); + assert_eq!(coll.live_count(), 1, "only the displaced node is live"); + assert_eq!( + coll.local_for_surrogate(surrogate), + Some(direct), + "the displaced node is bound to the surrogate again" + ); + assert!( + core.vector_doc_map.is_empty(), + "the put's `vector_doc_map` entry is gone, and none existed before" + ); + } + /// A schemaless vector arriving as an SQL string literal must be parsed /// and indexed like an `ARRAY[...]` literal. /// @@ -543,14 +688,17 @@ mod tests { )]))) .expect("encode doc"); - core.apply_point_put_vector_indexes(VectorIndexPutParams { - database_id: db_id, - tid, - collection, - storage_key, - value: &body, - wal_lsn: 0, - }) + core.apply_point_put_vector_indexes( + VectorIndexPutParams { + database_id: db_id, + tid, + collection, + storage_key, + value: &body, + wal_lsn: 0, + }, + &mut Vec::new(), + ) .expect("vector indexing must accept a JSON-string embedding"); assert_eq!( @@ -583,14 +731,17 @@ mod tests { ]))) .expect("encode doc"); - let res = core.apply_point_put_vector_indexes(VectorIndexPutParams { - database_id: db_id, - tid, - collection, - storage_key, - value: &body, - wal_lsn: 0, - }); + let res = core.apply_point_put_vector_indexes( + VectorIndexPutParams { + database_id: db_id, + tid, + collection, + storage_key, + value: &body, + wal_lsn: 0, + }, + &mut Vec::new(), + ); assert!( matches!(res, Err(crate::Error::DataException { .. })), @@ -625,14 +776,17 @@ mod tests { Surrogate::new(surrogate), ); let doc = doc_with_vectors(&[("embedding", &[1.0, 0.0, 0.0])]); - core.apply_point_put_vector_indexes(VectorIndexPutParams { - database_id: db_id, - tid, - collection, - storage_key, - value: &doc, - wal_lsn, - }) + core.apply_point_put_vector_indexes( + VectorIndexPutParams { + database_id: db_id, + tid, + collection, + storage_key, + value: &doc, + wal_lsn, + }, + &mut Vec::new(), + ) .expect("vector indexing must accept this fixture"); } assert_eq!( diff --git a/nodedb/src/data/executor/handlers/point/apply_put/vector/remove.rs b/nodedb/src/data/executor/handlers/point/apply_put/vector/remove.rs index e9c3f9b5c..526b00c55 100644 --- a/nodedb/src/data/executor/handlers/point/apply_put/vector/remove.rs +++ b/nodedb/src/data/executor/handlers/point/apply_put/vector/remove.rs @@ -40,11 +40,12 @@ impl CoreLoop { // replaces the recorded node, so the recorded id can name a node that // is already soft-deleted while the live one keeps scoring. let mut vector_id = recorded; + let mut node_deleted = false; if let Some(coll) = self.vector_collections.get_mut(&index_key) { if let Some(bound) = coll.local_for_surrogate(storage_key.surrogate()) { vector_id = bound; } - coll.delete(vector_id); + node_deleted = coll.delete(vector_id); } Some(VectorIndexDelta { index_key, @@ -52,6 +53,7 @@ impl CoreLoop { collection: collection.to_string(), field: field.to_string(), doc_id: storage_key, + node_deleted, }) } @@ -62,8 +64,8 @@ impl CoreLoop { /// the surrogate's old embedding before inserting the new one, since /// `insert_with_surrogate` appends rather than replaces). /// - /// Candidate fields come from the same strict-schema / `vector_params` - /// enumeration the put path uses, so each `vector_doc_map` entry is looked + /// Candidate fields come from the same strict-schema / `vector_params` / + /// declared-column enumeration the put path uses, so each `vector_doc_map` entry is looked /// up by its exact key (via `remove_document_vector_index_field`) instead /// of scanning the whole map. Returns the removed `(index_key, vector_id)` /// deltas so a transactional caller can push `UndoEntry::DeleteVector` diff --git a/nodedb/src/data/executor/handlers/point/apply_put/vector/types.rs b/nodedb/src/data/executor/handlers/point/apply_put/vector/types.rs index 1ab9f9db0..3f4a3865a 100644 --- a/nodedb/src/data/executor/handlers/point/apply_put/vector/types.rs +++ b/nodedb/src/data/executor/handlers/point/apply_put/vector/types.rs @@ -29,6 +29,25 @@ pub(in crate::data::executor) struct VectorIndexDelta { pub collection: String, pub field: String, pub doc_id: crate::engine::document::store::StorageKey, + /// Whether the removal tombstoned `vector_id`. `false` when the node was + /// dead already and only the `vector_doc_map` entry went. + pub node_deleted: bool, +} + +impl VectorIndexDelta { + /// The undo entry that reverses this removal. + pub(in crate::data::executor) fn into_delete_undo( + self, + ) -> crate::data::executor::handlers::transaction::undo::UndoEntry { + crate::data::executor::handlers::transaction::undo::UndoEntry::DeleteVector { + index_key: self.index_key, + vector_id: self.vector_id, + collection: self.collection, + field: self.field, + doc_id: Some(self.doc_id), + node_deleted: self.node_deleted, + } + } } /// Inputs to `remove_then_insert_vector_field`, the shared per-field diff --git a/nodedb/src/data/executor/handlers/point/delete.rs b/nodedb/src/data/executor/handlers/point/delete.rs index 92e70f54a..9dbb4a651 100644 --- a/nodedb/src/data/executor/handlers/point/delete.rs +++ b/nodedb/src/data/executor/handlers/point/delete.rs @@ -8,6 +8,7 @@ use tracing::debug; use crate::bridge::envelope::{ErrorCode, Response, WriteSetEntry}; use crate::data::executor::core_loop::CoreLoop; use crate::data::executor::core_loop::redo_image::versioned_point_images; +use crate::data::executor::enforcement::chain_guard::{AbandonedWrite, abandon_write}; use crate::data::executor::enforcement::write_hook::{self, HookCtx, ImageBody, WriteImages}; use crate::data::executor::handlers::partial_refusal::refusal_after_partial_apply; use crate::data::executor::handlers::point::apply_delete::PointDeleteParams; @@ -88,7 +89,7 @@ impl CoreLoop { Ok(txn) => txn, Err(e) => return self.response_error(task, e), }; - let outcome = match self.apply_point_delete( + let mut outcome = match self.apply_point_delete( &txn, PointDeleteParams { database_id, @@ -104,6 +105,11 @@ impl CoreLoop { Ok(outcome) => outcome, Err(e) => return self.response_error(task, e), }; + // Every abort below drops `txn` uncommitted, which reverses the + // durable writes only. `abandon_write` reverses the in-memory + // cascades and the target rows' cache and index entries. + let storage_key = crate::engine::document::store::StorageKey::for_surrogate(surrogate); + let memory_undo = std::mem::take(&mut outcome.memory_undo); // Image-folding enforcement, inside the SAME transaction the removal was // staged in: a materialized-sum target write is itself a document write, // so the debit and the row's removal land or roll back together. A @@ -125,15 +131,19 @@ impl CoreLoop { ) { Ok(enforcement) => enforcement, Err(e) => { - // `apply_point_delete` already invalidated this row's cache - // entry, and dropping `txn` reverses every durable write it - // staged, so nothing else has to be undone here. + let e = abandon_write( + self, + AbandonedWrite::row(database_id, tid, collection, &storage_key) + .undo(memory_undo), + e, + ); return self.response_error(task, e); } }, None => Default::default(), }; let target_write_set = write_hook::target_write_set(&enforcement.target_writes); + let target_writes = enforcement.target_writes; // A delete subtracts the removed row's amount, so removing one leg of a // balanced journal on its own is a violation. Settled before the commit, @@ -141,16 +151,27 @@ impl CoreLoop { if let Err(e) = self.settle_balanced_entries(database_id, tid, collection, enforcement.balanced_entries) { + let e = abandon_write( + self, + AbandonedWrite::row(database_id, tid, collection, &storage_key) + .undo(memory_undo) + .targets(target_writes), + e, + ); return self.response_error(task, e); } if let Err(e) = txn.commit() { - return self.response_error( - task, - ErrorCode::Internal { + let e = abandon_write( + self, + AbandonedWrite::row(database_id, tid, collection, &storage_key) + .undo(memory_undo) + .targets(target_writes), + crate::Error::DataPlane(ErrorCode::Internal { detail: format!("commit: {e}"), - }, + }), ); + return self.response_error(task, e); } let prior = outcome.prior_value; @@ -197,20 +218,15 @@ impl CoreLoop { }; write_set.extend(target_write_set); if let Some(prior_bytes) = prior.as_deref() { - let old_converted = self.resolve_event_payload( - task.request.database_id.as_u64(), - tid, - collection, - prior_bytes, - ); // `document_identity` is read again below for `RETURNING`'s `id` // field, so the event-emit boundary gets a clone rather than the // move. self.emit_document_delete_event( task, + tid, collection, document_identity.clone(), - Some(old_converted.as_deref().unwrap_or(prior_bytes)), + Some(prior_bytes), ); } @@ -221,6 +237,7 @@ impl CoreLoop { // document with none of the row's real columns. The schema borrow is // scoped so the response build below can take `self` mutably. let doc = { + let identity_column = self.identity_column(database_id, tid, collection); let strict_schema = self .doc_configs .get(&( @@ -232,7 +249,12 @@ impl CoreLoop { StorageMode::Strict { schema } => Some(schema), StorageMode::Schemaless => None, }); - returning_doc::from_stored(prior_bytes, &document_identity, strict_schema) + returning_doc::from_stored( + prior_bytes, + &document_identity, + strict_schema, + &identity_column, + ) }; let doc = match doc { Ok(doc) => doc, @@ -342,6 +364,7 @@ impl CoreLoop { &body, &storage_key.to_identity(), strict_schema, + &self.identity_column(database_id, tid, collection), tid, collection, ) @@ -573,4 +596,240 @@ mod tests { "a refused delete must leave the chained row in place" ); } + + /// The in-memory index state of one source row. + #[derive(Debug, PartialEq)] + struct RowIndexes { + rtree: Vec<(u64, nodedb_types::BoundingBox)>, + spatial_doc: Option, + live_vectors: usize, + bound_vector: Option, + recorded_vector: Option, + node_deleted: bool, + } + + fn row_indexes(core: &CoreLoop, surrogate: Surrogate, document_id: &str) -> RowIndexes { + let storage_key = nodedb_types::StorageKey::for_surrogate(surrogate); + let entry_id = + crate::data::executor::handlers::point::apply_put::SpatialEntryId::from_storage_key( + storage_key, + ) + .as_u64(); + let db = DatabaseId::DEFAULT; + let tenant = TenantId::new(TID); + let rtree: Vec<(u64, nodedb_types::BoundingBox)> = core + .spatial_indexes + .get(&(db, tenant, SOURCE.to_string(), "loc".to_string())) + .map(|rt| rt.entries().into_iter().map(|e| (e.id, e.bbox)).collect()) + .unwrap_or_default(); + let spatial_doc = core + .spatial_doc_map + .get(&(db, tenant, SOURCE.to_string(), "loc".to_string(), entry_id)) + .cloned(); + let vector_key = CoreLoop::vector_index_key(DB, TID, SOURCE, "embedding"); + let vectors = core.vector_collections.get(&vector_key); + RowIndexes { + rtree, + spatial_doc, + live_vectors: vectors.map(|c| c.live_count()).unwrap_or(0), + bound_vector: vectors.and_then(|c| c.local_for_surrogate(surrogate)), + recorded_vector: core + .vector_doc_map + .get(&( + db, + tenant, + SOURCE.to_string(), + "embedding".to_string(), + storage_key, + )) + .copied(), + node_deleted: core.is_node_deleted(DB, TID, SOURCE, document_id), + } + } + + /// A delete refused after `apply_point_delete` ran leaves the row's + /// R-tree entry, vector node, reverse maps and node mark as they were. + /// The refusal here is the fold's: the target row it debits is gone. + #[test] + fn a_delete_refused_after_apply_leaves_every_index() { + let dir = tempfile::tempdir().expect("tempdir"); + let mut core = seeded_core(dir.path()); + let task = make_default_task(); + let targets = resolved(); + core.vector_params.insert( + (DatabaseId::DEFAULT, TenantId::new(TID), SOURCE.to_string()), + crate::engine::vector::hnsw::HnswParams::default(), + ); + + let surrogate = Surrogate(31); + let body = doc_format::encode_to_msgpack(&serde_json::json!({ + "account_id": A1, + "amount": 40, + "loc": {"type": "Point", "coordinates": [1.0, 2.0]}, + "embedding": [1.0, 0.0, 0.0], + })); + assert_eq!(insert(&mut core, &task, surrogate, &body), Status::Ok); + let before = row_indexes(&core, surrogate, "e31"); + assert_eq!(before.rtree.len(), 1, "the insert indexed the geometry"); + assert!(before.spatial_doc.is_some()); + assert_eq!(before.live_vectors, 1, "the insert indexed the vector"); + assert!(before.bound_vector.is_some()); + assert_eq!(before.recorded_vector, before.bound_vector); + assert!(!before.node_deleted); + + let target_key = nodedb_types::StorageKey::for_surrogate(T1); + core.sparse + .delete(DB, TID, TARGET, &target_key) + .expect("drop the target row"); + core.doc_cache.invalidate(DB, TID, TARGET, &target_key); + + let resp = core.execute_point_delete( + &task, + PointDeleteExec { + tid: TID, + collection: SOURCE, + document_id: "e31", + surrogate: Some(surrogate), + returning: None, + rls_filters: &[], + rls_write_check: &nodedb_types::RlsWriteCheck::NoPolicyApplies, + resolved_sum_targets: &targets, + }, + ); + assert_eq!( + resp.status, + Status::Error, + "the fold must refuse the delete" + ); + assert!( + !matches!( + resp.error_code.as_deref(), + Some(ErrorCode::RollbackFailed { .. }) + ), + "every in-memory entry must reverse, got {:?}", + resp.error_code + ); + assert!( + core.sparse + .get( + DB, + TID, + SOURCE, + &nodedb_types::StorageKey::for_surrogate(surrogate) + ) + .expect("read back") + .is_some(), + "the refused delete leaves the row" + ); + assert_eq!( + row_indexes(&core, surrogate, "e31"), + before, + "the refused delete leaves every in-memory index as it was" + ); + } + + /// [`seeded_core`] with a strict source that declares a `SparseVector` + /// column. + fn strict_sparse_core(dir: &std::path::Path) -> CoreLoop { + use nodedb_types::columnar::{ColumnDef, ColumnType, StrictSchema}; + let mut core = seeded_core(dir); + let schema = StrictSchema::new(vec![ + ColumnDef::required("_rowid", ColumnType::Int64), + ColumnDef::nullable("account_id", ColumnType::String), + ColumnDef::nullable("amount", ColumnType::Int64), + ColumnDef::nullable("terms", ColumnType::SparseVector), + ]) + .expect("schema"); + let mut source = CollectionConfig::new(SOURCE) + .with_storage_mode(nodedb_physical::physical_plan::StorageMode::Strict { schema }); + source.enforcement.materialized_sum_sources = vec![binding()]; + core.doc_configs.insert(config_key(SOURCE), source); + core + } + + /// A delete refused after `apply_point_delete` ran leaves the row's + /// sparse-vector postings on a strict collection. The refusal here is the + /// fold's: the target row it debits is gone. + #[test] + fn a_delete_refused_after_apply_leaves_the_sparse_postings() { + let dir = tempfile::tempdir().expect("tempdir"); + let mut core = strict_sparse_core(dir.path()); + let task = make_default_task(); + let targets = resolved(); + + let surrogate = Surrogate(32); + let body = doc_format::encode_to_msgpack(&serde_json::json!({ + "account_id": A1, + "amount": 40, + "terms": "{3:0.5, 7:1.5}", + })); + assert_eq!(insert(&mut core, &task, surrogate, &body), Status::Ok); + let sparse_key = CoreLoop::sparse_index_key(DB, TID, SOURCE, "terms"); + let row_key = nodedb_types::StorageKey::for_surrogate(surrogate).to_string(); + let postings = |core: &CoreLoop| { + core.sparse_vector_indexes.get(&sparse_key).map(|index| { + ( + index.doc_count(), + index.doc_image(&row_key), + index.next_internal_id(), + ) + }) + }; + let before = postings(&core); + let Some((doc_count, image, _)) = &before else { + panic!("the insert must create the sparse index"); + }; + assert_eq!(*doc_count, 1, "the insert indexed the sparse vector"); + assert!(image.is_some(), "the row holds sparse postings"); + + let target_key = nodedb_types::StorageKey::for_surrogate(T1); + core.sparse + .delete(DB, TID, TARGET, &target_key) + .expect("drop the target row"); + core.doc_cache.invalidate(DB, TID, TARGET, &target_key); + + let resp = core.execute_point_delete( + &task, + PointDeleteExec { + tid: TID, + collection: SOURCE, + document_id: "e32", + surrogate: Some(surrogate), + returning: None, + rls_filters: &[], + rls_write_check: &nodedb_types::RlsWriteCheck::NoPolicyApplies, + resolved_sum_targets: &targets, + }, + ); + assert_eq!( + resp.status, + Status::Error, + "the fold must refuse the delete" + ); + assert!( + !matches!( + resp.error_code.as_deref(), + Some(ErrorCode::RollbackFailed { .. }) + ), + "every in-memory entry must reverse, got {:?}", + resp.error_code + ); + assert!( + core.sparse + .get( + DB, + TID, + SOURCE, + &nodedb_types::StorageKey::for_surrogate(surrogate) + ) + .expect("read back") + .is_some(), + "the refused delete leaves the row" + ); + assert_eq!( + postings(&core), + before, + "the refused delete leaves the sparse postings as they were" + ); + } } diff --git a/nodedb/src/data/executor/handlers/point/insert.rs b/nodedb/src/data/executor/handlers/point/insert.rs index 21323e8b2..8ad6005e8 100644 --- a/nodedb/src/data/executor/handlers/point/insert.rs +++ b/nodedb/src/data/executor/handlers/point/insert.rs @@ -13,7 +13,7 @@ use tracing::debug; use crate::bridge::envelope::{Response, WriteSetEntry}; use crate::data::executor::core_loop::CoreLoop; use crate::data::executor::core_loop::redo_image::versioned_point_images; -use crate::data::executor::enforcement::chain_guard::{self, ChainGuard}; +use crate::data::executor::enforcement::chain_guard::{self, AbandonedWrite, ChainGuard}; use crate::data::executor::enforcement::write_hook::{self, HookCtx, ImageBody, WriteImages}; use crate::data::executor::handlers::point::apply_put::PointPutParams; use crate::data::executor::task::ExecutionTask; @@ -143,7 +143,16 @@ impl CoreLoop { // shape by the RETURNING renderer. let mut response = match returning { Some(spec) => { - self.stored_returning_response(task, spec, rls_filters, None, &[]) + let identity_column = + self.identity_column(database_id, tid, collection); + self.stored_returning_response( + task, + spec, + rls_filters, + None, + &identity_column, + &[], + ) } None => self.response_affected(task, 0), }; @@ -197,13 +206,11 @@ impl CoreLoop { ) { Ok(o) => o, Err(e) => { - chain_guard::abort_after_apply( + let e = chain_guard::abort_after_apply( self, &mut chain, - database_id, - tid, - collection, - &storage_key, + AbandonedWrite::row(database_id, tid, collection, &storage_key), + e, ); return self.response_error(task, e); } @@ -215,13 +222,12 @@ impl CoreLoop { .settle(self, surrogate, &outcome.stored_value) .and_then(|()| chain.persist_head(self, &txn)) { - chain_guard::abort_after_apply( + let e = chain_guard::abort_after_apply( self, &mut chain, - database_id, - tid, - collection, - &storage_key, + AbandonedWrite::row(database_id, tid, collection, &storage_key) + .undo(std::mem::take(&mut outcome.memory_undo)), + e, ); return self.response_error(task, e); } @@ -239,15 +245,14 @@ impl CoreLoop { new: ImageBody::Submitted(value), }, ) { - Ok(outcome) => outcome, + Ok(enforcement) => enforcement, Err(e) => { - chain_guard::abort_after_apply( + let e = chain_guard::abort_after_apply( self, &mut chain, - database_id, - tid, - collection, - &storage_key, + AbandonedWrite::row(database_id, tid, collection, &storage_key) + .undo(std::mem::take(&mut outcome.memory_undo)), + e, ); return self.response_error(task, e); } @@ -257,6 +262,7 @@ impl CoreLoop { // against its OWN collection — the statement's own redo names only the // source row. let target_write_set = write_hook::target_write_set(&enforcement.target_writes); + let target_writes = enforcement.target_writes; // BALANCED is settled before the commit, so a single-row insert of one // journal leg — unbalanced by the constraint's own definition when the @@ -264,33 +270,30 @@ impl CoreLoop { if let Err(e) = self.settle_balanced_entries(database_id, tid, collection, enforcement.balanced_entries) { - chain_guard::abort_after_apply( + let e = chain_guard::abort_after_apply( self, &mut chain, - database_id, - tid, - collection, - &storage_key, + AbandonedWrite::row(database_id, tid, collection, &storage_key) + .undo(std::mem::take(&mut outcome.memory_undo)) + .targets(target_writes), + e, ); return self.response_error(task, e); } if let Err(e) = txn.commit() { - chain_guard::abort_after_apply( + let e = chain_guard::abort_after_apply( self, &mut chain, - database_id, - tid, - collection, - &storage_key, - ); - return self.response_error( - task, + AbandonedWrite::row(database_id, tid, collection, &storage_key) + .undo(std::mem::take(&mut outcome.memory_undo)) + .targets(target_writes), crate::Error::Storage { engine: "sparse".into(), detail: format!("commit: {e}"), }, ); + return self.response_error(task, e); } self.checkpoint_coordinator.mark_dirty("sparse", 1); @@ -348,6 +351,7 @@ impl CoreLoop { spec, rls_filters, strict_schema.as_ref(), + &self.identity_column(database_id, tid, collection), &[(&document_identity, stored_value.as_slice())], ) } else { diff --git a/nodedb/src/data/executor/handlers/point/put.rs b/nodedb/src/data/executor/handlers/point/put.rs index 8d00a4e6e..ff7b64753 100644 --- a/nodedb/src/data/executor/handlers/point/put.rs +++ b/nodedb/src/data/executor/handlers/point/put.rs @@ -5,10 +5,10 @@ use tracing::debug; -use crate::bridge::envelope::{ErrorCode, Response, WriteSetEntry}; +use crate::bridge::envelope::{Response, WriteSetEntry}; use crate::data::executor::core_loop::CoreLoop; use crate::data::executor::core_loop::redo_image::versioned_point_images; -use crate::data::executor::enforcement::chain_guard::{self, ChainGuard}; +use crate::data::executor::enforcement::chain_guard::{self, AbandonedWrite, ChainGuard}; use crate::data::executor::enforcement::write_hook::{self, HookCtx, ImageBody, WriteImages}; use crate::data::executor::handlers::point::apply_put::PointPutParams; use crate::data::executor::task::ExecutionTask; @@ -114,13 +114,11 @@ impl CoreLoop { ) { Ok(p) => p, Err(e) => { - chain_guard::abort_after_apply( + let e = chain_guard::abort_after_apply( self, &mut chain, - database_id, - tid, - collection, - &storage_key, + AbandonedWrite::row(database_id, tid, collection, &storage_key), + e, ); return self.response_error(task, e); } @@ -130,13 +128,12 @@ impl CoreLoop { .settle(self, surrogate, &prior.stored_value) .and_then(|()| chain.persist_head(self, &txn)) { - chain_guard::abort_after_apply( + let e = chain_guard::abort_after_apply( self, &mut chain, - database_id, - tid, - collection, - &storage_key, + AbandonedWrite::row(database_id, tid, collection, &storage_key) + .undo(std::mem::take(&mut prior.memory_undo)), + e, ); return self.response_error(task, e); } @@ -157,18 +154,18 @@ impl CoreLoop { let enforcement = match write_hook::run(self, &txn, &hook_ctx, images) { Ok(outcome) => outcome, Err(e) => { - chain_guard::abort_after_apply( + let e = chain_guard::abort_after_apply( self, &mut chain, - database_id, - tid, - collection, - &storage_key, + AbandonedWrite::row(database_id, tid, collection, &storage_key) + .undo(std::mem::take(&mut prior.memory_undo)), + e, ); return self.response_error(task, e); } }; let target_write_set = write_hook::target_write_set(&enforcement.target_writes); + let target_writes = enforcement.target_writes; // Settled before the commit: an autocommit statement is its own // transaction boundary, so a put that leaves a journal group unbalanced @@ -176,32 +173,30 @@ impl CoreLoop { if let Err(e) = self.settle_balanced_entries(database_id, tid, collection, enforcement.balanced_entries) { - chain_guard::abort_after_apply( + let e = chain_guard::abort_after_apply( self, &mut chain, - database_id, - tid, - collection, - &storage_key, + AbandonedWrite::row(database_id, tid, collection, &storage_key) + .undo(std::mem::take(&mut prior.memory_undo)) + .targets(target_writes), + e, ); return self.response_error(task, e); } if let Err(e) = txn.commit() { - chain_guard::abort_after_apply( + let e = chain_guard::abort_after_apply( self, &mut chain, - database_id, - tid, - collection, - &storage_key, - ); - return self.response_error( - task, - ErrorCode::Internal { + AbandonedWrite::row(database_id, tid, collection, &storage_key) + .undo(std::mem::take(&mut prior.memory_undo)) + .targets(target_writes), + crate::Error::Storage { + engine: "sparse".into(), detail: format!("commit: {e}"), }, ); + return self.response_error(task, e); } // Record the committed write's version against its surrogate + collection. @@ -249,6 +244,7 @@ impl CoreLoop { spec, rls_filters, strict_schema.as_ref(), + &self.identity_column(database_id, tid, collection), &[(&document_identity, prior.stored_value.as_slice())], ) } else { @@ -271,7 +267,7 @@ impl CoreLoop { #[cfg(test)] mod tests { use super::*; - use crate::bridge::envelope::Status; + use crate::bridge::envelope::{ErrorCode, Status}; use crate::data::executor::core_loop::tests::{make_core_with_dir, make_default_task}; use crate::data::executor::doc_format; use crate::engine::document::store::CollectionConfig; @@ -406,6 +402,75 @@ mod tests { ); } + /// A put refused after `apply_point_put` ran puts the row's in-memory + /// index back. The second put replaces the row's vector node, then its + /// materialized-sum fold is refused: the plan resolves no target row. + /// The prior node must be live and bound to the row again, the new node + /// gone, and the target balance unmoved. + #[test] + fn a_put_refused_after_apply_restores_the_prior_vector_node() { + let dir = tempfile::tempdir().expect("tempdir"); + let mut core = seeded_core(dir.path()); + core.vector_params.insert( + config_key(SOURCE), + crate::engine::vector::hnsw::HnswParams::default(), + ); + let task = make_default_task(); + let body = |amount: i64, embedding: &[f64]| { + doc_format::encode_to_msgpack(&serde_json::json!({ + "account_id": A1, + "amount": amount, + "embedding": embedding, + })) + }; + let row = Surrogate(21); + + assert_eq!( + put(&mut core, &task, &body(10, &[1.0, 0.0, 0.0])), + Status::Ok + ); + let key = CoreLoop::vector_index_key(DB, TID, SOURCE, "embedding"); + let prior = core.vector_collections[&key] + .local_for_surrogate(row) + .expect("the first put binds a node"); + + let refused = core.execute_point_put( + &task, + PointPutExec { + tid: TID, + collection: SOURCE, + document_id: "e1", + surrogate: row, + value: &body(30, &[0.0, 1.0, 0.0]), + returning: None, + rls_filters: &[], + resolved_sum_targets: &[], + }, + ); + assert_eq!( + refused.status, + Status::Error, + "the unresolved fold must refuse" + ); + + let coll = &core.vector_collections[&key]; + assert_eq!(coll.live_count(), 1, "the new node must be gone"); + assert_eq!( + coll.local_for_surrogate(row), + Some(prior), + "the prior node must be live and bound to the row again" + ); + let doc_key = ( + DatabaseId::DEFAULT, + TenantId::new(TID), + SOURCE.to_string(), + "embedding".to_string(), + StorageKey::for_surrogate(row), + ); + assert_eq!(core.vector_doc_map.get(&doc_key).copied(), Some(prior)); + assert_eq!(balance(&core, T1), "10", "the refused put moves no total"); + } + /// A put that carries `Surrogate::ZERO` is refused before any write: no /// row lands under the zero key and no target total moves. #[test] diff --git a/nodedb/src/data/executor/handlers/point/update_reindex_sparse.rs b/nodedb/src/data/executor/handlers/point/update_reindex_sparse.rs index 885562ea3..6391339fb 100644 --- a/nodedb/src/data/executor/handlers/point/update_reindex_sparse.rs +++ b/nodedb/src/data/executor/handlers/point/update_reindex_sparse.rs @@ -65,7 +65,15 @@ impl CoreLoop { // clears the `SparseVector` field must not leave the stale literal // searchable, and the re-insert below only re-adds fields present in // the new body. - self.remove_document_sparse_indexes(p.database_id, p.tid, p.collection, p.storage_key); + // The row this re-indexes is already committed, so the undo entries + // go unused. + self.remove_document_sparse_indexes( + p.database_id, + p.tid, + p.collection, + p.storage_key, + &mut Vec::new(), + ); // Re-extract from the new body via the exact put-time path. Sparse // extraction reads MessagePack; strict bodies are stored as Binary @@ -94,7 +102,16 @@ impl CoreLoop { p.new_body }; - self.apply_point_put_sparse_indexes(p.database_id, p.tid, p.collection, p.storage_key, mp); + // The row this re-indexes is already committed, so the undo entries + // go unused. + self.apply_point_put_sparse_indexes( + p.database_id, + p.tid, + p.collection, + p.storage_key, + mp, + &mut Vec::new(), + ); Ok(()) } } diff --git a/nodedb/src/data/executor/handlers/point/update_reindex_vector.rs b/nodedb/src/data/executor/handlers/point/update_reindex_vector.rs index d3e6a9c04..4b2966b9d 100644 --- a/nodedb/src/data/executor/handlers/point/update_reindex_vector.rs +++ b/nodedb/src/data/executor/handlers/point/update_reindex_vector.rs @@ -101,15 +101,19 @@ impl CoreLoop { // Live in-memory maintenance only: `wal_lsn = 0` disables the // checkpoint-straddle guard (the WAL record for this update is appended - // in the Control Plane, not here). - self.apply_point_put_vector_indexes(VectorIndexPutParams { - database_id: p.database_id, - tid: p.tid, - collection: p.collection, - storage_key: p.storage_key, - value: mp, - wal_lsn: 0, - })?; + // in the Control Plane, not here). The row this re-indexes is already + // committed, so the undo entries go unused. + self.apply_point_put_vector_indexes( + VectorIndexPutParams { + database_id: p.database_id, + tid: p.tid, + collection: p.collection, + storage_key: p.storage_key, + value: mp, + wal_lsn: 0, + }, + &mut Vec::new(), + )?; Ok(()) } } diff --git a/nodedb/src/data/executor/handlers/transaction/redo_apply/document.rs b/nodedb/src/data/executor/handlers/transaction/redo_apply/document.rs index f34fa9b19..6ae8bb6f0 100644 --- a/nodedb/src/data/executor/handlers/transaction/redo_apply/document.rs +++ b/nodedb/src/data/executor/handlers/transaction/redo_apply/document.rs @@ -19,7 +19,9 @@ use nodedb_types::Surrogate; use redb::WriteTransaction; use crate::data::executor::core_loop::CoreLoop; -use crate::data::executor::enforcement::chain_guard::{ChainGuard, abort_after_apply}; +use crate::data::executor::enforcement::chain_guard::{ + AbandonedWrite, ChainGuard, abandon_write, abort_after_apply, +}; use crate::data::executor::enforcement::write_hook::{self, HookCtx, ImageBody, WriteImages}; use crate::data::executor::handlers::point::apply_delete::PointDeleteParams; use crate::data::executor::handlers::point::apply_put::PointPutParams; @@ -127,7 +129,7 @@ impl CoreLoop { return Err(error); } }; - let outcome = match self.apply_point_put( + let mut outcome = match self.apply_point_put( &txn, PointPutParams { database_id: row.database_id, @@ -146,16 +148,24 @@ impl CoreLoop { ) { Ok(outcome) => outcome, Err(error) => { - self.abort_committed_write(&mut chain, row, &storage_key); - return Err(error); + return Err(abort_after_apply( + self, + &mut chain, + abandoned(row, &storage_key), + error, + )); } }; if let Err(error) = chain .settle(self, surrogate, &outcome.stored_value) .and_then(|()| chain.persist_head(self, &txn)) { - self.abort_committed_write(&mut chain, row, &storage_key); - return Err(error); + return Err(abort_after_apply( + self, + &mut chain, + abandoned(row, &storage_key).undo(std::mem::take(&mut outcome.memory_undo)), + error, + )); } // The fold reads the SUBMITTED body: `_chain_hash` wraps the row and no @@ -169,16 +179,26 @@ impl CoreLoop { new: ImageBody::Submitted(value), }, }; - let target_writes = match write_hook::run(self, &txn, &hook_ctx, images) { + let mut target_writes = match write_hook::run(self, &txn, &hook_ctx, images) { Ok(enforcement) => enforcement.target_writes, Err(error) => { - self.abort_committed_write(&mut chain, row, &storage_key); - return Err(error); + return Err(abort_after_apply( + self, + &mut chain, + abandoned(row, &storage_key).undo(std::mem::take(&mut outcome.memory_undo)), + error, + )); } }; if let Err(error) = commit_row(txn) { - self.abort_committed_write(&mut chain, row, &storage_key); - return Err(error); + return Err(abort_after_apply( + self, + &mut chain, + abandoned(row, &storage_key) + .undo(std::mem::take(&mut outcome.memory_undo)) + .targets(target_writes), + error, + )); } self.checkpoint_coordinator.mark_dirty("sparse", 1); @@ -201,12 +221,10 @@ impl CoreLoop { }); if self.recording_redo_undo() { let mut undo = Vec::new(); - push_target_undo(&mut undo, &target_writes); + push_target_undo(&mut undo, &mut target_writes); push_put_undo( &mut undo, DocumentRow { - database_id: row.database_id, - tid: row.tenant_id, collection: row.collection, storage_key, }, @@ -234,7 +252,7 @@ impl CoreLoop { }; let txn = self.sparse.begin_write()?; - let outcome = self.apply_point_delete( + let mut outcome = self.apply_point_delete( &txn, PointDeleteParams { database_id: row.database_id, @@ -247,21 +265,39 @@ impl CoreLoop { resolved_targets: &resolved, }, )?; - let target_writes = match outcome.prior_value { - Some(ref old) => { - write_hook::run( - self, - &txn, - &hook_ctx, - WriteImages::Delete { - old: ImageBody::Stored(old), - }, - )? - .target_writes - } + // An abort below drops `txn` uncommitted, which reverses the durable + // writes only. `abandon_write` reverses the in-memory cascades. + let mut target_writes = match outcome.prior_value { + Some(ref old) => match write_hook::run( + self, + &txn, + &hook_ctx, + WriteImages::Delete { + old: ImageBody::Stored(old), + }, + ) { + Ok(enforcement) => enforcement.target_writes, + Err(error) => { + let undo = std::mem::take(&mut outcome.memory_undo); + return Err(abandon_write( + self, + abandoned(row, &storage_key).undo(undo), + error, + )); + } + }, None => Vec::new(), }; - commit_row(txn)?; + if let Err(error) = commit_row(txn) { + let undo = std::mem::take(&mut outcome.memory_undo); + return Err(abandon_write( + self, + abandoned(row, &storage_key) + .undo(undo) + .targets(target_writes), + error, + )); + } self.checkpoint_coordinator.mark_dirty("sparse", 1); let removed = outcome.prior_value.clone(); @@ -271,12 +307,10 @@ impl CoreLoop { // reverse, so its undo is recorded either way. if self.recording_redo_undo() { let mut undo = Vec::new(); - push_target_undo(&mut undo, &target_writes); + push_target_undo(&mut undo, &mut target_writes); push_delete_undo( &mut undo, DocumentRow { - database_id: row.database_id, - tid: row.tenant_id, collection: row.collection, storage_key, }, @@ -299,24 +333,6 @@ impl CoreLoop { Ok(true) } - /// Reverse the in-memory effects of a put abandoned after - /// `apply_point_put` ran; the caller drops its transaction uncommitted. - fn abort_committed_write( - &mut self, - chain: &mut ChainGuard, - row: &CommittedDocWrite<'_>, - storage_key: &StorageKey, - ) { - abort_after_apply( - self, - chain, - row.database_id, - row.tenant_id, - row.collection, - storage_key, - ); - } - /// The sum targets a write to `collection` folds into: the open scope's /// under a committed-redo apply, else the restart-replay folds of the /// record at `record_lsn`. @@ -356,6 +372,12 @@ impl CoreLoop { } } +/// A put of `row` abandoned after `apply_point_put` ran. The caller drops its +/// transaction uncommitted. +fn abandoned<'a>(row: &CommittedDocWrite<'a>, storage_key: &'a StorageKey) -> AbandonedWrite<'a> { + AbandonedWrite::row(row.database_id, row.tenant_id, row.collection, storage_key) +} + fn commit_row(txn: WriteTransaction) -> crate::Result<()> { txn.commit().map_err(|e| crate::Error::Storage { engine: "sparse".into(), diff --git a/nodedb/src/data/executor/handlers/transaction/redo_apply/passes.rs b/nodedb/src/data/executor/handlers/transaction/redo_apply/passes.rs index 2d081de03..3d244e057 100644 --- a/nodedb/src/data/executor/handlers/transaction/redo_apply/passes.rs +++ b/nodedb/src/data/executor/handlers/transaction/redo_apply/passes.rs @@ -119,12 +119,13 @@ impl CoreLoop { ) -> PassRefusal { match self.rollback_undo_log(target.database_id, target.tid, undo) { Ok(()) => PassRefusal::RolledBack(cause), - Err((entry_index, detail)) => PassRefusal::RollbackFailed(ErrorCode::RollbackFailed { - entry_index, - detail: format!( - "rolling back a committed redo install that failed with {cause:?}: {detail}" - ), - }), + Err(mut undo_error) => { + undo_error.action = format!( + "rolling back a committed redo install that failed with {cause:?}: {}", + undo_error.action + ); + PassRefusal::RollbackFailed(ErrorCode::from(undo_error)) + } } } @@ -168,6 +169,7 @@ impl CoreLoop { (_, None) => Err(PassRefusal::RollbackFailed(ErrorCode::RollbackFailed { entry_index: 0, detail: "committed transaction redo lost its apply scope and its undo log".into(), + cause: None, })), } } diff --git a/nodedb/src/data/executor/handlers/transaction/undo/apply.rs b/nodedb/src/data/executor/handlers/transaction/undo/apply.rs index 0406369de..45e145c2b 100644 --- a/nodedb/src/data/executor/handlers/transaction/undo/apply.rs +++ b/nodedb/src/data/executor/handlers/transaction/undo/apply.rs @@ -3,14 +3,14 @@ //! Per-engine undo entry application logic. //! //! Each `apply_undo_*` method handles one engine family's undo entries. -//! All methods return `Err((entry_index, detail))` on fatal failure so the -//! caller can escalate to a typed `RollbackFailed` response. +//! All methods return an [`UndoError`] on fatal failure, and the caller +//! escalates it to a typed `RollbackFailed` response. use tracing::error; use crate::data::executor::core_loop::CoreLoop; -use super::{TimeseriesIngestUndo, UndoEntry}; +use super::{TimeseriesIngestUndo, UndoEntry, UndoError}; impl CoreLoop { // ── Vector ─────────────────────────────────────────────────────────────── @@ -20,7 +20,7 @@ impl CoreLoop { _tid: u64, entry_index: usize, entry: UndoEntry, - ) -> Result<(), (usize, String)> { + ) -> Result<(), UndoError> { match entry { UndoEntry::InsertVector { index_key, @@ -50,17 +50,20 @@ impl CoreLoop { Ok(()) } None => { - let detail = format!( - "vector index {:?} not found during undo of vector insert {}", - index_key, vector_id + let err = UndoError::mismatch( + entry_index, + format!( + "vector index {index_key:?} not found during undo of vector \ + insert {vector_id}" + ), ); error!( core = self.core_id, entry_index, - error = %detail, + error = %err, "transaction undo: vector index missing; shard state unknown" ); - Err((entry_index, detail)) + Err(err) } }, UndoEntry::DeleteVector { @@ -69,42 +72,83 @@ impl CoreLoop { collection, field, doc_id, - } => match self.vector_collections.get_mut(&index_key) { - Some(index) => { - index.undelete(vector_id); - // Restore the `vector_doc_map` entry the forward delete - // removed — without this a rolled-back delete leaves the - // doc→vector reverse lookup missing, so a later delete of - // the same document can never find (and soft-delete) its - // vector: a permanent orphan. Mirrors - // `apply_undo_spatial`'s `spatial_doc_map.insert`. `None` - // `doc_id` marks the direct primary-vector write path, - // which never populates `vector_doc_map` — skip it there. - if let Some(doc_id) = doc_id { - self.vector_doc_map.insert( - (index_key.0, index_key.1, collection, field, doc_id), - vector_id, + node_deleted, + } => { + // The forward delete dropped the node's surrogate binding + // with its tombstone, so a document node gets both back. + // A search hit on an unbound node resolves to no row. + // A node the forward delete found dead stays dead. + if node_deleted { + let Some(index) = self.vector_collections.get_mut(&index_key) else { + let err = UndoError::mismatch( + entry_index, + format!( + "vector index {index_key:?} not found during undo of vector \ + delete {vector_id}" + ), ); + error!( + core = self.core_id, + entry_index, + error = %err, + "transaction undo: vector index missing; shard state unknown" + ); + return Err(err); + }; + let restored = match doc_id { + Some(doc_id) => index.undelete_bound(vector_id, doc_id.surrogate()), + None => index.undelete(vector_id), + }; + if !restored { + return Err(UndoError::mismatch( + entry_index, + format!( + "vector index {index_key:?} holds no tombstone for node \ + {vector_id} during undo of vector delete" + ), + )); } - Ok(()) } - None => { - let detail = format!( - "vector index {:?} not found during undo of vector delete {}", - index_key, vector_id + // Restore the `vector_doc_map` entry the forward delete + // removed — without this a rolled-back delete leaves the + // doc→vector reverse lookup missing, so a later delete of + // the same document can never find (and soft-delete) its + // vector: a permanent orphan. Mirrors + // `apply_undo_spatial`'s `spatial_doc_map.insert`. `None` + // `doc_id` marks the direct primary-vector write path, + // which never populates `vector_doc_map` — skip it there. + if let Some(doc_id) = doc_id { + self.vector_doc_map.insert( + (index_key.0, index_key.1, collection, field, doc_id), + vector_id, ); - error!( - core = self.core_id, + } + Ok(()) + } + UndoEntry::DisplacedVector { + index_key, + vector_id, + surrogate, + } => { + let restored = self + .vector_collections + .get_mut(&index_key) + .is_some_and(|index| index.undelete_bound(vector_id, surrogate)); + if restored { + Ok(()) + } else { + Err(UndoError::mismatch( entry_index, - error = %detail, - "transaction undo: vector index missing; shard state unknown" - ); - Err((entry_index, detail)) + format!( + "vector index {index_key:?} cannot restore displaced node \ + {vector_id}: index missing or node holds no tombstone" + ), + )) } - }, - _ => Err(( + } + _ => Err(UndoError::mismatch( entry_index, - "apply_undo_vector called with non-vector entry".to_string(), + "apply_undo_vector called with non-vector entry", )), } } @@ -115,7 +159,7 @@ impl CoreLoop { &mut self, entry_index: usize, entry: UndoEntry, - ) -> Result<(), (usize, String)> { + ) -> Result<(), UndoError> { match entry { UndoEntry::ColumnarInsert { collection_key, @@ -165,9 +209,9 @@ impl CoreLoop { // Engine absent: no in-memory state to roll back. Ok(()) } - _ => Err(( + _ => Err(UndoError::mismatch( entry_index, - "apply_undo_columnar called with non-columnar entry".to_string(), + "apply_undo_columnar called with non-columnar entry", )), } } @@ -178,14 +222,14 @@ impl CoreLoop { &mut self, entry_index: usize, entry: UndoEntry, - ) -> Result<(), (usize, String)> { + ) -> Result<(), UndoError> { match entry { UndoEntry::TimeseriesIngest(token) => { self.restore_timeseries_ingest_preimage(entry_index, token) } - _ => Err(( + _ => Err(UndoError::mismatch( entry_index, - "apply_undo_timeseries called with non-timeseries entry".to_string(), + "apply_undo_timeseries called with non-timeseries entry", )), } } @@ -194,7 +238,7 @@ impl CoreLoop { &mut self, entry_index: usize, token: TimeseriesIngestUndo, - ) -> Result<(), (usize, String)> { + ) -> Result<(), UndoError> { let TimeseriesIngestUndo { collection_key, memtable_before, @@ -214,7 +258,7 @@ impl CoreLoop { .get(&collection_key) .map(nodedb_mem::ReservationToken::size); if reservation_now != reservation_bytes_before { - return Err(( + return Err(UndoError::mismatch( entry_index, format!( "timeseries reservation changed during deferred ingest for {:?}: before {:?}, now {:?}", @@ -234,9 +278,10 @@ impl CoreLoop { snapshot, config, ) .map_err(|error| { - ( + UndoError::failed( entry_index, - format!("timeseries memtable snapshot restore failed: {error}"), + "restoring the timeseries memtable snapshot", + error, ) })?; restored.restore_memory_bytes_for_undo(memory_bytes); @@ -247,9 +292,9 @@ impl CoreLoop { self.columnar_memtables.remove(&collection_key); } _ => { - return Err(( + return Err(UndoError::mismatch( entry_index, - "timeseries undo token has inconsistent memtable pre-image fields".into(), + "timeseries undo token has inconsistent memtable pre-image fields", )); } } @@ -654,6 +699,7 @@ mod tests { .expect("engine present"); let mut out: Vec<(i64, i64)> = engine .scan_memtable_rows() + .map(|row| row.expect("read")) .filter_map(|row| match (&row[0], &row[1]) { (Value::Integer(id), Value::Integer(v)) => Some((*id, *v)), _ => None, @@ -676,7 +722,9 @@ mod tests { // COMMIT of `UPDATE m SET v = 999` (empty filter = all rows). let updates = vec![( "v".to_string(), - nodedb_types::value_to_msgpack(&Value::Integer(999)).unwrap(), + nodedb_physical::physical_plan::UpdateValue::Literal( + nodedb_types::value_to_msgpack(&Value::Integer(999)).unwrap(), + ), )]; let plan = PhysicalPlan::Columnar(ColumnarOp::Update { collection: QualifiedCollection::new(DatabaseId::DEFAULT, "m"), @@ -831,6 +879,7 @@ mod tests { collection: "c".to_string(), field: "emb".to_string(), doc_id: Some(vector_doc_key().4), + node_deleted: true, }; core.apply_undo_vector(TID, 0, undo).unwrap(); @@ -839,5 +888,120 @@ mod tests { Some(vector_id), "vector_doc_map entry must be restored so a later delete can find the vector again" ); + let coll = core + .vector_collections + .get(&vector_index_key()) + .expect("the index stays"); + assert_eq!( + coll.local_for_surrogate(nodedb_types::Surrogate::new(1)), + Some(vector_id), + "the restored node must be bound to its row again, or a search hit on it \ + resolves to no row" + ); + } + + /// A delete that found its recorded node dead already restores only the + /// `vector_doc_map` entry. The node stays deleted. + #[test] + fn vector_delete_undo_leaves_a_node_that_was_dead_before() { + let dir = tempfile::tempdir().unwrap(); + let (mut core, _tx, _rx) = make_core_with_dir(dir.path()); + + let index_key = vector_index_key(); + let coll = core + .vector_collections + .entry(index_key.clone()) + .or_insert_with(|| nodedb_vector::VectorCollection::new(2, Default::default())); + let vector_id = coll + .insert_with_surrogate(vec![3.0, 4.0], nodedb_types::Surrogate::new(1)) + .unwrap(); + coll.delete(vector_id); + + let undo = UndoEntry::DeleteVector { + index_key, + vector_id, + collection: "c".to_string(), + field: "emb".to_string(), + doc_id: Some(vector_doc_key().4), + node_deleted: false, + }; + core.apply_undo_vector(TID, 0, undo).unwrap(); + + assert_eq!( + core.vector_doc_map.get(&vector_doc_key()).copied(), + Some(vector_id) + ); + let coll = core + .vector_collections + .get(&vector_index_key()) + .expect("the index stays"); + assert_eq!( + coll.live_count(), + 0, + "a node dead before the delete stays dead" + ); + assert_eq!( + coll.local_for_surrogate(nodedb_types::Surrogate::new(1)), + None + ); + } + + /// Undo of a displaced node un-deletes it and binds it to its surrogate, + /// and writes no `vector_doc_map` entry. + #[test] + fn displaced_vector_undo_restores_the_node_and_its_binding() { + let dir = tempfile::tempdir().unwrap(); + let (mut core, _tx, _rx) = make_core_with_dir(dir.path()); + + let index_key = vector_index_key(); + let surrogate = nodedb_types::Surrogate::new(1); + let coll = core + .vector_collections + .entry(index_key.clone()) + .or_insert_with(|| nodedb_vector::VectorCollection::new(2, Default::default())); + let vector_id = coll + .insert_with_surrogate(vec![3.0, 4.0], surrogate) + .unwrap(); + assert!(coll.delete(vector_id)); + + let undo = UndoEntry::DisplacedVector { + index_key, + vector_id, + surrogate, + }; + core.apply_undo_vector(TID, 0, undo).unwrap(); + + let coll = core + .vector_collections + .get(&vector_index_key()) + .expect("the index stays"); + assert_eq!(coll.live_count(), 1); + assert_eq!(coll.local_for_surrogate(surrogate), Some(vector_id)); + assert!(core.vector_doc_map.is_empty()); + } + + /// A displaced node that carries no tombstone at undo time leaves the + /// core's state unknown, so the undo fails. + #[test] + fn displaced_vector_undo_of_a_live_node_fails() { + let dir = tempfile::tempdir().unwrap(); + let (mut core, _tx, _rx) = make_core_with_dir(dir.path()); + + let index_key = vector_index_key(); + let surrogate = nodedb_types::Surrogate::new(1); + let coll = core + .vector_collections + .entry(index_key.clone()) + .or_insert_with(|| nodedb_vector::VectorCollection::new(2, Default::default())); + let vector_id = coll + .insert_with_surrogate(vec![3.0, 4.0], surrogate) + .unwrap(); + + let undo = UndoEntry::DisplacedVector { + index_key, + vector_id, + surrogate, + }; + assert!(core.apply_undo_vector(TID, 3, undo).is_err()); } } diff --git a/nodedb/src/data/executor/handlers/transaction/undo/crdt_collection.rs b/nodedb/src/data/executor/handlers/transaction/undo/crdt_collection.rs index b27d0cb9d..45ace6430 100644 --- a/nodedb/src/data/executor/handlers/transaction/undo/crdt_collection.rs +++ b/nodedb/src/data/executor/handlers/transaction/undo/crdt_collection.rs @@ -11,7 +11,7 @@ use crate::data::executor::core_loop::CoreLoop; use crate::types::{DatabaseId, TenantId}; -use super::UndoEntry; +use super::{UndoEntry, UndoError}; /// The Loro state of one CRDT collection before a write. pub(in crate::data::executor) struct CrdtCollectionUndo { @@ -51,14 +51,14 @@ impl CoreLoop { &mut self, entry_index: usize, undo: CrdtCollectionUndo, - ) -> Result<(), (usize, String)> { + ) -> Result<(), UndoError> { let key = (undo.database_id, undo.tenant_id); if !undo.engine_existed { self.crdt_engines.remove(&key); return Ok(()); } let Some(engine) = self.crdt_engines.get_mut(&key) else { - return Err(( + return Err(UndoError::mismatch( entry_index, format!( "the CRDT engine of '{}' vanished before its write was rolled back", @@ -69,9 +69,10 @@ impl CoreLoop { engine .restore_collection_snapshot(&undo.collection, undo.snapshot.as_deref()) .map_err(|e| { - ( + UndoError::failed( entry_index, - format!("restoring the CRDT collection '{}': {e}", undo.collection), + format!("restoring the CRDT collection '{}'", undo.collection), + e, ) }) } diff --git a/nodedb/src/data/executor/handlers/transaction/undo/crdt_row.rs b/nodedb/src/data/executor/handlers/transaction/undo/crdt_row.rs new file mode 100644 index 000000000..58fd95f71 --- /dev/null +++ b/nodedb/src/data/executor/handlers/transaction/undo/crdt_row.rs @@ -0,0 +1,80 @@ +// SPDX-License-Identifier: BUSL-1.1 + +//! Undo of one CRDT scalar write to a document row. +//! +//! A scalar write changes only the row's scalar fields. The pre-image is +//! those fields, and the undo writes them back with new Loro operations. +//! This costs one row, where the collection snapshot costs the whole +//! collection. + +use nodedb_crdt::state::RowImage; + +use crate::data::executor::core_loop::CoreLoop; +use crate::types::{DatabaseId, TenantId}; + +use super::{UndoEntry, UndoError}; + +/// The scalar state of one CRDT document row before a write. +pub(in crate::data::executor) struct CrdtRowUndo { + pub database_id: DatabaseId, + pub tenant_id: TenantId, + pub collection: String, + pub row_id: String, + /// The row before the write. `None`: the collection had no Loro state. + pub image: Option, +} + +impl CoreLoop { + /// Capture the state of `row_id` that a CRDT scalar write can change. + /// Opens the tenant's CRDT engine when it is not open. + pub(in crate::data::executor) fn capture_crdt_row_undo( + &mut self, + database_id: DatabaseId, + tenant_id: TenantId, + collection: &str, + row_id: &str, + ) -> crate::Result { + let image = self + .get_crdt_engine(database_id, tenant_id)? + .doc_row_image(collection, row_id)?; + Ok(UndoEntry::CrdtRow(Box::new(CrdtRowUndo { + database_id, + tenant_id, + collection: collection.to_string(), + row_id: row_id.to_string(), + image, + }))) + } + + /// Put the row's scalar fields back. + pub(super) fn apply_undo_crdt_row( + &mut self, + entry_index: usize, + undo: CrdtRowUndo, + ) -> Result<(), UndoError> { + let Some(engine) = self + .crdt_engines + .get_mut(&(undo.database_id, undo.tenant_id)) + else { + return Err(UndoError::mismatch( + entry_index, + format!( + "the CRDT engine of '{}' vanished before its row '{}' was rolled back", + undo.collection, undo.row_id + ), + )); + }; + engine + .restore_doc_row(&undo.collection, &undo.row_id, undo.image.as_ref()) + .map_err(|e| { + UndoError::failed( + entry_index, + format!( + "restoring the CRDT row '{}' of '{}'", + undo.row_id, undo.collection + ), + e, + ) + }) + } +} diff --git a/nodedb/src/data/executor/handlers/transaction/undo/document.rs b/nodedb/src/data/executor/handlers/transaction/undo/document.rs index d6f67bfa3..3f57ca1a3 100644 --- a/nodedb/src/data/executor/handlers/transaction/undo/document.rs +++ b/nodedb/src/data/executor/handlers/transaction/undo/document.rs @@ -3,8 +3,8 @@ //! Document-engine undo entry application logic. //! //! `apply_undo_document` handles document-engine undo entries. All methods -//! return `Err((entry_index, detail))` on fatal failure so the caller can -//! escalate to a typed `RollbackFailed` response. +//! return an [`UndoError`] on fatal failure, and the caller escalates it to +//! a typed `RollbackFailed` response. use nodedb_types::StorageKey; use tracing::error; @@ -12,7 +12,7 @@ use tracing::error; use crate::data::executor::core_loop::CoreLoop; use crate::engine::sparse::btree_versioned::VersionedIndexEntry; -use super::UndoEntry; +use super::{UndoEntry, UndoError}; #[derive(Clone, Copy)] pub(super) struct UndoDocumentContext<'a> { @@ -32,7 +32,7 @@ impl CoreLoop { tid: u64, entry_index: usize, entry: UndoEntry, - ) -> Result<(), (usize, String)> { + ) -> Result<(), UndoError> { match entry { UndoEntry::PutDocument { collection, @@ -65,26 +65,26 @@ impl CoreLoop { self.sparse .put(database_id, tid, &collection, &storage_key, &old) .map(|_| ()) - .map_err(|e| e.to_string()) } else { self.sparse .delete(database_id, tid, &collection, &storage_key) .map(|_| ()) - .map_err(|e| e.to_string()) }; result.map_err(|e| { + let err = UndoError::failed( + entry_index, + format!("document restore on {collection}/{document_id}"), + e, + ); error!( core = self.core_id, entry_index, collection = %collection, document_id = %document_id, - error = %e, + error = %err, "transaction undo: document restore failed; shard state unknown" ); - ( - entry_index, - format!("document restore on {collection}/{document_id}: {e}"), - ) + err })?; } // Reverse plain secondary-index mutations: undo the inserts and @@ -104,18 +104,20 @@ impl CoreLoop { surrogate, ) .map_err(|e| { + let err = UndoError::failed( + entry_index, + format!("fts posting removal on {collection}/{document_id}"), + e, + ); error!( core = self.core_id, entry_index, collection = %collection, document_id = %document_id, - error = %e, + error = %err, "transaction undo: FTS posting removal failed; shard state unknown" ); - ( - entry_index, - format!("fts posting removal on {collection}/{document_id}: {e}"), - ) + err })?; // Evict any cached copy of the reversed document. Always safe: // a stale hit would otherwise resurrect a rolled-back put; the @@ -150,18 +152,20 @@ impl CoreLoop { .put(database_id, tid, &collection, &storage_key, &old_value) .map(|_| ()) .map_err(|e| { + let err = UndoError::failed( + entry_index, + format!("document re-insert on {collection}/{document_id}"), + e, + ); error!( core = self.core_id, entry_index, collection = %collection, document_id = %document_id, - error = %e, + error = %err, "transaction undo: document re-insert failed; shard state unknown" ); - ( - entry_index, - format!("document re-insert on {collection}/{document_id}: {e}"), - ) + err })?; } // Restore the plain secondary-index entries the forward delete @@ -196,7 +200,7 @@ impl CoreLoop { ctx: UndoDocumentContext<'_>, sys_from_ms: i64, index_tuples: &[(String, String)], - ) -> Result<(), (usize, String)> { + ) -> Result<(), UndoError> { let UndoDocumentContext { database_id, tid, @@ -204,28 +208,29 @@ impl CoreLoop { collection, document_id, } = ctx; - let map_err = |stage: &str, e: String| { + let map_err = |stage: &str, e: crate::Error| { + let err = UndoError::failed( + entry_index, + format!("bitemporal {stage} on {collection}/{document_id}"), + e, + ); error!( core = self.core_id, entry_index, collection = %collection, document_id = %document_id, - error = %e, + error = %err, "transaction undo: bitemporal version removal failed; shard state unknown" ); - ( - entry_index, - format!("bitemporal {stage} on {collection}/{document_id}: {e}"), - ) + err }; let txn = self .sparse - .db() .begin_write() - .map_err(|e| map_err("begin_write", e.to_string()))?; + .map_err(|e| map_err("begin_write", e))?; self.sparse .versioned_remove_in_txn(&txn, database_id, tid, collection, document_id, sys_from_ms) - .map_err(|e| map_err("version remove", e.to_string()))?; + .map_err(|e| map_err("version remove", e))?; for (field, value) in index_tuples { self.sparse .versioned_index_remove_in_txn( @@ -240,9 +245,17 @@ impl CoreLoop { sys_from_ms, }, ) - .map_err(|e| map_err("index remove", e.to_string()))?; + .map_err(|e| map_err("index remove", e))?; } - txn.commit().map_err(|e| map_err("commit", e.to_string()))?; + txn.commit().map_err(|e| { + map_err( + "commit", + crate::Error::Storage { + engine: "sparse".into(), + detail: e.to_string(), + }, + ) + })?; Ok(()) } @@ -259,7 +272,7 @@ impl CoreLoop { ctx: UndoDocumentContext<'_>, to_remove: &[(String, String)], to_restore: &[(String, String)], - ) -> Result<(), (usize, String)> { + ) -> Result<(), UndoError> { let UndoDocumentContext { database_id, tid, @@ -267,29 +280,31 @@ impl CoreLoop { collection, document_id, } = ctx; - let map_err = |stage: &str, e: String| { + let map_err = |stage: &str, e: crate::Error| { + let err = UndoError::failed( + entry_index, + format!("secondary-index {stage} on {collection}/{document_id}"), + e, + ); error!( core = self.core_id, entry_index, collection = %collection, document_id = %document_id, - error = %e, + error = %err, "transaction undo: secondary-index reversal failed; shard state unknown" ); - ( - entry_index, - format!("secondary-index {stage} on {collection}/{document_id}: {e}"), - ) + err }; for (field, value) in to_remove { self.sparse .index_remove(database_id, tid, collection, field, value, document_id) - .map_err(|e| map_err("remove", e.to_string()))?; + .map_err(|e| map_err("remove", e))?; } for (field, value) in to_restore { self.sparse .index_put(database_id, tid, collection, field, value, document_id) - .map_err(|e| map_err("restore", e.to_string()))?; + .map_err(|e| map_err("restore", e))?; } Ok(()) } @@ -311,7 +326,7 @@ impl CoreLoop { collection: &str, entry_index: usize, chain_hash_prior: Option>, - ) -> Result<(), (usize, String)> { + ) -> Result<(), UndoError> { let key = ( crate::types::DatabaseId::new(database_id), crate::types::TenantId::new(tid), @@ -330,17 +345,19 @@ impl CoreLoop { } }; persisted.map_err(|e| { + let err = UndoError::failed( + entry_index, + format!("hash-chain head restore on {collection}"), + e, + ); error!( core = self.core_id, entry_index, collection = %collection, - error = %e, + error = %err, "transaction undo: hash-chain head restore failed; shard state unknown" ); - ( - entry_index, - format!("hash-chain head restore on {collection}: {e}"), - ) + err }) } } diff --git a/nodedb/src/data/executor/handlers/transaction/undo/document_fts.rs b/nodedb/src/data/executor/handlers/transaction/undo/document_fts.rs index 67414169a..502173d05 100644 --- a/nodedb/src/data/executor/handlers/transaction/undo/document_fts.rs +++ b/nodedb/src/data/executor/handlers/transaction/undo/document_fts.rs @@ -7,28 +7,29 @@ //! rollback restores the document body into the primary store, so it must also //! recompute and re-insert the FTS postings — otherwise the row comes back //! restored-but-unsearchable. `nodedb_fts::analyze` is deterministic, so the -//! recomputed text (extracted via the same [`extract_fts_text`] helper the -//! forward PUT path uses) reproduces byte-identical postings. +//! recomputed text (extracted via the same [`extract_fts_fields`] helper the +//! forward PUT path uses) reproduces byte-identical postings in every index. use tracing::error; use crate::data::executor::core_loop::CoreLoop; -use crate::data::executor::fts_text::extract_fts_text; +use crate::data::executor::fts_text::extract_fts_fields; +use super::UndoError; use super::document::UndoDocumentContext; impl CoreLoop { /// Re-index a restored document's text into the inverted index during /// DELETE rollback. Decodes the restored body through the storage-mode-aware /// helper (strict → Binary Tuple, schemaless → MessagePack) so both modes - /// recompute their real text. Returns `Err((entry_index, detail))` on - /// failure so a partial FTS restore escalates to `RollbackFailed`. + /// recompute their real text. Returns an [`UndoError`] on failure, so a + /// partial FTS restore escalates to `RollbackFailed`. pub(super) fn reindex_restored_document_fts( &self, ctx: UndoDocumentContext<'_>, surrogate: nodedb_types::Surrogate, old_value: &[u8], - ) -> Result<(), (usize, String)> { + ) -> Result<(), UndoError> { let UndoDocumentContext { database_id, tid, @@ -52,8 +53,14 @@ impl CoreLoop { // unsearchable. That is a failed rollback, not a no-op. let doc = self .decode_stored_document(config, old_value) - .map_err(|e| (entry_index, e.to_string()))?; - let text = extract_fts_text(&doc); + .map_err(|e| { + UndoError::failed( + entry_index, + format!("decoding the restored body of {collection}/{document_id}"), + e, + ) + })?; + let text = extract_fts_fields(&doc); if text.is_empty() { return Ok(()); } @@ -66,18 +73,20 @@ impl CoreLoop { &text, ) .map_err(|e| { + let err = UndoError::failed( + entry_index, + format!("fts re-index on {collection}/{document_id}"), + e, + ); error!( core = self.core_id, entry_index, collection = %collection, document_id = %document_id, - error = %e, + error = %err, "transaction undo: FTS re-index failed; shard state unknown" ); - ( - entry_index, - format!("fts re-index on {collection}/{document_id}: {e}"), - ) + err }) } } diff --git a/nodedb/src/data/executor/handlers/transaction/undo/document_outcome.rs b/nodedb/src/data/executor/handlers/transaction/undo/document_outcome.rs index a16b45687..6e931263e 100644 --- a/nodedb/src/data/executor/handlers/transaction/undo/document_outcome.rs +++ b/nodedb/src/data/executor/handlers/transaction/undo/document_outcome.rs @@ -22,8 +22,6 @@ use super::UndoEntry; /// The document row one write touched. pub(in crate::data::executor::handlers) struct DocumentRow<'a> { - pub database_id: u64, - pub tid: u64, pub collection: &'a str, pub storage_key: StorageKey, } @@ -31,12 +29,14 @@ pub(in crate::data::executor::handlers) struct DocumentRow<'a> { /// Push the undo entries that reverse every materialized-sum target row the /// write folded into. A target write is a full document write, so it has /// index, vector, spatial and stats side effects of its own. The targets stay -/// with the caller, which also reports them in its response. +/// with the caller, which also reports them in its response. Their in-memory +/// undo entries move onto `undo_log`. pub(in crate::data::executor::handlers) fn push_target_undo( undo_log: &mut Vec, - targets: &[TargetWrite], + targets: &mut [TargetWrite], ) { for target in targets { + let memory_undo = std::mem::take(&mut target.outcome.memory_undo); let outcome = &target.outcome; undo_log.push(UndoEntry::PutDocument { collection: target.collection.clone(), @@ -48,12 +48,7 @@ pub(in crate::data::executor::handlers) fn push_target_undo( secondary_index_removed: outcome.secondary_index_removed.clone(), chain_hash_prior: None, }); - push_put_side_effects( - undo_log, - outcome.vector_inserts.clone(), - outcome.spatial_inserts.clone(), - outcome.stats_prior.clone(), - ); + push_put_side_effects(undo_log, memory_undo, outcome.stats_prior.clone()); } } @@ -76,12 +71,7 @@ pub(in crate::data::executor::handlers) fn push_put_undo( secondary_index_removed: outcome.secondary_index_removed, chain_hash_prior, }); - push_put_side_effects( - undo_log, - outcome.vector_inserts, - outcome.spatial_inserts, - outcome.stats_prior, - ); + push_put_side_effects(undo_log, outcome.memory_undo, outcome.stats_prior); } /// Push the undo entries that reverse one document delete. A delete that @@ -102,53 +92,15 @@ pub(in crate::data::executor::handlers) fn push_delete_undo( chain_hash_prior: None, }); } - for delta in outcome.vector_deletes { - undo_log.push(UndoEntry::DeleteVector { - index_key: delta.index_key, - vector_id: delta.vector_id, - collection: delta.collection, - field: delta.field, - doc_id: Some(delta.doc_id), - }); - } - for (key, entry_id, bbox, document_id) in outcome.spatial_deletes { - undo_log.push(UndoEntry::SpatialDelete { - key, - entry_id, - bbox, - document_id, - }); - } - // `Some` only when this delete newly marked the node: a tombstone a prior - // committed write left is never un-marked. - if let Some(node_id) = outcome.mark_node_deleted { - undo_log.push(UndoEntry::MarkNodeDeleted { - database_id: row.database_id, - tid: row.tid, - collection: row.collection.to_string(), - node_id, - }); - } + undo_log.extend(outcome.memory_undo); } fn push_put_side_effects( undo_log: &mut Vec, - vector_inserts: Vec, - spatial_inserts: Vec<(crate::data::executor::spatial_key::SpatialIndexKey, u64)>, + memory_undo: Vec, stats_prior: Vec, ) { - for delta in vector_inserts { - undo_log.push(UndoEntry::InsertVector { - index_key: delta.index_key, - vector_id: delta.vector_id, - collection: delta.collection, - field: delta.field, - doc_id: Some(delta.doc_id), - }); - } - for (key, entry_id) in spatial_inserts { - undo_log.push(UndoEntry::SpatialInsert { key, entry_id }); - } + undo_log.extend(memory_undo); for (key, prior) in stats_prior { undo_log.push(UndoEntry::StatsRestore { key, prior }); } diff --git a/nodedb/src/data/executor/handlers/transaction/undo/edge_cut.rs b/nodedb/src/data/executor/handlers/transaction/undo/edge_cut.rs index 8ce4062b8..73581c74b 100644 --- a/nodedb/src/data/executor/handlers/transaction/undo/edge_cut.rs +++ b/nodedb/src/data/executor/handlers/transaction/undo/edge_cut.rs @@ -8,6 +8,7 @@ use tracing::error; +use super::UndoError; use crate::data::executor::core_loop::CoreLoop; use crate::engine::graph::edge_store::EdgeCutInstall; use crate::types::{DatabaseId, TenantId}; @@ -25,29 +26,33 @@ impl CoreLoop { &mut self, entry_index: usize, undo: EdgeCutUndo, - ) -> Result<(), (usize, String)> { + ) -> Result<(), UndoError> { let EdgeCutUndo { database_id, tid, install, } = undo; let core = self.core_id; - let fail = |detail: String| { + let fail = |action: String, cause: crate::Error| { + let err = UndoError::failed(entry_index, action, cause); error!( core, entry_index, - error = %detail, + error = %err, "transaction undo: edge cut rollback failed; shard state unknown" ); - (entry_index, detail) + err }; self.edge_store .remove_edge_cut(DatabaseId::new(database_id), TenantId::new(tid), &install) .map_err(|e| { - fail(format!( - "removing the cut of '{}' at {}: {e}", - install.collection, install.cut - )) + fail( + format!( + "removing the cut of '{}' at {}", + install.collection, install.cut + ), + e, + ) })?; for flip in &install.flips { self.mirror_edge_csr( @@ -58,10 +63,13 @@ impl CoreLoop { flip.before.as_deref(), ) .map_err(|e| { - fail(format!( - "restoring the CSR edge {} {}-[{}]->{}: {e}", - install.collection, flip.src, flip.label, flip.dst - )) + fail( + format!( + "restoring the CSR edge {} {}-[{}]->{}", + install.collection, flip.src, flip.label, flip.dst + ), + e.into(), + ) })?; } Ok(()) diff --git a/nodedb/src/data/executor/handlers/transaction/undo/edge_write.rs b/nodedb/src/data/executor/handlers/transaction/undo/edge_write.rs index 45dabb546..2c6912b72 100644 --- a/nodedb/src/data/executor/handlers/transaction/undo/edge_write.rs +++ b/nodedb/src/data/executor/handlers/transaction/undo/edge_write.rs @@ -10,6 +10,7 @@ use tracing::error; +use super::UndoError; use crate::data::executor::core_loop::CoreLoop; use crate::engine::graph::edge_store::{EdgeRef, EdgeVersionWrite}; use crate::types::{DatabaseId, TenantId}; @@ -105,11 +106,11 @@ impl CoreLoop { } /// Reverse one edge write. - pub(super) fn apply_undo_edge_write( + pub(in crate::data::executor) fn apply_undo_edge_write( &mut self, entry_index: usize, undo: EdgeWriteUndo, - ) -> Result<(), (usize, String)> { + ) -> Result<(), UndoError> { let EdgeWriteUndo { database_id, tid, @@ -121,14 +122,15 @@ impl CoreLoop { csr, } = undo; let core = self.core_id; - let fail = |detail: String| { + let fail = |action: String, cause: crate::Error| { + let err = UndoError::failed(entry_index, action, cause); error!( core, entry_index, - error = %detail, + error = %err, "transaction undo: edge rollback failed; shard state unknown" ); - (entry_index, detail) + err }; let edge_name = format!("{collection} {src_id}-[{label}]->{dst_id}"); let database = DatabaseId::new(database_id); @@ -139,19 +141,19 @@ impl CoreLoop { EdgeRef::new(database, tenant, &collection, &src_id, &label, &dst_id), &version, ) - .map_err(|e| fail(format!("removing the version of {edge_name}: {e}")))?; + .map_err(|e| fail(format!("removing the version of {edge_name}"), e))?; let partition = self.csr_partition_mut(database_id, tid); partition .restore_edge_in_collection(&src_id, &label, &dst_id, &collection, csr.weight) - .map_err(|e| fail(format!("restoring the CSR edge {edge_name}: {e}")))?; + .map_err(|e| fail(format!("restoring the CSR edge {edge_name}"), e.into()))?; for (node, prior) in &csr.surrogates { partition.restore_node_surrogate(node, *prior); } for node in csr.created_nodes.iter().rev() { partition .withdraw_newest_node(node) - .map_err(|e| fail(format!("withdrawing node '{node}': {e}")))?; + .map_err(|e| fail(format!("withdrawing node '{node}'"), e.into()))?; } Ok(()) } @@ -303,7 +305,7 @@ mod tests { assert_eq!(resolve(&core, "alice", "bob", 250), Some(weighted(2.5))); assert_eq!( core.csr_partition(DB, TID) - .map(|p| p.neighbors("alice", None, Direction::Out)), + .map(|p| p.neighbors("alice", &[], Direction::Out)), Some(vec![("KNOWS".to_string(), "bob".to_string())]) ); } diff --git a/nodedb/src/data/executor/handlers/transaction/undo/entry.rs b/nodedb/src/data/executor/handlers/transaction/undo/entry.rs index b3b7ab0e3..43fc796e6 100644 --- a/nodedb/src/data/executor/handlers/transaction/undo/entry.rs +++ b/nodedb/src/data/executor/handlers/transaction/undo/entry.rs @@ -156,6 +156,12 @@ pub(in crate::data::executor) enum UndoEntry { field: String, doc_id: Option, }, + /// Undo the creation of a vector index by a document write: the index did + /// not exist before the write. An index left behind fixes the vector + /// width for every later write, so undo removes it. + VectorCollectionCreated { + index_key: (nodedb_types::DatabaseId, TenantId, String), + }, /// Undo a VectorDelete by un-deleting (clearing tombstone) and restoring /// the `vector_doc_map` entry the forward delete removed — mirroring /// `SpatialDelete`'s reverse-map restore. Without this, a rolled-back @@ -172,6 +178,19 @@ pub(in crate::data::executor) enum UndoEntry { collection: String, field: String, doc_id: Option, + /// Whether the forward delete tombstoned `vector_id`. `false` when + /// the node carried a tombstone already: undo then restores only the + /// `vector_doc_map` entry and leaves the node deleted. + node_deleted: bool, + }, + /// Undo the displacement of a node bound to a document's surrogate that + /// no `vector_doc_map` entry recorded, such as a node a direct vector + /// write bound. A document put soft-deletes it before it binds its own + /// node. Undo un-deletes it and binds it to `surrogate` again. + DisplacedVector { + index_key: (nodedb_types::DatabaseId, TenantId, String), + vector_id: u32, + surrogate: nodedb_types::Surrogate, }, /// Undo a spatial R-tree insert by removing the entry from the per-field /// R-tree and deleting its reverse `spatial_doc_map` record. @@ -256,6 +275,9 @@ pub(in crate::data::executor) enum UndoEntry { }, /// Undo a CRDT write by putting the collection's Loro document back. CrdtCollection(Box), + /// Undo a CRDT scalar write to one document row by putting the row's + /// scalar fields back. + CrdtRow(Box), /// Undo an array cell write by putting back the memtable tiles it /// touched. ArrayTiles { diff --git a/nodedb/src/data/executor/handlers/transaction/undo/error.rs b/nodedb/src/data/executor/handlers/transaction/undo/error.rs new file mode 100644 index 000000000..c2015806a --- /dev/null +++ b/nodedb/src/data/executor/handlers/transaction/undo/error.rs @@ -0,0 +1,116 @@ +// SPDX-License-Identifier: BUSL-1.1 + +//! The typed error of an undo entry whose reverse write did not apply. + +use crate::bridge::envelope::ErrorCode; + +/// An undo entry whose reverse write did not apply. +/// +/// The core's state is unknown after it. The caller answers +/// `ErrorCode::RollbackFailed`, and the core fail-stops when that response +/// leaves it. +#[derive(Debug, thiserror::Error)] +#[error( + "undo entry {entry_index}: {action}{}", + .cause.as_ref().map(|c| format!(": {c}")).unwrap_or_default() +)] +pub struct UndoError { + /// Forward-order position of the entry in the undo log. + pub entry_index: usize, + /// The reverse write that failed, with the object it targets. + pub action: String, + /// The error the reverse write returned. `None` when the engine state + /// does not match the entry, so no reverse write ran. Boxed to keep the + /// `Result` of every undo function small. + #[source] + pub cause: Option>, +} + +impl UndoError { + /// A reverse write that returned `cause`. + pub fn failed( + entry_index: usize, + action: impl Into, + cause: impl Into, + ) -> Self { + Self { + entry_index, + action: action.into(), + cause: Some(Box::new(cause.into())), + } + } + + /// An entry the engine state does not match. No reverse write ran. + pub fn mismatch(entry_index: usize, action: impl Into) -> Self { + Self { + entry_index, + action: action.into(), + cause: None, + } + } +} + +impl From for ErrorCode { + fn from(e: UndoError) -> Self { + ErrorCode::RollbackFailed { + entry_index: e.entry_index, + detail: e.action, + cause: e.cause.map(|cause| Box::new(ErrorCode::from(*cause))), + } + } +} + +impl From for crate::Error { + fn from(e: UndoError) -> Self { + crate::Error::DataPlane(ErrorCode::from(e)) + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn a_failed_reverse_write_keeps_its_typed_cause() { + let cause = crate::Error::Storage { + engine: "sparse".into(), + detail: "commit".into(), + }; + let code = ErrorCode::from(UndoError::failed(3, "restoring row r1", cause)); + match code { + ErrorCode::RollbackFailed { + entry_index, + detail, + cause: Some(cause), + } => { + assert_eq!(entry_index, 3); + assert_eq!(detail, "restoring row r1"); + assert!(matches!(*cause, ErrorCode::Internal { .. })); + } + other => panic!("expected RollbackFailed with a cause, got {other:?}"), + } + } + + #[test] + fn a_state_mismatch_carries_no_cause() { + let code = ErrorCode::from(UndoError::mismatch(0, "vector index missing")); + assert!(matches!( + code, + ErrorCode::RollbackFailed { + entry_index: 0, + cause: None, + .. + } + )); + } + + #[test] + fn a_data_plane_cause_keeps_its_code() { + let cause = crate::Error::DataPlane(ErrorCode::DivisionByZero); + let code = ErrorCode::from(UndoError::failed(1, "x", cause)); + assert!(matches!( + code, + ErrorCode::RollbackFailed { cause: Some(c), .. } if *c == ErrorCode::DivisionByZero + )); + } +} diff --git a/nodedb/src/data/executor/handlers/transaction/undo/fts_doc.rs b/nodedb/src/data/executor/handlers/transaction/undo/fts_doc.rs index 4b8a1b5e9..9a19a5eab 100644 --- a/nodedb/src/data/executor/handlers/transaction/undo/fts_doc.rs +++ b/nodedb/src/data/executor/handlers/transaction/undo/fts_doc.rs @@ -14,7 +14,7 @@ use crate::data::executor::core_loop::CoreLoop; use crate::engine::sparse::inverted::FtsDocImage; use crate::types::{DatabaseId, TenantId}; -use super::UndoEntry; +use super::{UndoEntry, UndoError}; /// The pre-image of one document's index footprint. pub(in crate::data::executor) struct FtsDocUndo { @@ -59,7 +59,7 @@ impl CoreLoop { &mut self, entry_index: usize, undo: FtsDocUndo, - ) -> Result<(), (usize, String)> { + ) -> Result<(), UndoError> { self.inverted .restore_document_image( undo.database_id, @@ -69,13 +69,14 @@ impl CoreLoop { undo.prior.as_ref(), ) .map_err(|e| { - ( + UndoError::failed( entry_index, format!( - "restoring the full-text footprint of {} in '{}': {e}", + "restoring the full-text footprint of {} in '{}'", undo.surrogate.as_u32(), undo.collection ), + e, ) }) } diff --git a/nodedb/src/data/executor/handlers/transaction/undo/graph_node.rs b/nodedb/src/data/executor/handlers/transaction/undo/graph_node.rs index 892c01e48..90eb2b1ee 100644 --- a/nodedb/src/data/executor/handlers/transaction/undo/graph_node.rs +++ b/nodedb/src/data/executor/handlers/transaction/undo/graph_node.rs @@ -16,19 +16,19 @@ //! The node-label undo puts each label back, then withdraws the label names //! and the node the op interned, so the CSR holds what it held before. //! -//! Returns `Err((entry_index, detail))` on fatal failure so the caller can -//! escalate to a typed `RollbackFailed` response. +//! Returns an [`UndoError`] on fatal failure, and the caller escalates it to +//! a typed `RollbackFailed` response. use crate::data::executor::core_loop::CoreLoop; -use super::UndoEntry; +use super::{UndoEntry, UndoError}; impl CoreLoop { pub(super) fn apply_undo_mark_node( &mut self, _entry_index: usize, entry: UndoEntry, - ) -> Result<(), (usize, String)> { + ) -> Result<(), UndoError> { match entry { UndoEntry::MarkNodeDeleted { database_id, @@ -87,7 +87,7 @@ impl CoreLoop { &mut self, entry_index: usize, undo: NodeLabelsUndo, - ) -> Result<(), (usize, String)> { + ) -> Result<(), UndoError> { let NodeLabelsUndo { database_id, tid, @@ -100,9 +100,10 @@ impl CoreLoop { for (label, carried) in prior { if carried { partition.add_node_label(&node_id, &label).map_err(|e| { - ( + UndoError::failed( entry_index, - format!("restoring label '{label}' on node '{node_id}': {e}"), + format!("restoring label '{label}' on node '{node_id}'"), + e, ) })?; } else { @@ -110,14 +111,14 @@ impl CoreLoop { } } for label in interned_labels.iter().rev() { - partition - .withdraw_newest_node_label(label) - .map_err(|e| (entry_index, format!("withdrawing label '{label}': {e}")))?; + partition.withdraw_newest_node_label(label).map_err(|e| { + UndoError::failed(entry_index, format!("withdrawing label '{label}'"), e) + })?; } if created_node { - partition - .withdraw_newest_node(&node_id) - .map_err(|e| (entry_index, format!("withdrawing node '{node_id}': {e}")))?; + partition.withdraw_newest_node(&node_id).map_err(|e| { + UndoError::failed(entry_index, format!("withdrawing node '{node_id}'"), e) + })?; } Ok(()) } diff --git a/nodedb/src/data/executor/handlers/transaction/undo/kv.rs b/nodedb/src/data/executor/handlers/transaction/undo/kv.rs index 3ad39d3c9..e864e8e72 100644 --- a/nodedb/src/data/executor/handlers/transaction/undo/kv.rs +++ b/nodedb/src/data/executor/handlers/transaction/undo/kv.rs @@ -5,7 +5,7 @@ use crate::data::executor::core_loop::CoreLoop; use crate::engine::kv::current_ms; -use super::UndoEntry; +use super::{UndoEntry, UndoError}; fn kv_key<'a>( did: u64, @@ -28,7 +28,7 @@ impl CoreLoop { tid: u64, entry_index: usize, entry: UndoEntry, - ) -> Result<(), (usize, String)> { + ) -> Result<(), UndoError> { match entry { UndoEntry::KvPut { collection, @@ -41,7 +41,7 @@ impl CoreLoop { prior.as_ref(), current_ms(), ) - .map_err(|e| (entry_index, e.to_string())), + .map_err(|e| UndoError::failed(entry_index, "reinstating a KV entry", e)), UndoEntry::KvDelete { collection, key, @@ -49,7 +49,7 @@ impl CoreLoop { } => self .kv_engine .restore_entry_image(kv_key(did, tid, &collection, &key), &prior, current_ms()) - .map_err(|e| (entry_index, e.to_string())), + .map_err(|e| UndoError::failed(entry_index, "restoring a deleted KV entry", e)), UndoEntry::KvTtl { collection, key, @@ -85,13 +85,19 @@ impl CoreLoop { }, now_ms, ) - .map_err(|e| (entry_index, e.to_string()))?; + .map_err(|e| { + UndoError::failed( + entry_index, + format!("restoring a truncated KV row of '{collection}'"), + e, + ) + })?; } Ok(()) } - _ => Err(( + _ => Err(UndoError::mismatch( entry_index, - "apply_undo_kv called with non-kv entry".to_string(), + "apply_undo_kv called with non-kv entry", )), } } diff --git a/nodedb/src/data/executor/handlers/transaction/undo/memory.rs b/nodedb/src/data/executor/handlers/transaction/undo/memory.rs new file mode 100644 index 000000000..c742c7a15 --- /dev/null +++ b/nodedb/src/data/executor/handlers/transaction/undo/memory.rs @@ -0,0 +1,75 @@ +// SPDX-License-Identifier: BUSL-1.1 + +//! Reverse the in-memory side effects of document writes whose redb +//! transaction drops uncommitted. +//! +//! A dropped redb transaction reverses the row, its btree and FTS entries +//! and its column stats. It does not reverse what the write did in memory: +//! R-tree and `spatial_doc_map` entries, vector nodes and their +//! `vector_doc_map` entries, sparse-vector postings, and the deleted-node +//! marks. `apply_point_put` and `apply_point_delete` return those as undo +//! entries, and an autocommit write that aborts reverses them here, through +//! the undo driver a rolled-back transaction uses. + +use crate::data::executor::core_loop::CoreLoop; +use crate::data::executor::enforcement::materialized_sum::apply::TargetWrite; +use crate::engine::document::store::StorageKey; + +use super::UndoEntry; + +impl CoreLoop { + /// Reverse `undo`, last entry first. + /// + /// An entry that does not reverse leaves the core's state unknown. The + /// result is then `RollbackFailed`, and the core fail-stops when that + /// error leaves it in a response. + pub(in crate::data::executor) fn undo_memory_effects( + &mut self, + database_id: u64, + tid: u64, + undo: Vec, + ) -> crate::Result<()> { + if undo.is_empty() { + return Ok(()); + } + self.rollback_undo_log(database_id, tid, undo) + .map_err(crate::Error::from) + } + + /// Abandon the target rows a materialized-sum pass wrote into a + /// transaction that drops uncommitted. + /// + /// Each target write is a full document write: it cached its row and can + /// have in-memory index entries of its own. This drops the cache entries + /// and reverses the entries, last target first. + pub(in crate::data::executor) fn abandon_target_writes( + &mut self, + database_id: u64, + tid: u64, + targets: Vec, + ) -> crate::Result<()> { + let mut undo = Vec::new(); + for target in targets { + let key = StorageKey::for_surrogate(target.surrogate); + self.doc_cache + .invalidate(database_id, tid, &target.collection, &key); + undo.extend(target.outcome.memory_undo); + } + self.undo_memory_effects(database_id, tid, undo) + } +} + +/// The error an abandoned write reports. +/// +/// This is `original`, unless reversing the write's in-memory effects +/// failed. That failure leaves the core's state unknown, so it outranks the +/// error that caused the abort. +pub(in crate::data::executor) fn abort_error( + original: crate::Error, + undo: crate::Result<()>, +) -> crate::Error { + match undo { + Ok(()) => original, + Err(fatal) => fatal, + } +} diff --git a/nodedb/src/data/executor/handlers/transaction/undo/mod.rs b/nodedb/src/data/executor/handlers/transaction/undo/mod.rs index 3c4b05c79..e20a75644 100644 --- a/nodedb/src/data/executor/handlers/transaction/undo/mod.rs +++ b/nodedb/src/data/executor/handlers/transaction/undo/mod.rs @@ -5,15 +5,18 @@ pub(super) mod apply; pub(super) mod columnar_insert; pub(in crate::data::executor) mod crdt_collection; +pub(in crate::data::executor) mod crdt_row; pub(super) mod document; pub(super) mod document_fts; pub(in crate::data::executor::handlers) mod document_outcome; pub(in crate::data::executor) mod edge_cut; pub(in crate::data::executor) mod edge_write; pub(super) mod entry; +pub(in crate::data::executor) mod error; pub(in crate::data::executor) mod fts_doc; pub(super) mod graph_node; pub(super) mod kv; +pub(in crate::data::executor) mod memory; pub(super) mod rollback; pub(super) mod spatial; pub(in crate::data::executor) mod spatial_row; @@ -27,3 +30,4 @@ pub(in crate::data::executor) mod vector_write; pub(in crate::data::executor) use entry::{ ColumnarTruncateUndo, TimeseriesIngestUndo, TimeseriesTruncateUndo, UndoEntry, }; +pub(in crate::data::executor) use error::UndoError; diff --git a/nodedb/src/data/executor/handlers/transaction/undo/rollback.rs b/nodedb/src/data/executor/handlers/transaction/undo/rollback.rs index 46ddd4f37..3c0670c6e 100644 --- a/nodedb/src/data/executor/handlers/transaction/undo/rollback.rs +++ b/nodedb/src/data/executor/handlers/transaction/undo/rollback.rs @@ -2,8 +2,8 @@ //! Rollback driver for the undo log. -use super::UndoEntry; use super::graph_node::NodeLabelsUndo; +use super::{UndoEntry, UndoError}; use crate::data::executor::core_loop::CoreLoop; impl CoreLoop { @@ -11,18 +11,17 @@ impl CoreLoop { /// /// Returns `Ok(())` if all undo entries were applied successfully. /// - /// Returns `Err((entry_index, detail))` on the first undo failure — - /// the entry index is the original forward-order position of the failed - /// entry (before reversal). On failure the caller **must** return a - /// `RollbackFailed` error. The core's state is then unknown: the core - /// fail-stops when that response leaves it, and a restart rebuilds the - /// state through WAL replay. - pub(in crate::data::executor::handlers) fn rollback_undo_log( + /// Returns the [`UndoError`] of the first entry that does not reverse. + /// Its index is the forward-order position of that entry. On error the + /// caller **must** answer `RollbackFailed`. The core's state is then + /// unknown: the core fail-stops when that response leaves it, and a + /// restart rebuilds the state through WAL replay. + pub(in crate::data::executor) fn rollback_undo_log( &mut self, did: u64, tid: u64, undo_log: Vec, - ) -> Result<(), (usize, String)> { + ) -> Result<(), UndoError> { let total = undo_log.len(); for (rev_idx, entry) in undo_log.into_iter().rev().enumerate() { // Convert reversed index back to original forward-order index for @@ -34,16 +33,15 @@ impl CoreLoop { Ok(()) } - /// Apply a single undo entry. Returns `Err((entry_index, detail))` if the - /// undo cannot be applied — this is a fatal condition: the shard's in-memory - /// state is now partially rolled back and must not serve writes. + /// Apply a single undo entry. An error is fatal: the shard's in-memory + /// state is partially rolled back and must not serve writes. fn apply_undo_entry( &mut self, did: u64, tid: u64, entry_index: usize, entry: UndoEntry, - ) -> Result<(), (usize, String)> { + ) -> Result<(), UndoError> { match entry { UndoEntry::ChainHead { collection, prior } => { self.undo_chain_hash(did, tid, &collection, entry_index, Some(prior)) @@ -51,8 +49,12 @@ impl CoreLoop { UndoEntry::PutDocument { .. } | UndoEntry::DeleteDocument { .. } => { self.apply_undo_document(did, tid, entry_index, entry) } - UndoEntry::InsertVector { .. } | UndoEntry::DeleteVector { .. } => { - self.apply_undo_vector(tid, entry_index, entry) + UndoEntry::InsertVector { .. } + | UndoEntry::DeleteVector { .. } + | UndoEntry::DisplacedVector { .. } => self.apply_undo_vector(tid, entry_index, entry), + UndoEntry::VectorCollectionCreated { index_key } => { + self.vector_collections.remove(&index_key); + Ok(()) } UndoEntry::SpatialInsert { .. } | UndoEntry::SpatialDelete { .. } => { self.apply_undo_spatial(entry_index, entry) @@ -84,16 +86,15 @@ impl CoreLoop { UndoEntry::SpatialRow(undo) => self.apply_undo_spatial_row(entry_index, *undo), UndoEntry::VectorWrite(undo) => self.apply_undo_vector_write(entry_index, *undo), UndoEntry::CrdtCollection(undo) => self.apply_undo_crdt_collection(entry_index, *undo), + UndoEntry::CrdtRow(undo) => self.apply_undo_crdt_row(entry_index, *undo), UndoEntry::ArrayTiles { array_id, snapshot } => self .array_engine .restore_tiles(&array_id, snapshot) .map_err(|e| { - ( + UndoError::failed( entry_index, - format!( - "restoring the memtable tiles of array '{}': {e}", - array_id.name - ), + format!("restoring the memtable tiles of array '{}'", array_id.name), + e, ) }), UndoEntry::SparseDoc { @@ -439,7 +440,7 @@ mod tests { .is_some(); let csr = !core .csr_partition_mut(DB, TID) - .neighbors(PK, None, Direction::Out) + .neighbors(PK, &[], Direction::Out) .is_empty(); store && csr } diff --git a/nodedb/src/data/executor/handlers/transaction/undo/spatial.rs b/nodedb/src/data/executor/handlers/transaction/undo/spatial.rs index 7a0b412a4..ebcf268ba 100644 --- a/nodedb/src/data/executor/handlers/transaction/undo/spatial.rs +++ b/nodedb/src/data/executor/handlers/transaction/undo/spatial.rs @@ -7,19 +7,19 @@ //! write transaction does NOT reverse them — they require explicit undo. This //! mirrors the vector-index undo path (`apply_undo_vector`). //! -//! Returns `Err((entry_index, detail))` on fatal failure so the caller can -//! escalate to a typed `RollbackFailed` response. +//! Returns an [`UndoError`] on fatal failure, and the caller escalates it to +//! a typed `RollbackFailed` response. use crate::data::executor::core_loop::CoreLoop; -use super::UndoEntry; +use super::{UndoEntry, UndoError}; impl CoreLoop { pub(super) fn apply_undo_spatial( &mut self, _entry_index: usize, entry: UndoEntry, - ) -> Result<(), (usize, String)> { + ) -> Result<(), UndoError> { match entry { UndoEntry::SpatialInsert { key, entry_id } => { // Reverse a forward spatial insert: drop the R-tree entry and diff --git a/nodedb/src/data/executor/handlers/transaction/undo/spatial_row.rs b/nodedb/src/data/executor/handlers/transaction/undo/spatial_row.rs index a67b5bd34..304e2ccf7 100644 --- a/nodedb/src/data/executor/handlers/transaction/undo/spatial_row.rs +++ b/nodedb/src/data/executor/handlers/transaction/undo/spatial_row.rs @@ -14,7 +14,7 @@ use crate::data::executor::spatial_key::SpatialIndexKey; use crate::types::{DatabaseId, TenantId}; use crate::util::fnv1a_hash; -use super::UndoEntry; +use super::{UndoEntry, UndoError}; /// The pre-image of one spatial row. pub(in crate::data::executor) struct SpatialRowUndo { @@ -79,7 +79,7 @@ impl CoreLoop { &mut self, entry_index: usize, undo: SpatialRowUndo, - ) -> Result<(), (usize, String)> { + ) -> Result<(), UndoError> { let SpatialRowUndo { key, storage_key, @@ -100,9 +100,10 @@ impl CoreLoop { .map(drop), }; restored.map_err(|e| { - ( + UndoError::failed( entry_index, - format!("restoring spatial row in '{}': {e}", key.2), + format!("restoring spatial row in '{}'", key.2), + e, ) })?; diff --git a/nodedb/src/data/executor/handlers/transaction/undo/stats.rs b/nodedb/src/data/executor/handlers/transaction/undo/stats.rs index 5f90d3389..44be354bb 100644 --- a/nodedb/src/data/executor/handlers/transaction/undo/stats.rs +++ b/nodedb/src/data/executor/handlers/transaction/undo/stats.rs @@ -9,21 +9,21 @@ //! must explicitly restore the captured pre-image (mirroring the vector/spatial //! undo paths, which reverse side-effects an aborted redb txn leaves behind). //! -//! Returns `Err((entry_index, detail))` on fatal failure so the caller can -//! escalate to a typed `RollbackFailed` response. +//! Returns an [`UndoError`] on fatal failure, and the caller escalates it to +//! a typed `RollbackFailed` response. use tracing::error; use crate::data::executor::core_loop::CoreLoop; -use super::UndoEntry; +use super::{UndoEntry, UndoError}; impl CoreLoop { pub(super) fn apply_undo_stats( &mut self, entry_index: usize, entry: UndoEntry, - ) -> Result<(), (usize, String)> { + ) -> Result<(), UndoError> { match entry { UndoEntry::StatsRestore { key, prior } => { // Restore the exact pre-image via the stats store's own write @@ -33,14 +33,14 @@ impl CoreLoop { self.stats_store .restore(&key, prior.as_deref()) .map_err(|e| { - let detail = format!("stats restore {key}: {e}"); + let err = UndoError::failed(entry_index, format!("stats restore {key}"), e); error!( core = self.core_id, entry_index, - error = %detail, + error = %err, "transaction undo: column stats restore failed; shard state unknown" ); - (entry_index, detail) + err }) } _ => unreachable!("apply_undo_stats called with non-stats entry"), diff --git a/nodedb/src/data/executor/handlers/transaction/undo/timeseries.rs b/nodedb/src/data/executor/handlers/transaction/undo/timeseries.rs index 23a646441..3f84f8e25 100644 --- a/nodedb/src/data/executor/handlers/transaction/undo/timeseries.rs +++ b/nodedb/src/data/executor/handlers/transaction/undo/timeseries.rs @@ -14,7 +14,7 @@ use crate::data::executor::core_loop::CoreLoop; use crate::types::{DatabaseId, TenantId}; -use super::{TimeseriesIngestUndo, TimeseriesTruncateUndo}; +use super::{TimeseriesIngestUndo, TimeseriesTruncateUndo, UndoError}; impl CoreLoop { /// The complete in-memory pre-image of a timeseries collection before an @@ -45,7 +45,7 @@ impl CoreLoop { &mut self, entry_index: usize, undo: TimeseriesTruncateUndo, - ) -> Result<(), (usize, String)> { + ) -> Result<(), UndoError> { let TimeseriesTruncateUndo { collection_key, original_dir: original, @@ -64,24 +64,23 @@ impl CoreLoop { if original.exists() && let Err(e) = std::fs::remove_dir_all(&original) { - return Err(( + return Err(UndoError::failed( entry_index, - format!( - "timeseries truncate undo: remove {}: {e}", - original.display() - ), + format!("timeseries truncate undo: remove {}", original.display()), + e, )); } if let Some(moved) = moved_dir && let Err(e) = std::fs::rename(&moved, &original) { - return Err(( + return Err(UndoError::failed( entry_index, format!( - "timeseries truncate undo: rename {} back to {}: {e}", + "timeseries truncate undo: rename {} back to {}", moved.display(), original.display() ), + e, )); } diff --git a/nodedb/src/data/executor/handlers/transaction/undo/truncate_columnar.rs b/nodedb/src/data/executor/handlers/transaction/undo/truncate_columnar.rs index 6688cd11d..e36fb7aaa 100644 --- a/nodedb/src/data/executor/handlers/transaction/undo/truncate_columnar.rs +++ b/nodedb/src/data/executor/handlers/transaction/undo/truncate_columnar.rs @@ -10,14 +10,14 @@ use crate::data::executor::core_loop::CoreLoop; -use super::ColumnarTruncateUndo; +use super::{ColumnarTruncateUndo, UndoError}; impl CoreLoop { pub(super) fn apply_undo_columnar_truncate( &mut self, entry_index: usize, undo: ColumnarTruncateUndo, - ) -> Result<(), (usize, String)> { + ) -> Result<(), UndoError> { let ColumnarTruncateUndo { collection_key, rows, @@ -27,7 +27,7 @@ impl CoreLoop { spatial_doc_map, } = undo; let Some(engine) = self.columnar_engines.get_mut(&collection_key) else { - return Err(( + return Err(UndoError::mismatch( entry_index, format!( "columnar truncate undo: engine for {:?} vanished before rollback", diff --git a/nodedb/src/data/executor/handlers/transaction/undo/vector_truncate.rs b/nodedb/src/data/executor/handlers/transaction/undo/vector_truncate.rs index 2ef9cd200..9d8b161a4 100644 --- a/nodedb/src/data/executor/handlers/transaction/undo/vector_truncate.rs +++ b/nodedb/src/data/executor/handlers/transaction/undo/vector_truncate.rs @@ -16,7 +16,7 @@ use crate::data::executor::core_loop::CoreLoop; use crate::data::executor::handlers::vector_direct_row::VectorIndexKey; use crate::engine::vector::collection::VectorCollection; -use super::UndoEntry; +use super::{UndoEntry, UndoError}; /// The collection and sidecar rows a truncate would remove. pub(in crate::data::executor) struct VectorTruncateUndo { @@ -70,7 +70,7 @@ impl CoreLoop { &mut self, entry_index: usize, undo: VectorTruncateUndo, - ) -> Result<(), (usize, String)> { + ) -> Result<(), UndoError> { let VectorTruncateUndo { index_key, tid, @@ -92,9 +92,10 @@ impl CoreLoop { self.sparse .put(database_id, tid, &collection, &key, &bytes) .map_err(|e| { - ( + UndoError::failed( entry_index, - format!("restoring the sidecar of {key} in '{collection}': {e}"), + format!("restoring the sidecar of {key} in '{collection}'"), + e, ) })?; self.doc_cache diff --git a/nodedb/src/data/executor/handlers/transaction/undo/vector_write.rs b/nodedb/src/data/executor/handlers/transaction/undo/vector_write.rs index 34221689f..7be8de193 100644 --- a/nodedb/src/data/executor/handlers/transaction/undo/vector_write.rs +++ b/nodedb/src/data/executor/handlers/transaction/undo/vector_write.rs @@ -19,7 +19,7 @@ use crate::data::executor::core_loop::CoreLoop; use crate::data::executor::handlers::vector_direct_row::VectorIndexKey; use crate::engine::vector::collection::VectorWriteMark; -use super::UndoEntry; +use super::{UndoEntry, UndoError}; /// The pre-image of one vector write. pub(in crate::data::executor) struct VectorWriteUndo { @@ -115,7 +115,7 @@ impl CoreLoop { &mut self, entry_index: usize, undo: VectorWriteUndo, - ) -> Result<(), (usize, String)> { + ) -> Result<(), UndoError> { let VectorWriteUndo { index_key, tid, @@ -126,7 +126,6 @@ impl CoreLoop { payload_rows, } = undo; let database_id = index_key.0.as_u64(); - let fail = |detail: String| (entry_index, detail); // The bitmap entries of the rows the write left behind go first: their // fields sit in the sidecars the write stored. @@ -141,7 +140,9 @@ impl CoreLoop { let current = self .sparse .get(database_id, tid, &collection, &key) - .map_err(|e| fail(format!("reading the sidecar of {key}: {e}")))?; + .map_err(|e| { + UndoError::failed(entry_index, format!("reading the sidecar of {key}"), e) + })?; if let Some(bytes) = current && let Ok(fields) = crate::data::executor::handlers::vector_upsert::decode_payload_lowercased( @@ -156,17 +157,21 @@ impl CoreLoop { match mark { Some(mark) => { let Some(coll) = self.vector_collections.get_mut(&index_key) else { - return Err(fail(format!( - "vector index {:?} vanished before its write was rolled back", - index_key - ))); + return Err(UndoError::mismatch( + entry_index, + format!( + "vector index {index_key:?} vanished before its write was rolled back" + ), + )); }; if !coll.roll_back_to(mark) { - return Err(fail(format!( - "vector index {:?} sealed or trained away the nodes a rolled-back \ - write inserted", - index_key - ))); + return Err(UndoError::mismatch( + entry_index, + format!( + "vector index {index_key:?} sealed or trained away the nodes a \ + rolled-back write inserted" + ), + )); } for (id, fields) in &payload_rows { coll.payload.insert_row(*id, fields); @@ -192,7 +197,9 @@ impl CoreLoop { .delete(database_id, tid, &collection, &key) .map(drop), }; - restored.map_err(|e| fail(format!("restoring the sidecar of {key}: {e}")))?; + restored.map_err(|e| { + UndoError::failed(entry_index, format!("restoring the sidecar of {key}"), e) + })?; self.doc_cache .invalidate(database_id, tid, &collection, &key); } diff --git a/nodedb/src/data/executor/handlers/upsert/exec/insert.rs b/nodedb/src/data/executor/handlers/upsert/exec/insert.rs index d028ea163..e929130e3 100644 --- a/nodedb/src/data/executor/handlers/upsert/exec/insert.rs +++ b/nodedb/src/data/executor/handlers/upsert/exec/insert.rs @@ -3,10 +3,10 @@ //! The upsert insert branch: no existing row was found, so insert fresh //! (identical in shape to a `PointPut`, plus chain + enforcement). -use crate::bridge::envelope::{ErrorCode, Response}; +use crate::bridge::envelope::Response; use crate::data::executor::core_loop::CoreLoop; use crate::data::executor::core_loop::redo_image::submitted_row_image; -use crate::data::executor::enforcement::chain_guard::{self, ChainGuard}; +use crate::data::executor::enforcement::chain_guard::{self, AbandonedWrite, ChainGuard}; use crate::data::executor::enforcement::write_hook::{self, HookCtx, ImageBody, WriteImages}; use crate::data::executor::handlers::point::apply_put::PointPutParams; use crate::data::executor::handlers::rls_write_gate; @@ -58,6 +58,7 @@ impl CoreLoop { // The plan's `document_id` is the row's client identity: the write // gate, the event, the redo entry, and `RETURNING` all name it. let document_identity = RowIdentity::from_user_key(document_id); + let identity_column = self.identity_column(database_id, tid, collection); // Insert: document doesn't exist, create new (same as PointPut). // The incoming body IS the post-image here, and the planner @@ -69,6 +70,7 @@ impl CoreLoop { value, &document_identity, None, + &identity_column, tid, collection, ) { @@ -96,7 +98,7 @@ impl CoreLoop { // existence probe just above found none, and apply_point_put // is the only writer on this core — prior must be None. We // pass it straight through so the emit resolves to Insert. - let prior = match self.apply_point_put( + let mut prior = match self.apply_point_put( &txn, PointPutParams { database_id, @@ -115,13 +117,11 @@ impl CoreLoop { ) { Ok(p) => p, Err(e) => { - chain_guard::abort_after_apply( + let e = chain_guard::abort_after_apply( self, &mut chain, - database_id, - tid, - collection, - &storage_key, + AbandonedWrite::row(database_id, tid, collection, &storage_key), + e, ); return self.response_error(task, e); } @@ -133,13 +133,12 @@ impl CoreLoop { .settle(self, surrogate, &prior.stored_value) .and_then(|()| chain.persist_head(self, &txn)) { - chain_guard::abort_after_apply( + let e = chain_guard::abort_after_apply( self, &mut chain, - database_id, - tid, - collection, - &storage_key, + AbandonedWrite::row(database_id, tid, collection, &storage_key) + .undo(std::mem::take(&mut prior.memory_undo)), + e, ); return self.response_error(task, e); } @@ -157,50 +156,48 @@ impl CoreLoop { ) { Ok(o) => o, Err(e) => { - chain_guard::abort_after_apply( + let e = chain_guard::abort_after_apply( self, &mut chain, - database_id, - tid, - collection, - &storage_key, + AbandonedWrite::row(database_id, tid, collection, &storage_key) + .undo(std::mem::take(&mut prior.memory_undo)), + e, ); return self.response_error(task, e); } }; let target_write_set = write_hook::target_write_set(&enforcement.target_writes); + let target_writes = enforcement.target_writes; // Settled before the commit, so an insert of one journal leg on // its own leaves nothing behind. if let Err(e) = self.settle_balanced_entries(database_id, tid, collection, enforcement.balanced_entries) { - chain_guard::abort_after_apply( + let e = chain_guard::abort_after_apply( self, &mut chain, - database_id, - tid, - collection, - &storage_key, + AbandonedWrite::row(database_id, tid, collection, &storage_key) + .undo(std::mem::take(&mut prior.memory_undo)) + .targets(target_writes), + e, ); return self.response_error(task, e); } if let Err(e) = txn.commit() { - chain_guard::abort_after_apply( + let e = chain_guard::abort_after_apply( self, &mut chain, - database_id, - tid, - collection, - &storage_key, - ); - return self.response_error( - task, - ErrorCode::Internal { + AbandonedWrite::row(database_id, tid, collection, &storage_key) + .undo(std::mem::take(&mut prior.memory_undo)) + .targets(target_writes), + crate::Error::Storage { + engine: "sparse".into(), detail: format!("commit: {e}"), }, ); + return self.response_error(task, e); } self.emit_put_event( @@ -219,6 +216,7 @@ impl CoreLoop { spec, rls_filters, strict_schema, + &identity_column, &[(&document_identity, prior.stored_value.as_slice())], ), None => self.response_affected(task, 1), diff --git a/nodedb/src/data/executor/handlers/upsert/exec/overwrite.rs b/nodedb/src/data/executor/handlers/upsert/exec/overwrite.rs index 2dd3a967f..671958506 100644 --- a/nodedb/src/data/executor/handlers/upsert/exec/overwrite.rs +++ b/nodedb/src/data/executor/handlers/upsert/exec/overwrite.rs @@ -7,9 +7,11 @@ use crate::bridge::envelope::{ErrorCode, Response}; use crate::data::executor::core_loop::CoreLoop; use crate::data::executor::core_loop::redo_image::submitted_row_image; +use crate::data::executor::enforcement::chain_guard::{AbandonedWrite, abandon_write}; use crate::data::executor::enforcement::write_hook::{self, HookCtx, ImageBody, WriteImages}; -use crate::data::executor::handlers::point::apply_put::PointPutParams; +use crate::data::executor::handlers::point::apply_put::{PointPutParams, VectorIndexDelta}; use crate::data::executor::handlers::rls_write_gate; +use crate::data::executor::handlers::transaction::undo::UndoEntry; use crate::data::executor::handlers::upsert::merge::{apply_on_conflict_updates, merge_values}; use crate::data::executor::task::ExecutionTask; use crate::engine::document::store::{RowIdentity, StorageKey}; @@ -155,26 +157,19 @@ impl CoreLoop { // post-image the policy never saw, which is why this branch // cannot be admitted at plan time. Decided on the MessagePack // form, exactly as the insert arm below decides its own body. + let identity_column = self.identity_column(database_id, tid, collection); if let Err(e) = rls_write_gate::admit_stored_row( rls_write_check, &merged_body, &document_identity, None, + &identity_column, tid, collection, ) { return self.response_error(task, e); } - // The surrogate is stable across an overwrite and - // `insert_with_surrogate` APPENDS an HNSW node rather than - // replacing one, so the prior embedding has to come out before - // the write below puts the new one in — otherwise KNN keeps - // scoring both. No-op when `has_vectors` is false. - if has_vectors { - self.remove_document_vector_indexes(database_id, tid, collection, storage_key); - } - // One transaction for the body, every index that describes it, // and every derived write the collection's constraints imply. // The bare `sparse.put` this replaces reconciled none of them: @@ -185,15 +180,25 @@ impl CoreLoop { let txn = match self.sparse.begin_write() { Ok(t) => t, Err(e) => { - return self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ); + return self.response_error(task, ErrorCode::from(e)); } }; - let outcome = match self.apply_point_put( + + // The surrogate is stable across an overwrite and + // `insert_with_surrogate` APPENDS an HNSW node rather than + // replacing one, so the prior embedding has to come out before + // the write below puts the new one in — otherwise KNN keeps + // scoring both. No-op when `has_vectors` is false. The removal is + // in memory, so an abort below puts the prior nodes back. + let mut memory_undo: Vec = Vec::new(); + if has_vectors { + memory_undo.extend( + self.remove_document_vector_indexes(database_id, tid, collection, storage_key) + .into_iter() + .map(VectorIndexDelta::into_delete_undo), + ); + } + let mut outcome = match self.apply_point_put( &txn, PointPutParams { database_id, @@ -216,11 +221,16 @@ impl CoreLoop { // dropping `txn` reverses the durable write but not // that entry, which would then serve a body that never // committed. - self.doc_cache - .invalidate(database_id, tid, collection, &storage_key); + let e = abandon_write( + self, + AbandonedWrite::row(database_id, tid, collection, &storage_key) + .undo(memory_undo), + e, + ); return self.response_error(task, e); } }; + memory_undo.append(&mut outcome.memory_undo); // An overwrite is an UPDATE, and the pre-image is what tells the // fold to take the row's old contribution off a total before @@ -237,12 +247,17 @@ impl CoreLoop { ) { Ok(o) => o, Err(e) => { - self.doc_cache - .invalidate(database_id, tid, collection, &storage_key); + let e = abandon_write( + self, + AbandonedWrite::row(database_id, tid, collection, &storage_key) + .undo(memory_undo), + e, + ); return self.response_error(task, e); } }; let target_write_set = write_hook::target_write_set(&enforcement.target_writes); + let target_writes = enforcement.target_writes; // Settled before the commit: the merged row's old amount comes // off the group and its new one goes on, so an overwrite that @@ -250,18 +265,28 @@ impl CoreLoop { if let Err(e) = self.settle_balanced_entries(database_id, tid, collection, enforcement.balanced_entries) { - self.doc_cache - .invalidate(database_id, tid, collection, &storage_key); + let e = abandon_write( + self, + AbandonedWrite::row(database_id, tid, collection, &storage_key) + .undo(memory_undo) + .targets(target_writes), + e, + ); return self.response_error(task, e); } if let Err(e) = txn.commit() { - return self.response_error( - task, - ErrorCode::Internal { + let e = abandon_write( + self, + AbandonedWrite::row(database_id, tid, collection, &storage_key) + .undo(memory_undo) + .targets(target_writes), + crate::Error::Storage { + engine: "sparse".into(), detail: format!("commit: {e}"), }, ); + return self.response_error(task, e); } // `current_bytes` is the pre-merge stored row, already read @@ -287,6 +312,7 @@ impl CoreLoop { spec, rls_filters, strict_schema, + &identity_column, &[(&document_identity, stored_bytes.as_slice())], ), None => self.response_affected(task, 1), diff --git a/nodedb/src/data/executor/handlers/write_batch.rs b/nodedb/src/data/executor/handlers/write_batch.rs index b195bb281..410648806 100644 --- a/nodedb/src/data/executor/handlers/write_batch.rs +++ b/nodedb/src/data/executor/handlers/write_batch.rs @@ -222,12 +222,14 @@ impl CoreLoop { // Don't commit — transaction is dropped (implicit rollback). // Send error responses for failed tasks, put successful ones back. drop(txn); + let fatal = self.abandon_batched_puts(&batch, &mut results); let count = batch.len(); let mut responses = Vec::with_capacity(count); for (task, result) in batch.iter().zip(results) { - let response = match result { - Err(err_response) => err_response, - Ok(_) => self.response_error( + let response = match (&fatal, result) { + (Some(code), _) => self.response_error(task, code.clone()), + (None, Err(err_response)) => err_response, + (None, Ok(_)) => self.response_error( task, ErrorCode::Internal { detail: "batch aborted due to sibling failure".into(), @@ -244,12 +246,17 @@ impl CoreLoop { // Commit once for all writes. let commit_result = txn.commit(); + let fatal = match commit_result { + Ok(()) => None, + Err(_) => self.abandon_batched_puts(&batch, &mut results), + }; let count = batch.len(); let mut responses = Vec::with_capacity(count); for (task, result) in batch.iter().zip(results.iter()) { - let response = match &commit_result { - Ok(()) => { + let response = match (&commit_result, &fatal) { + (Err(_), Some(code)) => self.response_error(task, code.clone()), + (Ok(()), _) => { // Emit write event for each successful batched PointPut. // The Insert vs Update tag is derived from the prior // bytes captured per row above. @@ -275,7 +282,7 @@ impl CoreLoop { } self.response_ok(task) } - Err(e) => self.response_error( + (Err(e), None) => self.response_error( task, ErrorCode::Internal { detail: format!("batch commit: {e}"), @@ -306,6 +313,45 @@ impl CoreLoop { count } + /// Reverse the in-memory side effects of batched puts whose shared + /// transaction drops uncommitted: the cache entry each put wrote and the + /// R-tree, vector and sparse entries of each applied put, last put first. + /// + /// Returns `Some(RollbackFailed)` when an entry did not reverse. The + /// core's state is then unknown, so this fail-stops it, and every task + /// of the batch reports that code. + fn abandon_batched_puts( + &mut self, + batch: &[ExecutionTask], + results: &mut [Result], + ) -> Option { + let mut undone: crate::Result<()> = Ok(()); + for (task, result) in batch.iter().zip(results.iter_mut()).rev() { + let PhysicalPlan::Document(DocumentOp::PointPut { + collection, + surrogate, + .. + }) = task.plan() + else { + continue; + }; + let database_id = task.request.database_id.as_u64(); + let tid = task.request.tenant_id.as_u64(); + let key = crate::engine::document::store::StorageKey::for_surrogate(*surrogate); + // A failed put can have cached its row before it failed. + self.doc_cache + .invalidate(database_id, tid, collection.as_str(), &key); + if let Ok(outcome) = result { + let memory_undo = std::mem::take(&mut outcome.memory_undo); + let row = self.undo_memory_effects(database_id, tid, memory_undo); + undone = undone.and(row); + } + } + let fatal = ErrorCode::from(undone.err()?); + self.fail_stop_on_rollback_code(&fatal); + Some(fatal) + } + /// Store the write sets of the journalled tasks of `batch`, finish the /// journalled run, then send every response. fn finish_journalled_batch( diff --git a/nodedb/src/data/executor/wal_replay_columnar_image.rs b/nodedb/src/data/executor/wal_replay_columnar_image.rs index 8e48f478f..075537626 100644 --- a/nodedb/src/data/executor/wal_replay_columnar_image.rs +++ b/nodedb/src/data/executor/wal_replay_columnar_image.rs @@ -79,7 +79,7 @@ fn image_values(schema: &ColumnarSchema, image: &Value) -> Result, Er schema .columns .iter() - .map(|col| ndb_field_to_value(fields.get(&col.name), &col.column_type)) + .map(|col| ndb_field_to_value(fields.get(&col.name), col)) .collect::>>() .map_err(ErrorCode::from) } @@ -231,6 +231,7 @@ impl CoreLoop { first, schema_bytes, ) + .map_err(ErrorCode::from)? } }; @@ -255,8 +256,9 @@ impl CoreLoop { } let mut undo = Vec::new(); - let removed = - self.apply_columnar_delete_pks(&key, &schema, &priors, recording.then_some(&mut undo)); + let removed = self + .apply_columnar_delete_pks(&key, &schema, &priors, recording.then_some(&mut undo)) + .map_err(ErrorCode::from)?; self.record_redo_undo(undo); if removed.affected != priors.len() as u64 { return Err(ErrorCode::Internal { diff --git a/nodedb/src/data/executor/wal_replay_redo_document_apply.rs b/nodedb/src/data/executor/wal_replay_redo_document_apply.rs index b071b2c3a..93b2a67fe 100644 --- a/nodedb/src/data/executor/wal_replay_redo_document_apply.rs +++ b/nodedb/src/data/executor/wal_replay_redo_document_apply.rs @@ -20,7 +20,9 @@ use nodedb_types::Surrogate; use super::core_loop::CoreLoop; -use super::enforcement::chain_guard::{ChainGuard, abort_after_apply}; +use super::enforcement::chain_guard::{ + AbandonedWrite, ChainGuard, abandon_write, abort_after_apply, +}; use super::handlers::point::apply_delete::PointDeleteParams; use super::handlers::point::apply_put::PointPutParams; use crate::engine::document::store::StorageKey; @@ -78,6 +80,9 @@ impl CoreLoop { return false; } }; + // The row's in-memory index entries, kept so a failed settle or + // commit can reverse them. + let mut memory_undo = Vec::new(); let applied = self .apply_point_put( &txn, @@ -98,7 +103,10 @@ impl CoreLoop { wal_lsn: (row.record_lsn != 0).then(|| crate::types::Lsn::new(row.record_lsn)), }, ) - .and_then(|outcome| chain.settle(self, surrogate, &outcome.stored_value)) + .and_then(|mut outcome| { + memory_undo = std::mem::take(&mut outcome.memory_undo); + chain.settle(self, surrogate, &outcome.stored_value) + }) .and_then(|()| chain.persist_head(self, &txn)); // An error drops the write txn un-committed, which rolls it back. let committed = applied.and_then(|()| { @@ -113,14 +121,14 @@ impl CoreLoop { true } Err(e) => { - abort_after_apply( + let e = abort_after_apply( self, &mut chain, - row.database_id, - row.tenant_id, - collection, - &storage_key, + AbandonedWrite::row(row.database_id, row.tenant_id, collection, &storage_key) + .undo(memory_undo), + e, ); + self.fail_stop_on_failed_undo(&e); tracing::warn!( core = self.core_id, %collection, @@ -157,6 +165,15 @@ impl CoreLoop { ) } + /// Fail-stop the core when `error` reports a failed undo. An in-memory + /// entry that did not reverse leaves the core's state unknown, so the + /// core stops serving. + fn fail_stop_on_failed_undo(&mut self, error: &crate::Error) { + if let crate::Error::DataPlane(code) = error { + self.fail_stop_on_rollback_code(code); + } + } + /// Apply one document DELETE through `apply_point_delete` in its own redb /// write transaction. `enforce = false` for the same reason as the put /// path. The row's node keeps its edges: the WAL names every tombstone @@ -201,6 +218,19 @@ impl CoreLoop { outcome.prior_value.is_some() } Err(e) => { + // The dropped txn reverses the durable writes only. The + // in-memory cascades are reversed here. + let storage_key = StorageKey::for_surrogate(surrogate); + let e = abandon_write( + self, + AbandonedWrite::row(database_id, tenant_id, collection, &storage_key) + .undo(outcome.memory_undo), + crate::Error::Storage { + engine: "sparse".into(), + detail: format!("WAL document redo commit: {e}"), + }, + ); + self.fail_stop_on_failed_undo(&e); tracing::warn!( core = self.core_id, %collection, From 77804e869d5e1e469e806a5c908b67157e5d5de8 Mon Sep 17 00:00:00 2001 From: Farhan Syah Date: Thu, 8 Oct 2026 16:51:47 +0800 Subject: [PATCH 11/24] fix(crdt): derive write sets from imported ops and keep dead letters An import's write set comes from the operations it adds to the oplog, so a blob with pending changes still reports the rows its ready changes wrote. A delta that writes a non-map root container is refused. CRDT document writes capture the row image and restore it on abort. A dead-letter entry the queue or store refuses fails or holds the write instead of being logged, restart replay stops at a record its writer cannot produce, and both cases file a black-box report. Snapshot imports run without the peer byte and operation ceilings. --- nodedb-crdt/src/error.rs | 41 +- nodedb-crdt/src/lib.rs | 5 +- nodedb-crdt/src/state/changed_rows.rs | 373 ++++++++ nodedb-crdt/src/state/mod.rs | 6 + nodedb-crdt/src/state/preview.rs | 17 +- nodedb-crdt/src/state/remove_fields.rs | 134 +++ nodedb-crdt/src/state/row_image.rs | 189 ++++ nodedb-crdt/src/state/snapshot.rs | 18 +- nodedb-crdt/src/state/write_set.rs | 840 +++++++++++++++--- .../executor/core_loop/crdt_dead_letters.rs | 51 +- .../handlers/control/crdt_apply/gated.rs | 180 +++- .../handlers/control/crdt_apply/local.rs | 142 ++- .../handlers/control/crdt_apply/write_set.rs | 47 +- .../executor/handlers/control/crdt_doc.rs | 465 ++++++---- nodedb/src/data/executor/wal_replay/crdt.rs | 734 ++++++++++++--- nodedb/src/diag/context/crdt.rs | 62 +- .../diag/context/crdt_dead_letter_store.rs | 133 +++ nodedb/src/diag/context/mod.rs | 10 +- nodedb/src/diag/mod.rs | 8 +- nodedb/src/diag/recording/crdt.rs | 78 +- nodedb/src/diag/recording/mod.rs | 9 +- .../engine/crdt/tenant_state/apply_target.rs | 73 ++ .../crdt/tenant_state/apply_validated.rs | 346 +++++--- .../engine/crdt/tenant_state/dead_letters.rs | 130 ++- .../engine/crdt/tenant_state/doc_mutate.rs | 31 + nodedb/src/engine/crdt/tenant_state/mod.rs | 4 +- .../engine/crdt/tenant_state/snapshot_io.rs | 40 + .../engine/sparse/btree/crdt_dead_letter.rs | 50 ++ 28 files changed, 3625 insertions(+), 591 deletions(-) create mode 100644 nodedb-crdt/src/state/changed_rows.rs create mode 100644 nodedb-crdt/src/state/remove_fields.rs create mode 100644 nodedb-crdt/src/state/row_image.rs create mode 100644 nodedb/src/diag/context/crdt_dead_letter_store.rs create mode 100644 nodedb/src/engine/crdt/tenant_state/apply_target.rs 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/src/data/executor/core_loop/crdt_dead_letters.rs b/nodedb/src/data/executor/core_loop/crdt_dead_letters.rs index d373bf39a..585f72fdb 100644 --- a/nodedb/src/data/executor/core_loop/crdt_dead_letters.rs +++ b/nodedb/src/data/executor/core_loop/crdt_dead_letters.rs @@ -19,8 +19,9 @@ impl CoreLoop { /// record already bound keeps its one entry. A write with no record has /// nothing to replay, so its entry stays in memory only. /// - /// A store error is returned: the entry is then in memory only, and the - /// caller must not report the rejection as durable. + /// A store error is returned, and the entry is removed from the queue: + /// nothing on this node records the rejection. The caller fails or holds + /// the write that needed the entry. pub(in crate::data::executor) fn store_crdt_dead_letter( &mut self, database_id: DatabaseId, @@ -36,36 +37,50 @@ impl CoreLoop { let Some(entry) = engine.bind_dead_letter_source(source_lsn.as_u64()) else { return Ok(()); }; - self.sparse.put_crdt_dead_letter( + let stored = self.sparse.put_crdt_dead_letter( database_id.as_u64(), tenant_id.as_u64(), source_lsn.as_u64(), &entry, - ) + ); + if let Err(error) = &stored { + crate::diag::crdt_dead_letter_not_stored( + error, + database_id.as_u64(), + tenant_id.as_u64(), + &entry, + source_lsn.as_u64(), + ); + engine.discard_dead_letter(entry.id); + } + stored } } impl CoreLoop { /// Store the entry the rejection of the replayed record at `lsn` - /// produced. The live apply of the record stored the same entry, so this - /// keeps one. A store error is logged, and the record keeps the entry in - /// memory. + /// produced. Returns whether replay may count the record as replayed. + /// + /// A store error stops restart replay at the record: the core + /// fail-stops, no checkpoint covers the record, and boot refuses to start + /// with a rejection nothing records. A committed-redo apply fails instead. pub(in crate::data::executor) fn store_replayed_dead_letter( &mut self, database_id: DatabaseId, tenant_id: TenantId, lsn: u64, - ) { - if let Err(error) = self.store_crdt_dead_letter(database_id, tenant_id, Some(Lsn::new(lsn))) - { - tracing::error!( - core = self.core_id, - %database_id, - %tenant_id, - lsn, - %error, - "a replayed CRDT rejection's dead-letter entry could not be stored" - ); + ) -> bool { + match self.store_crdt_dead_letter(database_id, tenant_id, Some(Lsn::new(lsn))) { + Ok(()) => true, + Err(error) => { + self.replay_record_unapplied( + "crdt", + "dead_letter_store", + lsn, + &format!("the rejected delta's dead-letter entry could not be stored: {error}"), + ); + false + } } } } diff --git a/nodedb/src/data/executor/handlers/control/crdt_apply/gated.rs b/nodedb/src/data/executor/handlers/control/crdt_apply/gated.rs index 7cde32526..e7113d0ea 100644 --- a/nodedb/src/data/executor/handlers/control/crdt_apply/gated.rs +++ b/nodedb/src/data/executor/handlers/control/crdt_apply/gated.rs @@ -22,7 +22,7 @@ use nodedb_types::Surrogate; use nodedb_types::sync::violation::ViolationType; use nodedb_types::sync::wire::{AckStatus, SyncProvenance}; -use crate::bridge::envelope::{ErrorCode, Response}; +use crate::bridge::envelope::{ErrorCode, Response, SyncHold}; use crate::data::executor::core_loop::CoreLoop; use crate::data::executor::sync_gate::SyncAdmit; use crate::data::executor::task::ExecutionTask; @@ -44,6 +44,10 @@ enum GateDisposition { /// to succeed later. The high-water-mark is held back so the re-push is /// admitted rather than deduplicated. Retryable, + /// Nothing was applied and nothing on this node records the delta. The + /// refusal is an error, so the frame's record is cancelled, and the mark + /// is held so the re-push at the same seq is admitted. + NotApplied, /// The delta will never apply. The sender must compensate, and the /// high-water-mark advances so it cannot wedge the stream. Terminal(ViolationType), @@ -177,12 +181,7 @@ impl CoreLoop { Ok(e) => e, Err(e) => { warn!(core = self.core_id, error = %e, "failed to create CRDT engine"); - return self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ); + return self.response_error(task, ErrorCode::from(e)); } }; let installed = engine.installed_constraint_version(collection); @@ -298,21 +297,48 @@ impl CoreLoop { } // The candidate was discarded, so authoritative state did not // move. The record is cancelled and replay never reaches this - // rejection, so the dead-letter entry is stored now. + // rejection, so the dead-letter entry is stored now. An entry the + // store refused is removed from the queue, so nothing on this node + // records the delta: retiring it as terminal would lose it. + // Holding the mark keeps it on the sender, and the store recorded + // the refusal in the black box. GateOutcome::Applied(ValidatedApplyOutcome::Rejected(vt)) => { - if let Err(error) = - self.store_crdt_dead_letter(task.request.database_id, tenant_id, task.wal_lsn()) - { - tracing::error!( - core = self.core_id, - %collection, - %document_id, - %error, - "crdt sync apply rejected a delta, and its dead-letter entry could not \ - be stored; it stays in memory only" - ); + match self.store_crdt_dead_letter( + task.request.database_id, + tenant_id, + task.wal_lsn(), + ) { + Ok(()) => GateDisposition::Terminal(vt), + Err(error) => { + warn!( + core = self.core_id, + %collection, + %document_id, + violation = %vt, + %error, + "crdt sync apply refused a delta whose dead-letter entry could not \ + be stored; nothing applied, high-water-mark held" + ); + GateDisposition::NotApplied + } } - GateDisposition::Terminal(vt) + } + // Rejected, but the dead-letter queue refused the delta, so nothing + // on this node records it. Retiring it as terminal would lose it. + // Holding the mark keeps it on the sender, and the re-push is + // dead-lettered once the queue has room. The apply recorded the + // refusal in the black box. + GateOutcome::Applied(ValidatedApplyOutcome::DeadLetterRefused { violation, error }) => { + warn!( + core = self.core_id, + %collection, + %document_id, + %violation, + %error, + "crdt sync apply refused a delta it could not dead-letter; nothing \ + applied, high-water-mark held" + ); + GateDisposition::NotApplied } GateOutcome::Applied(ValidatedApplyOutcome::Malformed) => { // The bytes did not decode, so nothing was imported. Acking @@ -333,6 +359,20 @@ impl CoreLoop { ), }) } + // This node could not build the candidate. The delta is not at + // fault and nothing records it, so the refusal is an error that + // holds the mark and the sender re-pushes. + GateOutcome::Applied(ValidatedApplyOutcome::CandidateUnavailable { error }) => { + warn!( + core = self.core_id, + %collection, + %document_id, + %error, + "crdt sync apply refused: no apply candidate; nothing applied, \ + high-water-mark held" + ); + GateDisposition::NotApplied + } GateOutcome::Applied(ValidatedApplyOutcome::PendingDependencies) => { // Well-formed operations that arrived without their causal // history: Loro buffered them and the applied state did not @@ -380,6 +420,13 @@ impl CoreLoop { // exists to prevent. Report the unchanged mark instead. self.sync_ack_response(task, AckStatus::Gap { expected: prov.seq }, current_hwm) } + GateDisposition::NotApplied => self.response_error( + task, + ErrorCode::SyncNotApplied { + hold: SyncHold::Gap { expected: prov.seq }, + applied_seq: current_hwm, + }, + ), GateDisposition::Terminal(violation) => { // Permanently refused: it will never succeed on a re-push, so // holding the stream for it buys nothing. @@ -466,6 +513,99 @@ mod tests { assert_eq!(entries[0].source_lsn, Some(11)); } + /// A peer delta a constraint refuses while the dead-letter queue is full + /// is an error that holds the mark: nothing records it, so the sender + /// keeps it and its re-push at the same seq is admitted. + #[test] + fn a_refusal_the_full_dead_letter_queue_refuses_holds_the_mark() { + let dir = tempfile::tempdir().expect("tempdir"); + let (mut core, _request_tx, _response_rx) = make_core_with_dir(dir.path()); + let seed = task_at(10); + install_unique_email(&mut core, &seed); + let first = user_delta(2, "a", "x@y.com"); + assert_eq!( + core.execute_crdt_apply(&seed, users_params("a", &first)) + .status, + Status::Ok + ); + let capacity = core + .get_crdt_engine(seed.request.database_id, seed.request.tenant_id) + .expect("engine") + .fill_dead_letter_queue_for_test(); + + let second = user_delta(3, "b", "x@y.com"); + let prov = provenance(1); + let mut params = users_params("b", &second); + params.peer_id = 3; + params.provenance = Some(&prov); + let response = core.execute_crdt_apply(&task_at(11), params); + + assert_eq!(response.status, Status::Error); + assert!( + matches!( + response.error_code.as_deref(), + Some(ErrorCode::SyncNotApplied { + hold: crate::bridge::envelope::SyncHold::Gap { expected: 1 }, + applied_seq: 0, + }) + ), + "got {:?}", + response.error_code + ); + assert_eq!(core.sync_hwm_value(9, 1), 0, "the mark is held"); + assert!( + core.crdt_engines + .get(&(seed.request.database_id, seed.request.tenant_id)) + .is_some_and(|engine| !engine.row_exists("users", "b")) + ); + assert_eq!(dead_letters(&mut core).len(), capacity, "nothing queued"); + } + + /// A peer delta a constraint refuses while the store refuses its + /// dead-letter entry is an error that holds the mark, and the queue keeps + /// no entry storage lacks. + #[test] + fn a_refusal_whose_dead_letter_the_store_refuses_holds_the_mark() { + let dir = tempfile::tempdir().expect("tempdir"); + let (mut core, _request_tx, _response_rx) = make_core_with_dir(dir.path()); + let seed = task_at(10); + install_unique_email(&mut core, &seed); + let first = user_delta(2, "a", "x@y.com"); + assert_eq!( + core.execute_crdt_apply(&seed, users_params("a", &first)) + .status, + Status::Ok + ); + core.sparse.break_crdt_dead_letter_table_for_test(); + + let second = user_delta(3, "b", "x@y.com"); + let prov = provenance(1); + let mut params = users_params("b", &second); + params.peer_id = 3; + params.provenance = Some(&prov); + let response = core.execute_crdt_apply(&task_at(11), params); + + assert_eq!(response.status, Status::Error); + assert!( + matches!( + response.error_code.as_deref(), + Some(ErrorCode::SyncNotApplied { + hold: crate::bridge::envelope::SyncHold::Gap { expected: 1 }, + applied_seq: 0, + }) + ), + "got {:?}", + response.error_code + ); + assert_eq!(core.sync_hwm_value(9, 1), 0, "the mark is held"); + assert!( + core.crdt_engines + .get(&(seed.request.database_id, seed.request.tenant_id)) + .is_some_and(|engine| !engine.row_exists("users", "b")) + ); + assert!(dead_letters(&mut core).is_empty(), "nothing queued"); + } + /// A delta with no target document is refused before anything installs. #[test] fn a_delta_without_a_target_document_is_refused_before_it_installs() { diff --git a/nodedb/src/data/executor/handlers/control/crdt_apply/local.rs b/nodedb/src/data/executor/handlers/control/crdt_apply/local.rs index 75aa6362a..68eaab503 100644 --- a/nodedb/src/data/executor/handlers/control/crdt_apply/local.rs +++ b/nodedb/src/data/executor/handlers/control/crdt_apply/local.rs @@ -29,6 +29,13 @@ enum LocalRefusal { /// Permanent: a row the delta writes violates a constraint. Nothing /// applied, and the delta is in the dead-letter queue. Constraint(ViolationType), + /// A row the delta writes violates a constraint, and the dead-letter + /// queue refused the delta. Nothing applied and nothing records it, so + /// the sender keeps it and re-sends once the queue has room. + DeadLetterRefused { detail: String }, + /// This node could not build the candidate the delta validates in. + /// Nothing applied, and the delta is not at fault, so the sender re-sends. + CandidateUnavailable { detail: String }, } impl CoreLoop { @@ -89,12 +96,7 @@ impl CoreLoop { Ok(e) => e, Err(e) => { warn!(core = self.core_id, error = %e, "failed to create CRDT engine"); - return self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ); + return self.response_error(task, ErrorCode::from(e)); } }; let outcome = engine.apply_committed_delta_validated( @@ -112,7 +114,23 @@ impl CoreLoop { Ok(Self::encode_crdt_row(engine, collection, document_id)) } ValidatedApplyOutcome::Rejected(vt) => Err(LocalRefusal::Constraint(vt)), + ValidatedApplyOutcome::DeadLetterRefused { violation, error } => { + Err(LocalRefusal::DeadLetterRefused { + detail: format!( + "delta for {collection}/{document_id} violates {violation}, and \ + the dead-letter queue refused it: {error}; nothing was applied" + ), + }) + } ValidatedApplyOutcome::Malformed => Err(LocalRefusal::Malformed), + ValidatedApplyOutcome::CandidateUnavailable { error } => { + Err(LocalRefusal::CandidateUnavailable { + detail: format!( + "no apply candidate for {collection}: {error}; delta for \ + {collection}/{document_id} was not applied" + ), + }) + } ValidatedApplyOutcome::PendingDependencies => { // Nothing was imported: the operations are buffered awaiting // predecessors this collection's document has never seen. @@ -180,17 +198,22 @@ impl CoreLoop { ); // The record is cancelled, so replay never reaches // this rejection. Its dead-letter entry is stored - // before the refusal is reported. + // before the refusal is reported. An entry the store + // refused is removed from the queue, so nothing + // records the delta and the sender keeps it, as for + // a queue refusal. The store recorded the refusal in + // the black box. match self.store_crdt_dead_letter( task.request.database_id, tenant_id, task.wal_lsn(), ) { Ok(()) => crdt_rejection(collection, document_id, &violation), - Err(error) => ErrorCode::Internal { - detail: format!( + Err(error) => ErrorCode::RetryableRefusal { + reason: format!( "delta for {collection}/{document_id} violates {violation}, \ - and its dead-letter entry could not be stored: {error}" + and its dead-letter entry could not be stored: {error}; \ + nothing was applied" ), }, } @@ -206,6 +229,20 @@ impl CoreLoop { ); ErrorCode::RetryableRefusal { reason: detail } } + // The apply recorded the refusal in the black box. + LocalRefusal::DeadLetterRefused { detail } => { + ErrorCode::RetryableRefusal { reason: detail } + } + LocalRefusal::CandidateUnavailable { detail } => { + warn!( + core = self.core_id, + %collection, + %document_id, + detail = %detail, + "crdt apply refused retryably: no apply candidate" + ); + ErrorCode::RetryableRefusal { reason: detail } + } }; return self.response_error(task, code); } @@ -474,6 +511,91 @@ pub(in crate::data::executor::handlers::control::crdt_apply) mod tests { ); } + /// A rejected delta whose dead-letter entry the store refuses is a + /// retryable error, and the queue keeps no entry storage lacks. + #[test] + fn a_rejection_whose_dead_letter_the_store_refuses_is_a_retryable_error() { + let dir = tempfile::tempdir().expect("tempdir"); + let (mut core, _request_tx, _response_rx) = make_core_with_dir(dir.path()); + let seed = task_at(20); + install_unique_email(&mut core, &seed); + let first = user_delta(2, "a", "x@y.com"); + assert_eq!( + core.apply_crdt_local(&seed, users_params("a", &first)) + .status, + Status::Ok + ); + core.sparse.break_crdt_dead_letter_table_for_test(); + + let second = user_delta(3, "b", "x@y.com"); + let response = core.apply_crdt_local(&task_at(21), users_params("b", &second)); + + assert_eq!(response.status, Status::Error); + assert!( + matches!( + response.error_code.as_deref(), + Some(ErrorCode::RetryableRefusal { reason }) if reason.contains("could not be stored") + ), + "got {:?}", + response.error_code + ); + assert!( + core.crdt_engines + .get(&(seed.request.database_id, seed.request.tenant_id)) + .is_some_and(|engine| !engine.row_exists("users", "b")), + "a refused delta must not reach the CRDT state" + ); + assert!(dead_letters(&mut core).is_empty(), "nothing queued"); + } + + /// A rejected delta the full dead-letter queue refuses is refused as a + /// retryable error, not as a constraint verdict: nothing records it, so + /// the sender keeps it. No entry is stored for its record. + #[test] + fn a_rejection_the_full_dead_letter_queue_refuses_is_a_retryable_error() { + let dir = tempfile::tempdir().expect("tempdir"); + let (mut core, _request_tx, _response_rx) = make_core_with_dir(dir.path()); + let seed = task_at(20); + install_unique_email(&mut core, &seed); + let first = user_delta(2, "a", "x@y.com"); + assert_eq!( + core.apply_crdt_local(&seed, users_params("a", &first)) + .status, + Status::Ok + ); + core.get_crdt_engine(seed.request.database_id, seed.request.tenant_id) + .expect("engine") + .fill_dead_letter_queue_for_test(); + + let second = user_delta(3, "b", "x@y.com"); + let response = core.apply_crdt_local(&task_at(21), users_params("b", &second)); + + assert_eq!(response.status, Status::Error); + assert!( + matches!( + response.error_code.as_deref(), + Some(ErrorCode::RetryableRefusal { reason }) if reason.contains("dead-letter") + ), + "got {:?}", + response.error_code + ); + let db = seed.request.database_id; + let tenant = seed.request.tenant_id; + assert!( + core.crdt_engines + .get(&(db, tenant)) + .is_some_and(|engine| !engine.row_exists("users", "b")), + "a refused delta must not reach the CRDT state" + ); + assert!( + core.sparse + .load_crdt_dead_letters(db.as_u64(), tenant.as_u64()) + .expect("load") + .is_empty(), + "no dead-letter entry exists for the refused record" + ); + } + /// After a restart the queue holds the entries the live path stored. A /// record that rejects again keeps its one entry. #[test] diff --git a/nodedb/src/data/executor/handlers/control/crdt_apply/write_set.rs b/nodedb/src/data/executor/handlers/control/crdt_apply/write_set.rs index 2cd346b55..a3d106db6 100644 --- a/nodedb/src/data/executor/handlers/control/crdt_apply/write_set.rs +++ b/nodedb/src/data/executor/handlers/control/crdt_apply/write_set.rs @@ -10,6 +10,9 @@ use crate::data::executor::core_loop::CoreLoop; +/// The constraint name a delta outside its frame target violates. +const ONE_DOCUMENT_PER_DELTA: &str = "one_document_per_delta"; + impl CoreLoop { /// Enforce the one-document-per-delta contract: every row a validated delta /// wrote must be exactly the frame-declared `(collection, document_id)`. @@ -17,29 +20,34 @@ impl CoreLoop { /// A client that coalesced N document upserts into one delta, or tagged the /// frame with a synthetic id matching no written row, has no surrogate for /// the extra rows; materializing just one would silently drop the rest. - /// Returns a human-readable detail naming the offending rows so the caller - /// surfaces the violation instead of losing data. + /// The error is a CRDT constraint violation whose detail names the + /// offending rows, so the caller surfaces it instead of losing data. pub(crate) fn single_document_write_set( collection: &str, document_id: &str, write_set: &[(String, String)], - ) -> Result<(), String> { + ) -> crate::Result<()> { let foreign: Vec = write_set .iter() .filter(|(coll, row)| coll != collection || row != document_id) .map(|(coll, row)| format!("{coll}/{row}")) .collect(); if foreign.is_empty() { - Ok(()) - } else { - Err(format!( - "delta for {collection}/{document_id} wrote {} row(s) outside its frame \ - target: [{}]; a delta must carry exactly one document (cross-engine \ - identity binds one surrogate per delta)", - foreign.len(), - foreign.join(", ") - )) + return Ok(()); } + Err(crate::Error::Crdt( + nodedb_crdt::CrdtError::ConstraintViolation { + constraint: ONE_DOCUMENT_PER_DELTA.to_owned(), + collection: collection.to_owned(), + detail: format!( + "delta for {collection}/{document_id} wrote {} row(s) outside its frame \ + target: [{}]; a delta must carry exactly one document (cross-engine \ + identity binds one surrogate per delta)", + foreign.len(), + foreign.join(", ") + ), + }, + )) } } @@ -77,7 +85,15 @@ mod tests { ) .expect_err("multi-row delta must be rejected"); assert!( - err.contains("users/b"), + matches!( + &err, + crate::Error::Crdt(nodedb_crdt::CrdtError::ConstraintViolation { constraint, .. }) + if constraint == ONE_DOCUMENT_PER_DELTA + ), + "{err:?}" + ); + assert!( + err.to_string().contains("users/b"), "detail names the offending row: {err}" ); } @@ -93,13 +109,14 @@ mod tests { &ws(&[("entries", "u1"), ("entries", "u2")]), ) .expect_err("synthetic frame id must be rejected"); - assert!(err.contains("entries/u1") && err.contains("entries/u2")); + let detail = err.to_string(); + assert!(detail.contains("entries/u1") && detail.contains("entries/u2")); } #[test] fn foreign_collection_is_rejected() { let err = CoreLoop::single_document_write_set("users", "a", &ws(&[("orders", "a")])) .expect_err("row in a different collection must be rejected"); - assert!(err.contains("orders/a")); + assert!(err.to_string().contains("orders/a")); } } diff --git a/nodedb/src/data/executor/handlers/control/crdt_doc.rs b/nodedb/src/data/executor/handlers/control/crdt_doc.rs index 9c9b65725..df8a3cbec 100644 --- a/nodedb/src/data/executor/handlers/control/crdt_doc.rs +++ b/nodedb/src/data/executor/handlers/control/crdt_doc.rs @@ -5,6 +5,10 @@ //! server-side, then materializes the merged row into the sparse store with //! `EventSource::User` + text indexing so scans, secondary/spatial/vector //! indexes, AFTER triggers, and CDC all observe it. +//! +//! An aborted write leaves the Loro row and the sparse row in agreement. An +//! upsert writes the row's captured scalar state back to Loro. A delete +//! tombstones the Loro row only after its storage delete commits. use loro::LoroValue; use tracing::debug; @@ -13,12 +17,15 @@ use nodedb_types::Surrogate; use crate::bridge::envelope::{ErrorCode, Response}; use crate::data::executor::core_loop::CoreLoop; +use crate::data::executor::enforcement::chain_guard::{AbandonedWrite, abandon_write}; use crate::data::executor::handlers::point::apply_delete::PointDeleteParams; use crate::data::executor::handlers::returning_doc; use crate::data::executor::handlers::returning_rows; +use crate::data::executor::handlers::transaction::undo::UndoEntry; use crate::data::executor::handlers::transaction::undo::document_outcome::{ DocumentRow, push_delete_undo, }; +use crate::data::executor::handlers::transaction::undo::memory::abort_error; use crate::data::executor::task::ExecutionTask; use crate::engine::document::store::{RowIdentity, StorageKey}; use nodedb_physical::physical_plan::ReturningSpec; @@ -87,99 +94,111 @@ impl CoreLoop { .map(|(k, v)| (k.as_str(), super::convert::json_to_loro_value(v))) .collect(); - let materialized = { - let engine = match self.get_crdt_engine(task.request.database_id, tenant_id) { - Ok(e) => e, - Err(e) => { - return self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ); + // The sparse row is built from the merged Loro row, so Loro changes + // first. The row's prior scalar state is captured before that, and an + // abort writes it back. + let row_undo = match self.capture_crdt_row_undo( + task.request.database_id, + tenant_id, + collection, + document_id, + ) { + Ok(entry) => entry, + Err(e) => return self.response_error(task, e), + }; + let mutated = self + .get_crdt_engine(task.request.database_id, tenant_id) + .and_then(|engine| { + if partial { + engine.doc_set_fields(collection, document_id, &fields)?; + } else { + engine.doc_upsert(collection, document_id, &fields)?; } - }; - let res = if partial { - engine.doc_set_fields(collection, document_id, &fields) - } else { - engine.doc_upsert(collection, document_id, &fields) - }; - if let Err(e) = res { - return self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ); + Ok(Self::encode_crdt_row(engine, collection, document_id)) + }); + // A Loro write that fails part-way leaves some fields written. + let bytes = match mutated { + Ok(Some(bytes)) => bytes, + Ok(None) => { + let error = crate::Error::Internal { + detail: format!( + "crdt doc upsert: row {document_id} of '{collection}' has no \ + document body after its write" + ), + }; + return self.abandon_crdt_row(task, row_undo, error); } - Self::encode_crdt_row(engine, collection, document_id) + Err(e) => return self.abandon_crdt_row(task, row_undo, e), }; - let response = if let Some(bytes) = materialized { - // The Loro row and its sparse projection answer the same reads, so - // a projection that did not land fails the write. - if let Err(error) = self.materialize_document_write( - task, - super::crdt_materialize::CrdtMaterializeWrite { - tid: tenant_id.as_u64(), + // The Loro row and its sparse projection answer the same reads, so a + // projection that did not land fails the write and puts the Loro row + // back. + if let Err(error) = self.materialize_document_write( + task, + super::crdt_materialize::CrdtMaterializeWrite { + tid: tenant_id.as_u64(), + collection, + document_id, + surrogate, + value: &bytes, + index_text: true, + }, + ) { + return self.abandon_crdt_row(task, row_undo, error); + } + self.checkpoint_coordinator.mark_dirty("crdt", 1); + if let Some(spec) = returning { + // No strict schema: a CRDT row's stored body is whatever + // `encode_crdt_row` materialized from Loro, which is always + // MessagePack regardless of the collection's storage mode. + let doc = match returning_doc::from_stored( + &bytes, + &RowIdentity::from_user_key(document_id), + None, + &self.identity_column( + task.request.database_id.as_u64(), + tenant_id.as_u64(), collection, - document_id, - surrogate, - value: &bytes, - index_text: true, - }, + ), ) { - return self.response_error(task, error); - } - if let Some(spec) = returning { - // No strict schema: a CRDT row's stored body is whatever - // `encode_crdt_row` materialized from Loro, which is always - // MessagePack regardless of the collection's storage mode. - let doc = match returning_doc::from_stored( - &bytes, - &RowIdentity::from_user_key(document_id), - None, - ) { - Ok(doc) => doc, - Err(e) => return self.response_error(task, e), - }; - match returning_rows::build_rows_payload(spec, rls_filters, &[doc]) { - Ok(payload) => self.response_with_payload(task, payload), - Err(e) => { - return self.response_error( - task, - ErrorCode::Internal { - detail: format!("RETURNING encode: {e}"), - }, - ); - } - } - } else { - self.response_affected(task, 1) - } - } else if let Some(spec) = returning { - match returning_rows::build_rows_payload(spec, rls_filters, &[]) { + Ok(doc) => doc, + Err(e) => return self.response_error(task, e), + }; + match returning_rows::build_rows_payload(spec, rls_filters, &[doc]) { Ok(payload) => self.response_with_payload(task, payload), - Err(e) => { - return self.response_error( - task, - ErrorCode::Internal { - detail: format!("RETURNING encode: {e}"), - }, - ); - } + Err(e) => self.response_error(task, ErrorCode::from(e)), } } else { self.response_affected(task, 1) - }; - self.checkpoint_coordinator.mark_dirty("crdt", 1); - response + } + } + + /// Write the CRDT row's captured pre-image back after its write is + /// abandoned. Answers with the error the abort reports: `error`, or + /// `RollbackFailed` when the row did not go back. + fn abandon_crdt_row( + &mut self, + task: &ExecutionTask, + row_undo: UndoEntry, + error: crate::Error, + ) -> Response { + let undo = self.undo_memory_effects( + task.request.database_id.as_u64(), + task.request.tenant_id.as_u64(), + vec![row_undo], + ); + self.response_error(task, abort_error(error, undo)) } - /// Delete a document row: tombstone in the collection's Loro doc, then - /// remove it from the sparse store with the full index cascade + CDC delete - /// event (mirrors the point-delete apply path with `enforce = false`, since - /// the write was already admitted on its origin). + /// Delete a document row: remove it from the sparse store with the full + /// index cascade, then tombstone it in the collection's Loro doc, then + /// emit the CDC delete event. Mirrors the point-delete apply path with + /// `enforce = false`, since the write was already admitted on its origin. + /// + /// The Loro tombstone runs only after the storage commit. A tombstone + /// error after the commit reverses the committed delete through the undo + /// driver. pub(in crate::data::executor) fn execute_crdt_doc_delete( &mut self, task: &ExecutionTask, @@ -203,26 +222,9 @@ impl CoreLoop { return self.response_error(task, refusal); } let tenant_id = task.request.tenant_id; - { - let engine = match self.get_crdt_engine(task.request.database_id, tenant_id) { - Ok(e) => e, - Err(e) => { - return self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ); - } - }; - if let Err(e) = engine.doc_delete(collection, document_id) { - return self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ); - } + // Opening the engine can fail, so it opens before storage changes. + if let Err(e) = self.get_crdt_engine(task.request.database_id, tenant_id) { + return self.response_error(task, e); } let tid = tenant_id.as_u64(); @@ -233,12 +235,7 @@ impl CoreLoop { let txn = match self.sparse.begin_write() { Ok(txn) => txn, Err(e) => { - return self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ); + return self.response_error(task, ErrorCode::from(e)); } }; let outcome = match self.apply_point_delete( @@ -256,36 +253,49 @@ impl CoreLoop { ) { Ok(outcome) => outcome, Err(e) => { - return self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ); + return self.response_error(task, ErrorCode::from(e)); } }; if let Err(e) = txn.commit() { - return self.response_error( - task, - ErrorCode::Internal { + // The dropped txn reverses the durable writes only. The + // in-memory cascades are reversed here. + let e = abandon_write( + self, + AbandonedWrite::row( + task.request.database_id.as_u64(), + tid, + collection, + &storage_key, + ) + .undo(outcome.memory_undo), + crate::Error::DataPlane(ErrorCode::Internal { detail: format!("commit: {e}"), - }, + }), ); + return self.response_error(task, e); } self.checkpoint_coordinator.mark_dirty("sparse", 1); + let row = DocumentRow { + collection, + storage_key, + }; + // The storage delete is committed, so the Loro tombstone lands now. A + // tombstone error reverses the committed delete: Loro and storage + // then both still hold the row. + let tombstone = self + .get_crdt_engine(task.request.database_id, tenant_id) + .and_then(|engine| engine.doc_delete(collection, document_id)); + if let Err(e) = tombstone { + let mut undo = Vec::new(); + push_delete_undo(&mut undo, row, outcome); + let reversed = self.undo_memory_effects(task.request.database_id.as_u64(), tid, undo); + return self.response_error(task, abort_error(e, reversed)); + } + self.checkpoint_coordinator.mark_dirty("crdt", 1); let prior_value = outcome.prior_value.clone(); if self.recording_redo_undo() { let mut undo = Vec::new(); - push_delete_undo( - &mut undo, - DocumentRow { - database_id: task.request.database_id.as_u64(), - tid, - collection, - storage_key, - }, - outcome, - ); + push_delete_undo(&mut undo, row, outcome); self.record_redo_undo(undo); } @@ -293,25 +303,20 @@ impl CoreLoop { // removed, threading the pre-delete bytes through as `old_value` so // CDC/change-stream consumers observe the prior state. if let Some(prior_bytes) = prior_value.as_deref() { - let old_converted = self.resolve_event_payload( - task.request.database_id.as_u64(), - tid, - collection, - prior_bytes, - ); self.emit_document_delete_event( task, + tid, collection, RowIdentity::from_user_key(document_id), - Some(old_converted.as_deref().unwrap_or(prior_bytes)), + Some(prior_bytes), ); } // Project the pre-deletion row for RETURNING. `prior_value` is // only borrowed by the CDC emit above (via `.as_deref()`), so it is - // still available here; the user-visible `document_id` is injected as - // `id` exactly like PointDelete. - let response = if let Some(spec) = returning { + // still available here; the user-visible `document_id` fills the + // identity column exactly like PointDelete. + if let Some(spec) = returning { if let Some(prior_bytes) = prior_value.as_deref() { // No strict schema — see the upsert path: a CRDT row is // materialized as MessagePack in either storage mode. @@ -319,40 +324,190 @@ impl CoreLoop { prior_bytes, &RowIdentity::from_user_key(document_id), None, + &self.identity_column( + task.request.database_id.as_u64(), + tenant_id.as_u64(), + collection, + ), ) { Ok(doc) => doc, Err(e) => return self.response_error(task, e), }; match returning_rows::build_rows_payload(spec, rls_filters, &[doc]) { Ok(payload) => self.response_with_payload(task, payload), - Err(e) => { - return self.response_error( - task, - ErrorCode::Internal { - detail: format!("RETURNING encode: {e}"), - }, - ); - } + Err(e) => self.response_error(task, ErrorCode::from(e)), } } else { match returning_rows::build_rows_payload(spec, rls_filters, &[]) { Ok(payload) => self.response_with_payload(task, payload), - Err(e) => { - return self.response_error( - task, - ErrorCode::Internal { - detail: format!("RETURNING encode: {e}"), - }, - ); - } + Err(e) => self.response_error(task, ErrorCode::from(e)), } } } else { // No RETURNING: report what the delete actually removed. A tombstone // written over an already-absent document removes nothing. self.response_affected(task, u64::from(prior_value.is_some())) + } + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::bridge::envelope::Status; + use crate::data::executor::core_loop::tests::{make_core_with_dir, make_default_task}; + use crate::engine::document::store::{CollectionConfig, IndexPath}; + use crate::types::{DatabaseId, TenantId}; + + const DB: u64 = 0; + const TID: u64 = 1; + const COLL: &str = "notes"; + const DOC: &str = "n1"; + const SURROGATE: Surrogate = Surrogate(7); + + fn upsert(core: &mut CoreLoop, task: &ExecutionTask, fields_json: &str) -> Response { + core.execute_crdt_doc_upsert( + task, + CrdtDocUpsert { + collection: COLL, + document_id: DOC, + fields_json, + surrogate: SURROGATE, + partial: false, + returning: None, + rls_filters: &[], + }, + ) + } + + fn delete(core: &mut CoreLoop, task: &ExecutionTask) -> Response { + core.execute_crdt_doc_delete( + task, + CrdtDocDelete { + collection: COLL, + document_id: DOC, + surrogate: Some(SURROGATE), + returning: None, + rls_filters: &[], + }, + ) + } + + fn crdt_row(core: &mut CoreLoop) -> Option { + core.get_crdt_engine(DatabaseId::DEFAULT, TenantId::new(TID)) + .expect("engine") + .read_row(COLL, DOC) + } + + fn i64_field(row: &loro::LoroValue, key: &str) -> Option { + let loro::LoroValue::Map(map) = row else { + return None; }; - self.checkpoint_coordinator.mark_dirty("crdt", 1); - response + match map.get(key) { + Some(loro::LoroValue::I64(n)) => Some(*n), + _ => None, + } + } + + /// Make the storage step of the next write fail: the collection gets an + /// index path, and its stored row an undecodable body. A write that reads + /// the prior row to diff its index entries then refuses. + fn break_stored_row(core: &mut CoreLoop) { + let mut config = CollectionConfig::new(COLL); + config.index_paths.push(IndexPath::new("a")); + core.doc_configs.insert( + (DatabaseId::DEFAULT, TenantId::new(TID), COLL.to_string()), + config, + ); + let key = StorageKey::for_surrogate(SURROGATE); + core.sparse + .put(DB, TID, COLL, &key, &[0xc1]) + .expect("overwrite the stored row"); + core.doc_cache.invalidate(DB, TID, COLL, &key); + } + + /// A CRDT delete whose storage step fails leaves the row readable from + /// the CRDT state, as storage still holds it. + #[test] + fn a_refused_crdt_delete_leaves_the_crdt_row() { + let dir = tempfile::tempdir().expect("tempdir"); + let (mut core, _req, _resp) = make_core_with_dir(dir.path()); + let task = make_default_task(); + assert_eq!( + upsert(&mut core, &task, r#"{"a":1,"b":2}"#).status, + Status::Ok + ); + break_stored_row(&mut core); + + let resp = delete(&mut core, &task); + assert_eq!(resp.status, Status::Error, "the storage step must refuse"); + assert!( + !matches!( + resp.error_code.as_deref(), + Some(ErrorCode::RollbackFailed { .. }) + ), + "the refusal must reverse cleanly, got {:?}", + resp.error_code + ); + let row = crdt_row(&mut core).expect("the CRDT row must survive the refused delete"); + assert_eq!(i64_field(&row, "a"), Some(1)); + assert_eq!(i64_field(&row, "b"), Some(2)); + assert!( + core.sparse + .get(DB, TID, COLL, &StorageKey::for_surrogate(SURROGATE)) + .expect("read back") + .is_some(), + "the refused delete leaves the stored row" + ); + } + + /// A committed CRDT delete removes the row from both the CRDT state and + /// storage. + #[test] + fn a_crdt_delete_removes_the_row_from_both_stores() { + let dir = tempfile::tempdir().expect("tempdir"); + let (mut core, _req, _resp) = make_core_with_dir(dir.path()); + let task = make_default_task(); + assert_eq!(upsert(&mut core, &task, r#"{"a":1}"#).status, Status::Ok); + + assert_eq!(delete(&mut core, &task).status, Status::Ok); + assert!(crdt_row(&mut core).is_none()); + assert!( + core.sparse + .get(DB, TID, COLL, &StorageKey::for_surrogate(SURROGATE)) + .expect("read back") + .is_none() + ); + } + + /// A CRDT upsert whose sparse projection fails puts the Loro row back to + /// its fields before the write. + #[test] + fn a_refused_crdt_upsert_puts_the_crdt_row_back() { + let dir = tempfile::tempdir().expect("tempdir"); + let (mut core, _req, _resp) = make_core_with_dir(dir.path()); + let task = make_default_task(); + assert_eq!( + upsert(&mut core, &task, r#"{"a":1,"b":2}"#).status, + Status::Ok + ); + let before = crdt_row(&mut core); + break_stored_row(&mut core); + + let resp = upsert(&mut core, &task, r#"{"a":5,"c":3}"#); + assert_eq!(resp.status, Status::Error, "the projection must refuse"); + assert!( + !matches!( + resp.error_code.as_deref(), + Some(ErrorCode::RollbackFailed { .. }) + ), + "the refusal must reverse cleanly, got {:?}", + resp.error_code + ); + assert_eq!( + crdt_row(&mut core), + before, + "the refused upsert leaves the CRDT row as it was" + ); } } diff --git a/nodedb/src/data/executor/wal_replay/crdt.rs b/nodedb/src/data/executor/wal_replay/crdt.rs index da9955695..17c75acbd 100644 --- a/nodedb/src/data/executor/wal_replay/crdt.rs +++ b/nodedb/src/data/executor/wal_replay/crdt.rs @@ -47,7 +47,7 @@ impl CoreLoop { tombstones: &nodedb_wal::TombstoneSet, ) { use nodedb_wal::record::RecordType; - use tracing::{error, warn}; + use tracing::{debug, warn}; let mut replayed = 0usize; @@ -97,21 +97,17 @@ impl CoreLoop { } let database_id = crate::types::DatabaseId::new(record.header.database_id); - if let Some(signing) = payload.signing { - let Some(provenance) = payload.provenance.as_ref() else { - warn!(core = self.core_id, tenant = tid.as_u64(), %collection, "authenticated CRDT WAL record has no provenance"); - continue; - }; - if signing.auth_user_id == 0 - || signing.auth_device_id == 0 - || signing.auth_seq_no == 0 - || provenance.producer_id != signing.auth_device_id - || provenance.seq != signing.auth_seq_no - || (signing.required && signing.delta_signature == [0; 32]) - { - warn!(core = self.core_id, tenant = tid.as_u64(), %collection, "authenticated CRDT WAL admission metadata is inconsistent"); - continue; - } + // The live apply never reads this metadata, so it applied the + // delta. A record its writer cannot produce halts replay: skipping + // it drops a committed delta. + if let Err(error) = signed_record_metadata(&payload, record.header.lsn) { + self.replay_record_unapplied( + "crdt", + "signing_metadata", + record.header.lsn, + &error.to_string(), + ); + continue; } if let Some(expected) = payload.expected_frontier_digest { let actual = self @@ -127,65 +123,64 @@ impl CoreLoop { ) }); if actual != expected { - if payload.signing.is_some() - && let Some(provenance) = payload.provenance.as_ref() - { - self.sync_commit(provenance); - } - warn!( + // Replay rebuilds the collection in LSN order, so the + // frontier here is the one the live apply fenced against. + // The live apply refused the delta, applied nothing, and + // held the stream's mark so the sender's re-push is + // admitted. Replay skips it and holds the mark the same way. + debug!( core = self.core_id, tenant = tid.as_u64(), %collection, - "skipping stale fenced CRDT WAL delta during replay" + lsn = record.header.lsn, + "replay skips a CRDT delta its live apply refused at the frontier fence" ); continue; } } + // Opening the engine restores the tenant's stored dead-letter + // entries. A record one of them names was rejected by its live + // apply, which stored the entry and changed no state, so it + // replays as that rejection. Applying it again would enqueue a + // second entry, which a queue full of restored entries refuses. + let already_rejected = match self.get_crdt_engine(database_id, tid) { + Ok(engine) => engine.dead_letter_recorded(record.header.lsn), + Err(error) => { + self.replay_record_unapplied( + "crdt", + "engine_open", + record.header.lsn, + &error.to_string(), + ); + continue; + } + }; + let projection = match &payload.target { + _ if already_rejected => None, crate::wal::CrdtDeltaTarget::Document { document_id, surrogate, } => { let surrogate = *surrogate; let applied = match self.get_crdt_engine(database_id, tid) { - Ok(engine) => match payload.signing { - Some(signing) => engine.apply_committed_delta_authenticated( - collection, - &payload.bytes, - crate::engine::crdt::tenant_state::ApplyTarget::Document { - document_id, - surrogate, - }, - payload.peer_id, - crate::engine::crdt::tenant_state::DeltaSigningAdmission { - auth: nodedb_crdt::CrdtAuthContext { - user_id: signing.auth_user_id, - device_id: signing.auth_device_id, - seq_no: signing.auth_seq_no, - delta_signature: signing.delta_signature, - ..nodedb_crdt::CrdtAuthContext::default() - }, - required: signing.required, - preverified: true, - }, - ), - None => engine.apply_committed_delta_validated( - collection, - &payload.bytes, - crate::engine::crdt::tenant_state::ApplyTarget::Document { - document_id, - surrogate, - }, - payload.peer_id, - ), - }, - Err(e) => { - warn!( - core = self.core_id, - tenant = tid.as_u64(), - error = %e, - "failed to create CRDT engine during WAL replay" + Ok(engine) => engine.apply_committed_delta_authenticated( + collection, + &payload.bytes, + crate::engine::crdt::tenant_state::ApplyTarget::Document { + document_id, + surrogate, + }, + payload.peer_id, + replayed_admission(&payload), + ), + Err(error) => { + self.replay_record_unapplied( + "crdt", + "engine_open", + record.header.lsn, + &error.to_string(), ); continue; } @@ -195,31 +190,29 @@ impl CoreLoop { write_set, .. } => { - if let Err(detail) = + // The engine refuses a delta writing any row but + // its target before it installs, so a clean apply + // wrote the target alone. A wider write set is an + // engine invariant broken after Loro installed the + // delta: skipping its projection drops rows. + if let Err(error) = Self::single_document_write_set(collection, document_id, &write_set) { - self.note_replay_write_lsn( - record.header.database_id, - tid.as_u64(), - collection, - None, + self.replay_record_unapplied( + "crdt", + "one_document_contract", record.header.lsn, - ); - warn!( - core = self.core_id, - tenant = tid.as_u64(), - %collection, - %document_id, - %detail, - "CRDT WAL delta violates one-document replay contract" + &error.to_string(), ); continue; } let Some(engine) = self.crdt_engines.get(&(database_id, tid)) else { - warn!( - core = self.core_id, - tenant = tid.as_u64(), - "CRDT engine disappeared during WAL replay" + self.replay_record_unapplied( + "crdt", + "engine_missing", + record.header.lsn, + "the engine that applied the delta is gone before its row \ + projected", ); continue; }; @@ -236,31 +229,68 @@ impl CoreLoop { // The committed record remains a deterministic no-op // whose collection floor advances on every replica. warn!(core = self.core_id, tenant = tid.as_u64(), %collection, %reason, "CRDT WAL delta rejected during replay"); - self.store_replayed_dead_letter(database_id, tid, record.header.lsn); + if !self.store_replayed_dead_letter(database_id, tid, record.header.lsn) { + continue; + } None } + crate::engine::crdt::tenant_state::ValidatedApplyOutcome::DeadLetterRefused { + violation, + error, + } => { + // Rejected, and the queue refused its entry, so + // nothing but this record holds the delta. Replay + // stops at the record: no checkpoint covers it, and + // boot refuses until the queue takes the entry. + self.replay_record_unapplied( + "crdt", + "dead_letter_refused", + record.header.lsn, + &format!( + "delta for {collection} violates {violation}, and the \ + dead-letter queue refused it: {error}" + ), + ); + continue; + } crate::engine::crdt::tenant_state::ValidatedApplyOutcome::Malformed => { + // The bytes, their rows, or a missing required + // signature decide this, the same at replay as live. + // The live apply refused the delta, applied nothing, + // and advanced a sender's mark past it. Replay skips + // it and advances the mark the same way. if let Some(provenance) = payload.provenance.as_ref() { self.sync_commit(provenance); } - warn!(core = self.core_id, tenant = tid.as_u64(), %collection, "CRDT WAL delta malformed during replay"); + debug!(core = self.core_id, tenant = tid.as_u64(), %collection, lsn = record.header.lsn, "replay skips a CRDT delta its live apply refused as malformed"); continue; } crate::engine::crdt::tenant_state::ValidatedApplyOutcome::PendingDependencies => { - // Records replay in LSN order, so a delta whose - // causal predecessors are missing means the log is - // inconsistent with this collection's document — - // not a routine skip. Recovery continues so the - // node still starts, but this is an error-level - // event: the row it carried is NOT present. - error!( + // Replay rebuilds the collection in LSN order, so + // the predecessors are absent here exactly when they + // were absent live. The live apply applied nothing + // and held the sender's mark for the re-push. + // Replay skips the delta and holds the mark the + // same way. + debug!( core = self.core_id, tenant = tid.as_u64(), %collection, %document_id, lsn = record.header.lsn, - "CRDT WAL delta depends on operations absent from this \ - collection's document; row not recovered" + "replay skips a CRDT delta its live apply held for missing \ + predecessors" + ); + continue; + } + crate::engine::crdt::tenant_state::ValidatedApplyOutcome::CandidateUnavailable { error } => { + // This node failed, not the delta: skipping it + // drops a delta the live apply may have applied. + self.replay_record_unapplied( + "crdt", + "candidate_unavailable", + record.header.lsn, + &format!("no apply candidate for {collection}: {error}"), ); continue; } @@ -271,11 +301,12 @@ impl CoreLoop { // It rebuilds Loro state only; each row's projection comes // from the document records that carry its surrogate. match self.get_crdt_engine(database_id, tid) { - Ok(engine) => match engine.apply_committed_delta_validated( + Ok(engine) => match engine.apply_committed_delta_authenticated( collection, &payload.bytes, crate::engine::crdt::tenant_state::ApplyTarget::Collection, payload.peer_id, + replayed_admission(&payload), ) { crate::engine::crdt::tenant_state::ValidatedApplyOutcome::Clean { .. @@ -284,27 +315,55 @@ impl CoreLoop { reason, ) => { warn!(core = self.core_id, tenant = tid.as_u64(), %collection, %reason, "CRDT WAL snapshot import rejected during replay"); - self.store_replayed_dead_letter(database_id, tid, record.header.lsn); + if !self.store_replayed_dead_letter(database_id, tid, record.header.lsn) { + continue; + } None } + // See the document arm: nothing but this record + // holds the delta, so replay stops at it. + crate::engine::crdt::tenant_state::ValidatedApplyOutcome::DeadLetterRefused { + violation, + error, + } => { + self.replay_record_unapplied( + "crdt", + "dead_letter_refused", + record.header.lsn, + &format!( + "snapshot import for {collection} violates {violation}, \ + and the dead-letter queue refused it: {error}" + ), + ); + continue; + } + // See the document arm: the live import refused + // the snapshot on the same state, so replay skips it. crate::engine::crdt::tenant_state::ValidatedApplyOutcome::Malformed => { - if let Some(provenance) = payload.provenance.as_ref() { - self.sync_commit(provenance); - } - warn!(core = self.core_id, tenant = tid.as_u64(), %collection, "CRDT WAL snapshot import malformed during replay"); + debug!(core = self.core_id, tenant = tid.as_u64(), %collection, lsn = record.header.lsn, "replay skips a CRDT snapshot import its live apply refused as malformed"); continue; } crate::engine::crdt::tenant_state::ValidatedApplyOutcome::PendingDependencies => { - // Committed record whose causal predecessors are - // absent from this collection's document: the row - // did NOT apply. Loud, and not acknowledged as a - // clean replay. - error!(core = self.core_id, tenant = tid.as_u64(), %collection, lsn = record.header.lsn, "CRDT WAL snapshot import depends on operations absent from this collection's document; state not recovered"); + debug!(core = self.core_id, tenant = tid.as_u64(), %collection, lsn = record.header.lsn, "replay skips a CRDT snapshot import its live apply refused for missing predecessors"); + continue; + } + crate::engine::crdt::tenant_state::ValidatedApplyOutcome::CandidateUnavailable { error } => { + self.replay_record_unapplied( + "crdt", + "candidate_unavailable", + record.header.lsn, + &format!("no apply candidate for {collection}: {error}"), + ); continue; } }, - Err(e) => { - warn!(core = self.core_id, tenant = tid.as_u64(), error = %e, "failed to create CRDT engine during WAL replay"); + Err(error) => { + self.replay_record_unapplied( + "crdt", + "engine_open", + record.header.lsn, + &error.to_string(), + ); continue; } } @@ -364,6 +423,71 @@ impl CoreLoop { } } +/// Check the admission metadata of a signed record against what its writer +/// stamps. The writer copies the session's producer id and the frame's seq +/// into both the provenance and the signing fields, so a signed record +/// without provenance, or with fields that disagree, is not one it wrote. +/// Such a record at `lsn` is a corrupt WAL record. +fn signed_record_metadata( + payload: &crate::wal::CrdtDeltaWalPayload, + lsn: u64, +) -> crate::Result<()> { + let corrupt = + |detail: String| crate::Error::Wal(nodedb_wal::WalError::CorruptRecord { lsn, detail }); + let Some(signing) = payload.signing else { + return Ok(()); + }; + let Some(provenance) = payload.provenance.as_ref() else { + return Err(corrupt(format!( + "signed delta for {} carries no provenance", + payload.collection + ))); + }; + if provenance.producer_id != signing.auth_device_id || provenance.seq != signing.auth_seq_no { + return Err(corrupt(format!( + "signed delta for {} names producer {} seq {} in its provenance and device {} seq {} \ + in its signing fields", + payload.collection, + provenance.producer_id, + provenance.seq, + signing.auth_device_id, + signing.auth_seq_no + ))); + } + Ok(()) +} + +/// The signing admission a committed record replays under: the one its live +/// apply ran under. +/// +/// A signed record carries the signing fields the session admitted, already +/// verified, so it replays preverified. An absent required signature still +/// refuses it, as it did live. An unsigned record replays the live local +/// apply's admission: no signature, checked against the collection's signing +/// policy. +fn replayed_admission( + payload: &crate::wal::CrdtDeltaWalPayload, +) -> crate::engine::crdt::tenant_state::DeltaSigningAdmission { + match payload.signing { + Some(signing) => crate::engine::crdt::tenant_state::DeltaSigningAdmission { + auth: nodedb_crdt::CrdtAuthContext { + user_id: signing.auth_user_id, + device_id: signing.auth_device_id, + seq_no: signing.auth_seq_no, + delta_signature: signing.delta_signature, + ..nodedb_crdt::CrdtAuthContext::default() + }, + required: signing.required, + preverified: true, + }, + None => crate::engine::crdt::tenant_state::DeltaSigningAdmission { + auth: nodedb_crdt::CrdtAuthContext::default(), + required: false, + preverified: false, + }, + } +} + #[cfg(test)] mod crdt_replay_tests { use super::CoreLoop; @@ -483,8 +607,12 @@ mod crdt_replay_tests { ); } + /// The live apply refuses a stale fenced frame before admission, so the + /// sender's mark stays put and its re-push at the same seq is admitted. + /// Replay holds the mark the same way: advancing it would turn the + /// re-push after a restart into a `Duplicate` and drop the write. #[test] - fn replay_stale_v4_restores_authenticated_sequence_watermark() { + fn replay_stale_v4_holds_authenticated_sequence_watermark() { let tid = TenantId::new(9); let provenance = nodedb_types::sync::wire::SyncProvenance { producer_id: 77, @@ -529,10 +657,249 @@ mod crdt_replay_tests { let mut h = make_core(0); h.core .replay_crdt_wal(&[record], 1, &nodedb_wal::TombstoneSet::new()); - assert!(matches!( - h.core.sync_admit(&provenance), - crate::data::executor::sync_gate::SyncAdmit::Duplicate - )); + assert_eq!( + h.core + .sync_hwm_value(provenance.producer_id, provenance.stream_id), + 0, + "a stale fenced frame leaves the mark where the live apply left it" + ); + assert!(!h.core.is_fail_stopped(), "a stale fence is no halt"); + } + + /// A signed record of `secure_notes/doc` at LSN 5, unfenced. + fn signed_record( + tid: TenantId, + provenance: Option, + signing: crate::wal::CrdtDeltaSigning, + ) -> nodedb_wal::WalRecord { + let state = nodedb_crdt::state::CrdtState::new(77).expect("state"); + state + .upsert( + "secure_notes", + "doc", + &[("body", LoroValue::String("signed".into()))], + ) + .expect("upsert"); + let payload = crate::wal::CrdtDeltaWalPayload::new( + state.export_snapshot().expect("snapshot"), + "secure_notes".into(), + provenance, + None, + document_target("doc", 1), + ) + .with_signing(signing); + nodedb_wal::WalRecord::new(nodedb_wal::WalRecordArgs { + record_type: RecordType::CrdtDelta as u32, + lsn: 5, + tenant_id: tid.as_u64(), + vshard_id: 0, + database_id: DatabaseId::DEFAULT.as_u64(), + payload: payload.encode().expect("encode"), + encryption_key: None, + preamble_bytes: None, + }) + .expect("record") + } + + fn session_provenance() -> nodedb_types::sync::wire::SyncProvenance { + nodedb_types::sync::wire::SyncProvenance { + producer_id: 77, + epoch: 3, + stream_id: 5, + seq: 11, + } + } + + fn signing(user: u64, signature: [u8; 32], required: bool) -> crate::wal::CrdtDeltaSigning { + crate::wal::CrdtDeltaSigning { + auth_user_id: user, + auth_device_id: 77, + auth_seq_no: 11, + delta_signature: signature, + required, + } + } + + fn doc_replayed(h: &mut CoreHarness, tid: TenantId) -> bool { + h.core + .get_crdt_engine(DatabaseId::DEFAULT, tid) + .expect("engine") + .row_exists("secure_notes", "doc") + } + + /// The writer stamps provenance on every signed record and the live apply + /// applied it, so a signed record without provenance halts replay. + #[test] + fn a_signed_record_without_provenance_halts_replay() { + let tid = TenantId::new(21); + let mut h = make_core(0); + h.core.replay_crdt_wal( + &[signed_record(tid, None, signing(42, [7; 32], true))], + 1, + &nodedb_wal::TombstoneSet::new(), + ); + assert!(h.core.is_fail_stopped()); + assert!(h.core.replay_halt_error().is_some()); + } + + /// The writer copies one producer id and seq into the provenance and the + /// signing fields, so a record where they disagree halts replay. + #[test] + fn a_signed_record_with_disagreeing_metadata_halts_replay() { + let tid = TenantId::new(22); + let mut provenance = session_provenance(); + provenance.seq = 12; + let mut h = make_core(0); + h.core.replay_crdt_wal( + &[signed_record( + tid, + Some(provenance), + signing(42, [7; 32], true), + )], + 1, + &nodedb_wal::TombstoneSet::new(), + ); + assert!(h.core.is_fail_stopped()); + assert!(!doc_replayed(&mut h, tid)); + } + + /// The live apply reads no user id from the signing fields, so a record + /// whose user id is zero applied live and applies on replay. + #[test] + fn a_signed_record_with_a_zero_user_id_replays() { + let tid = TenantId::new(23); + let provenance = session_provenance(); + let mut h = make_core(0); + h.core.replay_crdt_wal( + &[signed_record( + tid, + Some(provenance.clone()), + signing(0, [7; 32], true), + )], + 1, + &nodedb_wal::TombstoneSet::new(), + ); + assert!(!h.core.is_fail_stopped()); + assert!(doc_replayed(&mut h, tid)); + assert_eq!( + h.core + .sync_hwm_value(provenance.producer_id, provenance.stream_id), + provenance.seq + ); + } + + /// A required signature that is absent made the live apply refuse the + /// delta as malformed and advance the mark. Replay reaches the same + /// refusal: nothing applies, the mark advances, and replay continues. + #[test] + fn a_record_missing_its_required_signature_replays_as_its_live_refusal() { + let tid = TenantId::new(24); + let provenance = session_provenance(); + let mut h = make_core(0); + h.core.replay_crdt_wal( + &[signed_record( + tid, + Some(provenance.clone()), + signing(42, [0; 32], true), + )], + 1, + &nodedb_wal::TombstoneSet::new(), + ); + assert!(!h.core.is_fail_stopped()); + assert!(!doc_replayed(&mut h, tid)); + assert_eq!( + h.core + .sync_hwm_value(provenance.producer_id, provenance.stream_id), + provenance.seq + ); + } + + /// An unsigned `secure_notes/doc` record at LSN 5 that carries `delta` + /// under the session's provenance. + fn provenance_record(tid: TenantId, delta: Vec) -> nodedb_wal::WalRecord { + let payload = crate::wal::CrdtDeltaWalPayload::new( + delta, + "secure_notes".into(), + Some(session_provenance()), + None, + document_target("doc", 1), + ); + nodedb_wal::WalRecord::new(nodedb_wal::WalRecordArgs { + record_type: RecordType::CrdtDelta as u32, + lsn: 5, + tenant_id: tid.as_u64(), + vshard_id: 0, + database_id: DatabaseId::DEFAULT.as_u64(), + payload: payload.encode().expect("encode"), + encryption_key: None, + preamble_bytes: None, + }) + .expect("record") + } + + /// Delta bytes that do not decode made the live apply refuse the delta + /// as malformed and advance the mark. Replay reaches the same refusal: + /// nothing applies, the mark advances, and replay continues. + #[test] + fn a_malformed_delta_replays_as_its_live_refusal() { + let tid = TenantId::new(25); + let provenance = session_provenance(); + let mut h = make_core(0); + h.core.replay_crdt_wal( + &[provenance_record(tid, b"not a loro delta".to_vec())], + 1, + &nodedb_wal::TombstoneSet::new(), + ); + assert!(!h.core.is_fail_stopped()); + assert!(!doc_replayed(&mut h, tid)); + assert_eq!( + h.core + .sync_hwm_value(provenance.producer_id, provenance.stream_id), + provenance.seq + ); + } + + /// A delta whose predecessors were absent made the live apply hold the + /// mark for the sender's re-push. Replay in LSN order meets the same + /// absence: nothing applies, the mark holds, and replay continues. The + /// re-push lands at a later LSN once the predecessors arrived. + #[test] + fn a_delta_missing_its_predecessors_replays_as_its_live_hold() { + let tid = TenantId::new(26); + let source = nodedb_crdt::state::CrdtState::new(77).expect("state"); + source + .upsert( + "secure_notes", + "doc", + &[("body", LoroValue::String("base".into()))], + ) + .expect("base write"); + let base_vv = source.oplog_version_vector(); + source + .upsert( + "secure_notes", + "doc", + &[("body", LoroValue::String("next".into()))], + ) + .expect("next write"); + let dependent = source + .export_updates_since(&base_vv) + .expect("dependent delta"); + let provenance = session_provenance(); + let mut h = make_core(0); + h.core.replay_crdt_wal( + &[provenance_record(tid, dependent)], + 1, + &nodedb_wal::TombstoneSet::new(), + ); + assert!(!h.core.is_fail_stopped(), "a held delta is no halt"); + assert!(!doc_replayed(&mut h, tid)); + assert_eq!( + h.core + .sync_hwm_value(provenance.producer_id, provenance.stream_id), + 0, + "the mark stays where the live apply held it" + ); } /// The WAL writer's own record for an apply carries the row's surrogate, @@ -843,4 +1210,157 @@ mod crdt_replay_tests { entries ); } + + /// A core whose `users` collection refuses a second row with one email. + fn unique_email_core(tid: TenantId) -> CoreHarness { + let mut h = make_core(0); + let engine = h + .core + .get_crdt_engine(DatabaseId::DEFAULT, tid) + .expect("engine"); + assert!(engine.set_collection_constraints( + "users", + 1, + vec![nodedb_crdt::Constraint { + name: "users_email_unique".into(), + collection: "users".into(), + field: "email".into(), + kind: nodedb_crdt::ConstraintKind::Unique, + }], + )); + engine + .set_collection_policy_typed("users", nodedb_crdt::policy::CollectionPolicy::strict()); + h + } + + fn users_write_lsn(h: &CoreHarness, tid: TenantId) -> Option { + h.core.write_index.collection_write_lsn( + &crate::data::executor::core_loop::write_index::CollKey { + db: DatabaseId::DEFAULT, + tenant: tid, + collection: Box::from("users"), + }, + ) + } + + /// A replayed rejection the full queue refuses stops replay at its + /// record: only the record holds the delta, so it stays uncovered. + #[test] + fn a_replayed_rejection_the_full_queue_refuses_halts_replay() { + let tid = TenantId::new(7); + let mut h = unique_email_core(tid); + let tombstones = nodedb_wal::TombstoneSet::new(); + h.core + .replay_crdt_wal(&[email_record(tid, "a", "x@y.com", 2, 19)], 1, &tombstones); + let capacity = h + .core + .get_crdt_engine(DatabaseId::DEFAULT, tid) + .expect("engine") + .fill_dead_letter_queue_for_test(); + + h.core + .replay_crdt_wal(&[email_record(tid, "b", "x@y.com", 3, 20)], 1, &tombstones); + + assert!(h.core.replay_halt_error().is_some(), "replay stopped"); + let engine = h + .core + .get_crdt_engine(DatabaseId::DEFAULT, tid) + .expect("engine"); + assert_eq!(engine.dead_letters().count(), capacity, "nothing queued"); + assert!(!engine.row_exists("users", "b")); + assert!( + h.core + .sparse + .load_crdt_dead_letters(DatabaseId::DEFAULT.as_u64(), tid.as_u64()) + .expect("load") + .is_empty() + ); + assert_eq!( + users_write_lsn(&h, tid), + Some(crate::types::Lsn::new(19)), + "the refused record is not counted as replayed" + ); + } + + /// A replayed rejection whose entry the store refuses stops replay at + /// its record, and the queue keeps no entry storage lacks. + #[test] + fn a_replayed_rejection_the_store_refuses_halts_replay() { + let tid = TenantId::new(7); + let mut h = unique_email_core(tid); + let tombstones = nodedb_wal::TombstoneSet::new(); + h.core + .replay_crdt_wal(&[email_record(tid, "a", "x@y.com", 2, 19)], 1, &tombstones); + h.core.sparse.break_crdt_dead_letter_table_for_test(); + + h.core + .replay_crdt_wal(&[email_record(tid, "b", "x@y.com", 3, 20)], 1, &tombstones); + + assert!(h.core.replay_halt_error().is_some(), "replay stopped"); + let engine = h + .core + .get_crdt_engine(DatabaseId::DEFAULT, tid) + .expect("engine"); + assert_eq!(engine.dead_letters().count(), 0, "nothing queued"); + assert_eq!(users_write_lsn(&h, tid), Some(crate::types::Lsn::new(19))); + } + + /// A record whose stored entry the engine restored replays as that + /// rejection, so a queue full of restored entries does not stop replay. + #[test] + fn a_replayed_record_with_a_restored_entry_replays_under_a_full_queue() { + let tid = TenantId::new(7); + let mut h = unique_email_core(tid); + let tombstones = nodedb_wal::TombstoneSet::new(); + let records = [ + email_record(tid, "a", "x@y.com", 2, 19), + email_record(tid, "b", "x@y.com", 3, 20), + ]; + h.core.replay_crdt_wal(&records, 1, &tombstones); + h.core + .get_crdt_engine(DatabaseId::DEFAULT, tid) + .expect("engine") + .fill_dead_letter_queue_for_test(); + + h.core.replay_crdt_wal(&records[1..], 1, &tombstones); + + assert!(h.core.replay_halt_error().is_none(), "replay continued"); + let engine = h + .core + .get_crdt_engine(DatabaseId::DEFAULT, tid) + .expect("engine"); + assert_eq!( + engine + .dead_letters() + .filter(|entry| entry.source_lsn == Some(20)) + .count(), + 1 + ); + assert_eq!(users_write_lsn(&h, tid), Some(crate::types::Lsn::new(20))); + } + + /// A stored entry that does not decode fails the engine open, and replay + /// stops at the first record that needs the engine. + #[test] + fn an_undecodable_stored_entry_halts_replay() { + let tid = TenantId::new(8); + let mut h = make_core(0); + h.core.sparse.put_raw_crdt_dead_letter_for_test( + DatabaseId::DEFAULT.as_u64(), + tid.as_u64(), + 4, + b"\xc1", + ); + + let record = make_crdt_record(0, tid, 0, "notes", "row1"); + h.core + .replay_crdt_wal(&[record], 1, &nodedb_wal::TombstoneSet::new()); + + assert!(h.core.replay_halt_error().is_some(), "replay stopped"); + assert!( + !h.core + .crdt_engines + .contains_key(&(DatabaseId::DEFAULT, tid)) + ); + } } diff --git a/nodedb/src/diag/context/crdt.rs b/nodedb/src/diag/context/crdt.rs index 6db3a378f..f40aa8c67 100644 --- a/nodedb/src/diag/context/crdt.rs +++ b/nodedb/src/diag/context/crdt.rs @@ -1,6 +1,7 @@ // SPDX-License-Identifier: BUSL-1.1 -//! Forensic payloads for CRDT history post-apply capture sites. +//! Forensic payloads for CRDT capture sites: history post-apply, and rejected +//! deltas the dead-letter queue refused. //! //! COMPACT HISTORY discards oplog entries on every node. A node that misses //! the compaction keeps a history its peers reclaimed, so a read at an old @@ -54,10 +55,69 @@ impl DomainContext for HistoryCompactionNotApplied<'_> { } } +/// A constraint-rejected CRDT delta whose dead-letter entry the queue +/// refused, so no node-local record of the rejection exists. +pub(in crate::diag) struct CrdtDeadLetterNotEnqueued<'a> { + pub tenant_id: u64, + /// Collection the delta wrote. + pub collection: &'a str, + /// Constraint the delta violated. + pub constraint: &'a str, + /// What failed, without the per-occurrence detail. + pub error_class: &'a str, +} + +impl DomainContext for CrdtDeadLetterNotEnqueued<'_> { + fn domain_kind(&self) -> &'static str { + "nodedb.crdt_dead_letter_not_enqueued" + } + + fn grouping_key(&self) -> String { + // The error class names the cause. Tenant, collection and constraint + // are the occurrence, so a full queue files one report. + format!("cause={}", self.error_class) + } + + fn to_json(&self) -> Value { + json!({ + "tenant_id": self.tenant_id, + "collection": self.collection, + "constraint": self.constraint, + "error_class": self.error_class, + "why_reported": "the delta violated a constraint and was refused, but the \ + dead-letter queue did not take its entry. The apply is \ + refused as an error so the sender keeps the delta, and every \ + further rejected delta is refused the same way until the \ + queue has room", + "operator_action": "inspect and drain the tenant's dead-letter queue, then let \ + the sender re-push. A full queue means rejected deltas \ + arrive faster than they are resolved", + }) + } +} + #[cfg(test)] mod tests { use super::*; + #[test] + fn dead_letter_grouping_ignores_the_occurrence() { + let first = CrdtDeadLetterNotEnqueued { + tenant_id: 1, + collection: "users", + constraint: "users_email_unique", + error_class: "dead-letter queue full", + }; + let second = CrdtDeadLetterNotEnqueued { + tenant_id: 7, + collection: "orders", + constraint: "orders_fk", + ..first + }; + assert_eq!(first.grouping_key(), second.grouping_key()); + assert_eq!(first.grouping_key(), "cause=dead-letter queue full"); + } + fn sample() -> HistoryCompactionNotApplied<'static> { HistoryCompactionNotApplied { stage: "compact_dispatch", diff --git a/nodedb/src/diag/context/crdt_dead_letter_store.rs b/nodedb/src/diag/context/crdt_dead_letter_store.rs new file mode 100644 index 000000000..135345430 --- /dev/null +++ b/nodedb/src/diag/context/crdt_dead_letter_store.rs @@ -0,0 +1,133 @@ +// SPDX-License-Identifier: BUSL-1.1 + +//! Forensic payloads for dead-letter entries that storage did not hold: an +//! entry the store refused, and a stored entry the queue refused on restore. + +use faultbox::DomainContext; +use faultbox::serde_json::{Value, json}; + +/// A rejected CRDT delta whose dead-letter entry the store refused, so the +/// rejection has no durable record. +pub(in crate::diag) struct CrdtDeadLetterNotStored<'a> { + pub database_id: u64, + pub tenant_id: u64, + /// Collection the delta wrote. + pub collection: &'a str, + /// Constraint the delta violated. + pub constraint: &'a str, + /// Record whose apply rejected the delta. + pub source_lsn: u64, + /// What failed, without the per-occurrence detail. + pub error_class: &'a str, +} + +impl DomainContext for CrdtDeadLetterNotStored<'_> { + fn domain_kind(&self) -> &'static str { + "nodedb.crdt_dead_letter_not_stored" + } + + fn grouping_key(&self) -> String { + // The error class names the cause. Tenant, collection, constraint and + // record are the occurrence, so a failing store files one report. + format!("cause={}", self.error_class) + } + + fn to_json(&self) -> Value { + json!({ + "database_id": self.database_id, + "tenant_id": self.tenant_id, + "collection": self.collection, + "constraint": self.constraint, + "source_lsn": self.source_lsn, + "error_class": self.error_class, + "why_reported": "the delta violated a constraint and was refused, but the store \ + did not take its dead-letter entry. The entry is removed from \ + the queue and the write that needed it fails: a live apply \ + refuses with an error, a sync apply holds its high-water-mark, \ + and restart replay stops", + "operator_action": "check the sparse store (disk space, permissions, redb \ + health) on this node, then let the sender re-push or \ + restart the node", + }) + } +} + +/// A stored dead-letter entry the in-memory queue refused while the tenant's +/// CRDT engine opened. +pub(in crate::diag) struct CrdtDeadLetterNotRestored<'a> { + pub tenant_id: u64, + /// Record that produced the refused entry, when it names one. + pub source_lsn: Option, + /// What failed, without the per-occurrence detail. + pub error_class: &'a str, +} + +impl DomainContext for CrdtDeadLetterNotRestored<'_> { + fn domain_kind(&self) -> &'static str { + "nodedb.crdt_dead_letter_not_restored" + } + + fn grouping_key(&self) -> String { + // The tenant and record are the occurrence, so every refused open of + // one oversized store files one report. + format!("cause={}", self.error_class) + } + + fn to_json(&self) -> Value { + json!({ + "tenant_id": self.tenant_id, + "source_lsn": self.source_lsn, + "error_class": self.error_class, + "why_reported": "storage holds more dead-letter entries for the tenant than its \ + queue takes. Every stored entry was accepted by the queue when \ + it was written, so this is a broken invariant. The tenant's \ + CRDT engine does not open, so no CRDT write for the tenant \ + applies and restart replay stops", + "operator_action": "compare the tenant's stored dead-letter entries against the \ + queue capacity, and purge the collections whose rejected \ + deltas are resolved", + }) + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn a_store_error_groups_by_its_cause_alone() { + let first = CrdtDeadLetterNotStored { + database_id: 1, + tenant_id: 2, + collection: "users", + constraint: "users_email_unique", + source_lsn: 11, + error_class: "redb", + }; + let second = CrdtDeadLetterNotStored { + database_id: 3, + tenant_id: 4, + collection: "orders", + constraint: "orders_fk", + source_lsn: 90, + ..first + }; + assert_eq!(first.grouping_key(), second.grouping_key()); + assert_eq!(first.grouping_key(), "cause=redb"); + } + + #[test] + fn a_refused_restore_groups_by_its_cause_alone() { + let first = CrdtDeadLetterNotRestored { + tenant_id: 2, + source_lsn: Some(11), + error_class: "dead-letter queue full", + }; + let second = CrdtDeadLetterNotRestored { + tenant_id: 9, + source_lsn: None, + ..first + }; + assert_eq!(first.grouping_key(), second.grouping_key()); + } +} diff --git a/nodedb/src/diag/context/mod.rs b/nodedb/src/diag/context/mod.rs index 2f9dc75a8..1a26fa9f0 100644 --- a/nodedb/src/diag/context/mod.rs +++ b/nodedb/src/diag/context/mod.rs @@ -7,9 +7,12 @@ //! report with a rising count instead of one per attempt. mod catalog; +mod columnar; mod continuous_agg; mod crdt; +mod crdt_dead_letter_store; mod data_plane; +mod event_image; mod index_rebuild; mod ingest; mod lease; @@ -26,12 +29,17 @@ pub(in crate::diag) use catalog::{ CatalogApplyOrphanRow, CollectionPurgeRowMissing, ConsumerGroupOffsetsRetained, MetadataApplyWedged, SynonymGroupNotApplied, }; +pub(in crate::diag) use columnar::{ColumnarSegmentCorrupt, TimeseriesPartitionUnreadable}; pub(in crate::diag) use continuous_agg::ContinuousAggregateNotApplied; -pub(in crate::diag) use crdt::HistoryCompactionNotApplied; +pub(in crate::diag) use crdt::{CrdtDeadLetterNotEnqueued, HistoryCompactionNotApplied}; +pub(in crate::diag) use crdt_dead_letter_store::{ + CrdtDeadLetterNotRestored, CrdtDeadLetterNotStored, +}; pub use data_plane::LostResponseWrite; pub(in crate::diag) use data_plane::{ CalvinApplyHalted, CalvinCompletionTimeout, CoreFailStopped, DataPlaneResponseLost, }; +pub(in crate::diag) use event_image::{StrictRowImageUnrendered, TimeseriesRowImageUndecodable}; pub(in crate::diag) use index_rebuild::IndexRebuildNotInstalled; pub(in crate::diag) use ingest::IlpAcceptedLinesDropped; pub use ingest::IlpFlushOutcome; diff --git a/nodedb/src/diag/mod.rs b/nodedb/src/diag/mod.rs index 607ff0c9d..c2f3ec656 100644 --- a/nodedb/src/diag/mod.rs +++ b/nodedb/src/diag/mod.rs @@ -12,7 +12,8 @@ pub use context::{DATABASE_SCOPE, IlpFlushOutcome, LostResponseWrite, TENANT_SCO pub use recording::{ IndexRebuildTarget, VectorBuildTarget, batch_insert_without_surrogates, calvin_apply_halted, calvin_completion_timeout, catalog_apply_orphan_row, collection_purge_row_missing, - consumer_group_offsets_retained, continuous_aggregate_not_applied, + columnar_segment_corrupt, consumer_group_offsets_retained, continuous_aggregate_not_applied, + crdt_dead_letter_not_enqueued, crdt_dead_letter_not_restored, crdt_dead_letter_not_stored, data_plane_core_fail_stopped, data_plane_response_lost, data_plane_responses_lost, descriptor_lease_not_renewed, entry_kind, fts_index_update_failed, history_compaction_not_applied, ilp_invalid_utf8_drop, ilp_line_read_drop, @@ -20,8 +21,9 @@ pub use recording::{ quota_row_invalid, quota_row_undecodable, quota_row_write_failed, quota_scope_purge_incomplete, quota_scope_replay_aborted, raft_entries_reapplied, raft_entry_reapplied, replay_record_unapplied, replicated_write_parked, replicated_writes_parked, - retention_autowire_orphaned, scope_quota_not_installed, strict_row_undecodable, - synonym_group_not_applied, vector_build_failed, vector_builder_disconnected, + retention_autowire_orphaned, scope_quota_not_installed, strict_row_image_unrendered, + strict_row_undecodable, synonym_group_not_applied, timeseries_partition_unreadable, + timeseries_row_image_undecodable, vector_build_failed, vector_builder_disconnected, vector_builder_spawn_failed, vector_index_not_applied, vector_rebuild_unreadable, wal_archival_failed_truncation_held, write_acked_without_durability, write_window_held, write_window_leaked, diff --git a/nodedb/src/diag/recording/crdt.rs b/nodedb/src/diag/recording/crdt.rs index a2d7b4592..1f64eeb36 100644 --- a/nodedb/src/diag/recording/crdt.rs +++ b/nodedb/src/diag/recording/crdt.rs @@ -1,6 +1,7 @@ // SPDX-License-Identifier: BUSL-1.1 -//! Capture sites for CRDT history work that never reached the Data Plane. +//! Capture sites for CRDT history work that never reached the Data Plane, and +//! for dead-letter entries the queue or the store could not hold. use faultbox::{Capture, EventKind, error_chain_of}; @@ -35,3 +36,78 @@ pub fn history_compaction_not_applied( .with_backtrace() .emit(); } + +/// Report a constraint-rejected CRDT delta the dead-letter queue refused. +/// Called from the validated apply, which refuses the delta as an error so +/// the sender keeps it. +pub fn crdt_dead_letter_not_enqueued( + err: &nodedb_crdt::CrdtError, + tenant_id: u64, + collection: &str, + constraint: &str, +) { + let class = error_class(err); + let ctx = context::CrdtDeadLetterNotEnqueued { + tenant_id, + collection, + constraint, + error_class: &class, + }; + let _ = Capture::new( + EventKind::Error, + "crdt apply: a constraint-rejected delta could not be dead-lettered", + ) + .error_chain(error_chain_of(err)) + .domain(&ctx) + .emit(); +} + +/// Report a dead-letter entry the store refused. Called only from +/// `store_crdt_dead_letter`, the one site every apply path stores through. +pub fn crdt_dead_letter_not_stored( + err: &crate::Error, + database_id: u64, + tenant_id: u64, + entry: &nodedb_crdt::DeadLetter, + source_lsn: u64, +) { + let class = error_class(err); + let ctx = context::CrdtDeadLetterNotStored { + database_id, + tenant_id, + collection: &entry.collection, + constraint: &entry.violated_constraint, + source_lsn, + error_class: &class, + }; + let _ = Capture::new( + EventKind::Error, + "crdt apply: a rejected delta's dead-letter entry could not be stored", + ) + .error_chain(error_chain_of(err)) + .domain(&ctx) + .emit(); +} + +/// Report a stored dead-letter entry the queue refused on restore. Called +/// only from `restore_dead_letters`. +pub fn crdt_dead_letter_not_restored( + err: &nodedb_crdt::CrdtError, + tenant_id: u64, + source_lsn: Option, +) { + let class = error_class(err); + let ctx = context::CrdtDeadLetterNotRestored { + tenant_id, + source_lsn, + error_class: &class, + }; + let _ = Capture::new( + EventKind::InvariantViolation, + "crdt engine open: a stored dead-letter entry does not fit the queue", + ) + .error_chain(error_chain_of(err)) + .domain(&ctx) + .with_backtrace() + .emit(); +} diff --git a/nodedb/src/diag/recording/mod.rs b/nodedb/src/diag/recording/mod.rs index b147a99ed..c34f9d202 100644 --- a/nodedb/src/diag/recording/mod.rs +++ b/nodedb/src/diag/recording/mod.rs @@ -9,9 +9,11 @@ //! discarded. mod catalog; +mod columnar; mod continuous_agg; mod crdt; mod data_plane; +mod event_image; mod index_rebuild; mod ingest; mod lease; @@ -28,12 +30,17 @@ pub use catalog::{ catalog_apply_orphan_row, collection_purge_row_missing, consumer_group_offsets_retained, metadata_apply_wedged, synonym_group_not_applied, }; +pub use columnar::{columnar_segment_corrupt, timeseries_partition_unreadable}; pub use continuous_agg::continuous_aggregate_not_applied; -pub use crdt::history_compaction_not_applied; +pub use crdt::{ + crdt_dead_letter_not_enqueued, crdt_dead_letter_not_restored, crdt_dead_letter_not_stored, + history_compaction_not_applied, +}; pub use data_plane::{ calvin_apply_halted, calvin_completion_timeout, data_plane_core_fail_stopped, data_plane_response_lost, data_plane_responses_lost, }; +pub use event_image::{strict_row_image_unrendered, timeseries_row_image_undecodable}; pub use index_rebuild::{IndexRebuildTarget, index_rebuild_not_installed}; pub use ingest::{ilp_invalid_utf8_drop, ilp_line_read_drop}; pub use lease::descriptor_lease_not_renewed; diff --git a/nodedb/src/engine/crdt/tenant_state/apply_target.rs b/nodedb/src/engine/crdt/tenant_state/apply_target.rs new file mode 100644 index 000000000..d00644d9b --- /dev/null +++ b/nodedb/src/engine/crdt/tenant_state/apply_target.rs @@ -0,0 +1,73 @@ +// SPDX-License-Identifier: BUSL-1.1 + +//! What a validated delta apply writes into, and how it imports the bytes. + +use nodedb_crdt::state::{CrdtState, ImportAdmission, WriteSetImport}; +use nodedb_types::Surrogate; + +/// What a validated delta apply writes into. +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum ApplyTarget<'a> { + /// One document's delta, under the row's bound surrogate. The caller + /// refuses `Surrogate::ZERO` before it builds this target. A delta that + /// writes any other row is refused as malformed. + /// + /// A document delta is a sync delta from a peer. It imports under the + /// peer byte and operation ceilings. + Document { + document_id: &'a str, + surrogate: Surrogate, + }, + /// A per-collection snapshot import. It can write any row of the + /// collection and binds no row identity. + /// + /// The snapshot is one this node or its backup exported: a checkpoint, + /// a restored collection, or the WAL record of either. It imports + /// without the peer ceilings and keeps every structural check, so a + /// collection that grew past the ceilings still loads. + Collection, +} + +impl ApplyTarget<'_> { + /// Whether a delta applied to this target can write `row`. + pub(super) fn admits_row(&self, row: &str) -> bool { + match self { + Self::Document { document_id, .. } => row == *document_id, + Self::Collection => true, + } + } + + /// The identity `row` validates under: the document's surrogate for a + /// document target. A snapshot import names no row identity, and the + /// validator's change record carries `Surrogate::ZERO` for it. The + /// validator never reads that field. + pub(super) fn validation_surrogate(&self, row: &str) -> Surrogate { + match self { + Self::Document { + document_id, + surrogate, + } if row == *document_id => *surrogate, + Self::Document { .. } | Self::Collection => Surrogate::ZERO, + } + } + + /// Import `delta` into `state` and report the rows it wrote. + pub(super) fn import_with_write_set(&self, state: &CrdtState, delta: &[u8]) -> WriteSetImport { + match self { + Self::Document { .. } => state.import_with_write_set(delta), + Self::Collection => state.import_local_with_write_set(delta), + } + } + + /// Import `delta` into `state`. + pub(super) fn import( + &self, + state: &CrdtState, + delta: &[u8], + ) -> nodedb_crdt::Result { + match self { + Self::Document { .. } => state.import(delta), + Self::Collection => state.import_local(delta), + } + } +} diff --git a/nodedb/src/engine/crdt/tenant_state/apply_validated.rs b/nodedb/src/engine/crdt/tenant_state/apply_validated.rs index 76378ce7b..ae9728867 100644 --- a/nodedb/src/engine/crdt/tenant_state/apply_validated.rs +++ b/nodedb/src/engine/crdt/tenant_state/apply_validated.rs @@ -10,9 +10,9 @@ use nodedb_crdt::state::CrdtState; use nodedb_crdt::validator::{ValidationOutcome, Violation}; -use nodedb_types::Surrogate; use nodedb_types::sync::violation::ViolationType; +use super::apply_target::ApplyTarget; use super::core::TenantCrdtEngine; /// Server-derived signing context for an externally synchronized delta. @@ -25,45 +25,6 @@ pub struct DeltaSigningAdmission { pub preverified: bool, } -/// What a validated delta apply writes into. -#[derive(Debug, Clone, Copy, PartialEq, Eq)] -pub enum ApplyTarget<'a> { - /// One document's delta, under the row's bound surrogate. The caller - /// refuses `Surrogate::ZERO` before it builds this target. A delta that - /// writes any other row is refused as malformed. - Document { - document_id: &'a str, - surrogate: Surrogate, - }, - /// A per-collection snapshot import. It can write any row of the - /// collection and binds no row identity. - Collection, -} - -impl ApplyTarget<'_> { - /// Whether a delta applied to this target can write `row`. - fn admits_row(&self, row: &str) -> bool { - match self { - Self::Document { document_id, .. } => row == *document_id, - Self::Collection => true, - } - } - - /// The identity `row` validates under: the document's surrogate for a - /// document target. A snapshot import names no row identity, and the - /// validator's change record carries `Surrogate::ZERO` for it. The - /// validator never reads that field. - fn validation_surrogate(&self, row: &str) -> Surrogate { - match self { - Self::Document { - document_id, - surrogate, - } if row == *document_id => *surrogate, - Self::Document { .. } | Self::Collection => Surrogate::ZERO, - } - } -} - /// Outcome of applying and validating one peer delta. #[derive(Debug)] pub enum ValidatedApplyOutcome { @@ -89,8 +50,15 @@ pub enum ValidatedApplyOutcome { imported_ops: usize, }, /// The candidate violated a constraint and was discarded. The violation - /// has been enqueued to the DLQ and translated for the caller. + /// is in the DLQ and translated for the caller. Rejected(ViolationType), + /// The candidate violated a constraint and was discarded, and the DLQ + /// refused its entry. Nothing applied and nothing on this node records + /// the delta, so the caller refuses it as an error and the sender keeps it. + DeadLetterRefused { + violation: ViolationType, + error: nodedb_crdt::CrdtError, + }, /// The delta bytes could not be imported (corrupt / undecodable). Treated /// as an idempotent no-op so the stream is not wedged. Malformed, @@ -104,6 +72,14 @@ pub enum ValidatedApplyOutcome { /// the high-water-mark past a write that was never applied — the silent /// data-loss path. The caller must refuse the delta instead. PendingDependencies, + /// This node could not build the detached candidate the delta validates + /// in: exporting or seeding the collection's document failed. Nothing + /// applied, and the failure says nothing about the delta, so the caller + /// refuses it retryably and never as malformed. + CandidateUnavailable { + /// Why the export or seed failed. + error: nodedb_crdt::CrdtError, + }, } impl TenantCrdtEngine { @@ -160,25 +136,28 @@ impl TenantCrdtEngine { } let candidate = match self.take_apply_candidate(collection) { Ok(candidate) => candidate, - Err(()) => return ValidatedApplyOutcome::Malformed, + Err(error) => return ValidatedApplyOutcome::CandidateUnavailable { error }, }; - let before = candidate.frontier(); - let admission = match candidate.import(delta) { + let imported = target.import_with_write_set(&candidate, delta); + let admission = match imported.outcome { Ok(admission) => admission, // Well-formed operations that arrived without their causal history // (their predecessors live in another collection's document, absent - // from this candidate). The candidate did not move, so this is NOT a - // no-op the caller may acknowledge — refusing it preserves the row - // instead of advancing the high-water-mark past a write that never - // applied. + // from this candidate). Some of the delta stayed pending, so this is + // NOT a no-op the caller may acknowledge — refusing it preserves the + // row instead of advancing the high-water-mark past a write that + // never applied. The candidate is dropped with whatever part did + // apply, so authoritative state stays unchanged. Err(nodedb_crdt::CrdtError::ImportPendingDependencies) => { return ValidatedApplyOutcome::PendingDependencies; } Err(_) => return ValidatedApplyOutcome::Malformed, }; - let write_set = match candidate.write_set_since(&before) { + // A write outside the root-map-of-row-maps shape names no row the + // constraints can check. The candidate is dropped unapplied. + let write_set = match imported.write_set { Ok(write_set) => write_set, - Err(_) => return ValidatedApplyOutcome::Malformed, + Err(fault) => return self.reject_shape_fault(collection, delta, peer_id, &fault), }; if write_set .iter() @@ -213,7 +192,8 @@ impl TenantCrdtEngine { }, }, }; - let violation = self.dlq_and_translate(coll, delta, peer_id, violation); + let wire = violation_to_type(&violation); + let outcome = self.dead_letter(coll, delta, peer_id, &violation, wire); match previous { Some(previous) => { self.collections.insert(collection.to_owned(), previous); @@ -222,7 +202,7 @@ impl TenantCrdtEngine { self.collections.remove(collection); } } - return ValidatedApplyOutcome::Rejected(violation); + return outcome; } // The candidate is now authoritative. Bring the doc it displaced up to @@ -235,7 +215,7 @@ impl TenantCrdtEngine { // land, no candidate is retained and the next apply rebuilds one — // slower, never wrong. if let Some(previous) = previous - && previous.import(delta).is_ok() + && target.import(&previous, delta).is_ok() { self.apply_candidates .insert(collection.to_owned(), previous); @@ -260,7 +240,10 @@ impl TenantCrdtEngine { /// delta that was refused for any reason leaves its operations in the /// candidate — Loro buffers even causally-pending ones — so a poisoned /// candidate is dropped rather than reused. - fn take_apply_candidate(&mut self, collection: &str) -> Result { + fn take_apply_candidate( + &mut self, + collection: &str, + ) -> Result { if let Some(candidate) = self.apply_candidates.remove(collection) { // A candidate is only usable while it still matches the state it // was cloned from. Anything else that touches this collection — a @@ -287,10 +270,10 @@ impl TenantCrdtEngine { // would fail to seed and every subsequent delta would report // `Malformed` — silently unwritable, with the sender blamed for it. Some(current) => { - let snapshot = current.export_snapshot().map_err(|_| ())?; - CrdtState::from_local_snapshot(peer_id, &snapshot).map_err(|_| ()) + let snapshot = current.export_snapshot()?; + CrdtState::from_local_snapshot(peer_id, &snapshot) } - None => CrdtState::new(peer_id).map_err(|_| ()), + None => CrdtState::new(peer_id), } } @@ -315,62 +298,61 @@ impl TenantCrdtEngine { self.collections.get(collection).map(|s| s.peer_id()) } - /// Enqueue a rejected delta to the DLQ and translate the internal violation - /// into the deterministic wire [`ViolationType`]. + /// Dead-letter a delta whose writes break the collection shape. The + /// caller gets the fault as a schema violation. + fn reject_shape_fault( + &mut self, + collection: &str, + delta: &[u8], + peer_id: u64, + fault: &nodedb_crdt::CrdtError, + ) -> ValidatedApplyOutcome { + let reason = fault.to_string(); + let violation = Violation { + constraint_name: ROW_SHAPE_CONSTRAINT.to_string(), + reason: reason.clone(), + hint: nodedb_crdt::CompensationHint::ManualIntervention { + reason: reason.clone(), + }, + }; + let wire = ViolationType::SchemaViolation { + field: shape_fault_target(fault), + reason, + }; + self.dead_letter(collection, delta, peer_id, &violation, wire) + } + + /// Enqueue a rejected delta to the DLQ and report it as `wire`. /// - /// The DLQ entry carries the INTERNAL compensation hint verbatim; the wire - /// hint the client eventually sees is derived from the returned - /// `ViolationType`, never from the DLQ. The DLQ id / timestamp are - /// node-local and non-deterministic and are deliberately not returned. - fn dlq_and_translate( + /// The DLQ entry carries the INTERNAL compensation hint verbatim. The wire + /// hint the client sees derives from `wire`, never from the DLQ. The DLQ + /// id and timestamp are node-local and non-deterministic, so they are not + /// returned. A delta the DLQ cannot hold is recorded nowhere, so the apply + /// reports [`ValidatedApplyOutcome::DeadLetterRefused`]. + fn dead_letter( &mut self, collection: &str, delta: &[u8], peer_id: u64, - violation: Violation, - ) -> ViolationType { + violation: &Violation, + wire: ViolationType, + ) -> ValidatedApplyOutcome { // No authenticated user identity is threaded to the apply path in this - // layer, so the DLQ records `0` (unauthenticated/legacy). A real - // user_id would come from the sync session's auth context once that is - // carried alongside `SyncProvenance` into the delta apply. + // layer, so the DLQ records `0` (unauthenticated/legacy). let user_id = 0u64; let tenant_id = self.tenant_id().as_u64(); - // Look up the violated constraint by name so the DLQ entry records the - // real collection/field. If it cannot be found, fall back to a - // best-effort ManualIntervention entry rather than panicking. - let constraint = self + // The violated constraint by name, so the DLQ entry records the real + // collection and field. An unresolved name records a + // ManualIntervention entry under a placeholder CHECK constraint. + let resolved = self .constraints_for_collection(collection) .into_iter() .find(|c| c.name == violation.constraint_name); - - let reason = violation.reason.clone(); - match constraint { - Some(constraint) => { - let enqueued = - self.validator - .dlq_mut() - .enqueue(nodedb_crdt::EnqueueDeadLetterArgs { - peer_id, - user_id, - tenant_id, - delta: delta.to_vec(), - constraint: &constraint, - reason, - hint: violation.hint.clone(), - }); - match enqueued { - Ok(id) => self.last_dead_letter = Some(id), - Err(e) => tracing::warn!( - tenant = tenant_id, - collection, - error = %e, - "crdt: failed to enqueue rejected delta to DLQ" - ), - } - } - None => { - let fallback = nodedb_crdt::Constraint { + let (constraint, hint) = match resolved { + Some(constraint) => (constraint, violation.hint.clone()), + None => ( + nodedb_crdt::Constraint { name: violation.constraint_name.clone(), collection: collection.to_string(), field: String::new(), @@ -378,35 +360,56 @@ impl TenantCrdtEngine { expr: String::new(), description: "unresolved constraint".to_string(), }, - }; - let hint = nodedb_crdt::CompensationHint::ManualIntervention { - reason: reason.clone(), - }; - let enqueued = - self.validator - .dlq_mut() - .enqueue(nodedb_crdt::EnqueueDeadLetterArgs { - peer_id, - user_id, - tenant_id, - delta: delta.to_vec(), - constraint: &fallback, - reason, - hint, - }); - match enqueued { - Ok(id) => self.last_dead_letter = Some(id), - Err(e) => tracing::warn!( - tenant = tenant_id, - collection, - error = %e, - "crdt: failed to enqueue rejected delta to DLQ (unresolved constraint)" - ), + }, + nodedb_crdt::CompensationHint::ManualIntervention { + reason: violation.reason.clone(), + }, + ), + }; + let enqueued = self + .validator + .dlq_mut() + .enqueue(nodedb_crdt::EnqueueDeadLetterArgs { + peer_id, + user_id, + tenant_id, + delta: delta.to_vec(), + constraint: &constraint, + reason: violation.reason.clone(), + hint, + }); + match enqueued { + Ok(id) => { + self.last_dead_letter = Some(id); + ValidatedApplyOutcome::Rejected(wire) + } + Err(error) => { + crate::diag::crdt_dead_letter_not_enqueued( + &error, + tenant_id, + collection, + &constraint.name, + ); + ValidatedApplyOutcome::DeadLetterRefused { + violation: wire, + error, } } } + } +} - violation_to_type(&violation) +/// Name the dead-letter entry of a delta that breaks the collection shape. +const ROW_SHAPE_CONSTRAINT: &str = "crdt_row_shape"; + +/// The container or row a shape fault names, as `collection/row` for a row. +fn shape_fault_target(fault: &nodedb_crdt::CrdtError) -> String { + match fault { + nodedb_crdt::CrdtError::NonMapRootContainer { container, .. } => container.clone(), + nodedb_crdt::CrdtError::NonMapRowValue { + collection, row_id, .. + } => format!("{collection}/{row_id}"), + other => other.to_string(), } } @@ -478,6 +481,49 @@ mod tests { state.export_snapshot().unwrap() } + /// A raw Loro snapshot built by `write` on a fresh document. + fn raw_delta(peer: u64, write: impl FnOnce(&loro::LoroDoc)) -> Vec { + let doc = loro::LoroDoc::new(); + doc.set_peer_id(peer).unwrap(); + write(&doc); + doc.export(loro::ExportMode::Snapshot).unwrap() + } + + #[test] + fn a_scalar_row_is_rejected_as_a_schema_violation() { + let mut engine = unique_engine(); + let delta = raw_delta(5, |doc| { + doc.get_map("users").insert("s1", 5).unwrap(); + }); + let outcome = + engine.apply_committed_delta_validated("users", &delta, ApplyTarget::Collection, 5); + match outcome { + ValidatedApplyOutcome::Rejected(ViolationType::SchemaViolation { field, reason }) => { + assert_eq!(field, "users/s1"); + assert!(reason.contains("s1"), "{reason}"); + } + other => panic!("expected SchemaViolation, got {other:?}"), + } + assert_eq!(engine.dlq_len(), 1); + assert!(!engine.row_exists("users", "s1")); + } + + #[test] + fn a_root_text_delta_is_rejected_as_a_schema_violation() { + let mut engine = unique_engine(); + let delta = raw_delta(6, |doc| { + doc.get_text("notes").insert(0, "hi").unwrap(); + }); + let outcome = + engine.apply_committed_delta_validated("users", &delta, ApplyTarget::Collection, 6); + match outcome { + ValidatedApplyOutcome::Rejected(ViolationType::SchemaViolation { field, .. }) => { + assert_eq!(field, "notes"); + } + other => panic!("expected SchemaViolation, got {other:?}"), + } + } + #[test] fn valid_delta_is_clean() { let mut engine = unique_engine(); @@ -854,6 +900,54 @@ mod tests { ); } + /// A rejected delta the full DLQ refuses surfaces as an error carrying + /// the DLQ failure. Nothing applies and nothing is left to bind. + #[test] + fn a_rejection_the_full_dlq_refuses_surfaces_as_an_error() { + let mut engine = unique_engine(); + let seed = row_delta(2, "a", "x@y.com", "A"); + let seeded = engine.apply_committed_delta_validated( + "users", + &seed, + ApplyTarget::Document { + document_id: "a", + surrogate: nodedb_types::Surrogate::new(10), + }, + 2, + ); + assert!(matches!(seeded, ValidatedApplyOutcome::Clean { .. })); + let capacity = engine.fill_dead_letter_queue_for_test(); + assert!(capacity > 0); + + let duplicate = row_delta(3, "b", "x@y.com", "B"); + let outcome = engine.apply_committed_delta_validated( + "users", + &duplicate, + ApplyTarget::Document { + document_id: "b", + surrogate: nodedb_types::Surrogate::new(11), + }, + 3, + ); + match outcome { + ValidatedApplyOutcome::DeadLetterRefused { + violation: ViolationType::UniqueViolation { field, .. }, + error: nodedb_crdt::CrdtError::DlqFull { .. }, + } => assert_eq!(field, "email"), + other => panic!("expected DeadLetterRefused(DlqFull), got {other:?}"), + } + assert_eq!( + engine.dlq_len(), + capacity, + "the refused entry is not queued" + ); + assert!(engine.read_row("users", "b").is_none()); + assert!( + engine.bind_dead_letter_source(12).is_none(), + "a refused apply leaves no entry to bind" + ); + } + /// The load-bearing safety property of retaining a candidate across /// applies: the candidate becomes authoritative on a clean apply, so a /// candidate that missed a local write would silently erase it. diff --git a/nodedb/src/engine/crdt/tenant_state/dead_letters.rs b/nodedb/src/engine/crdt/tenant_state/dead_letters.rs index 84a49a5e6..1fc7563e7 100644 --- a/nodedb/src/engine/crdt/tenant_state/dead_letters.rs +++ b/nodedb/src/engine/crdt/tenant_state/dead_letters.rs @@ -25,24 +25,138 @@ impl TenantCrdtEngine { .cloned() } - /// Put back entries read from durable storage. An entry the queue cannot - /// hold is logged and skipped: storage keeps it. - pub fn restore_dead_letters(&mut self, entries: Vec) { + /// Put back entries read from durable storage. + /// + /// An entry the queue cannot hold is an error: every stored entry was in + /// the queue when it was stored. The caller does not install an engine + /// that is missing an entry storage holds. + pub fn restore_dead_letters(&mut self, entries: Vec) -> crate::Result<()> { for entry in entries { let source_lsn = entry.source_lsn; if let Err(error) = self.validator.dlq_mut().restore(entry) { - tracing::warn!( - tenant = self.tenant_id.as_u64(), - ?source_lsn, - %error, - "crdt: a stored dead-letter entry does not fit the queue" + crate::diag::crdt_dead_letter_not_restored( + &error, + self.tenant_id.as_u64(), + source_lsn, ); + return Err(crate::Error::Crdt(error)); } } + Ok(()) + } + + /// Remove entry `id` from the queue. The caller removes an entry whose + /// store failed, so the queue holds no entry that storage lacks. + pub fn discard_dead_letter(&mut self, id: u64) { + self.validator.dlq_mut().remove(id); + } + + /// Whether the queue holds an entry produced by the record at + /// `source_lsn`. + pub fn dead_letter_recorded(&self, source_lsn: u64) -> bool { + self.validator + .dlq() + .iter() + .any(|entry| entry.source_lsn == Some(source_lsn)) } /// Every pending dead-letter entry, oldest first. pub fn dead_letters(&self) -> impl Iterator { self.validator.dlq().iter() } + + /// Enqueue placeholder entries until the dead-letter queue refuses one. + /// Returns how many it took. + #[cfg(test)] + pub(crate) fn fill_dead_letter_queue_for_test(&mut self) -> usize { + let filler = nodedb_crdt::Constraint { + name: "filler".into(), + collection: "filler".into(), + field: String::new(), + kind: nodedb_crdt::ConstraintKind::Check { + expr: String::new(), + description: "filler".into(), + }, + }; + let mut taken = 0; + while self + .validator + .dlq_mut() + .enqueue(nodedb_crdt::EnqueueDeadLetterArgs { + peer_id: 0, + user_id: 0, + tenant_id: self.tenant_id.as_u64(), + delta: Vec::new(), + constraint: &filler, + reason: String::new(), + hint: nodedb_crdt::CompensationHint::ManualIntervention { + reason: String::new(), + }, + }) + .is_ok() + { + taken += 1; + } + taken + } +} + +#[cfg(test)] +mod tests { + use nodedb_crdt::constraint::ConstraintSet; + use nodedb_crdt::{ + CompensationHint, Constraint, ConstraintKind, DeadLetterQueue, EnqueueDeadLetterArgs, + }; + + use super::*; + use crate::types::TenantId; + + fn stored_entry(source_lsn: u64) -> DeadLetter { + let mut dlq = DeadLetterQueue::new(1); + let id = dlq + .enqueue(EnqueueDeadLetterArgs { + peer_id: 1, + user_id: 0, + tenant_id: 3, + delta: b"delta".to_vec(), + constraint: &Constraint { + name: "email_unique".into(), + collection: "users".into(), + field: "email".into(), + kind: ConstraintKind::Unique, + }, + reason: "duplicate".into(), + hint: CompensationHint::ManualIntervention { + reason: "duplicate".into(), + }, + }) + .expect("enqueue"); + dlq.bind_source(id, source_lsn).cloned().expect("bound") + } + + /// A stored entry the full queue refuses fails the restore: the engine + /// is not installed missing an entry storage holds. + #[test] + fn a_stored_entry_the_full_queue_refuses_fails_the_restore() { + let mut engine = + TenantCrdtEngine::new(TenantId::new(3), 0, ConstraintSet::new()).expect("engine"); + engine.fill_dead_letter_queue_for_test(); + match engine.restore_dead_letters(vec![stored_entry(5)]) { + Err(crate::Error::Crdt(nodedb_crdt::CrdtError::DlqFull { .. })) => {} + other => panic!("expected DlqFull, got {other:?}"), + } + assert!(!engine.dead_letter_recorded(5)); + } + + #[test] + fn a_restored_entry_is_recorded_and_a_discarded_one_is_not() { + let mut engine = + TenantCrdtEngine::new(TenantId::new(3), 0, ConstraintSet::new()).expect("engine"); + let entry = stored_entry(5); + let id = entry.id; + engine.restore_dead_letters(vec![entry]).expect("restore"); + assert!(engine.dead_letter_recorded(5)); + engine.discard_dead_letter(id); + assert!(!engine.dead_letter_recorded(5)); + } } diff --git a/nodedb/src/engine/crdt/tenant_state/doc_mutate.rs b/nodedb/src/engine/crdt/tenant_state/doc_mutate.rs index 1c3775c1c..99a40c71f 100644 --- a/nodedb/src/engine/crdt/tenant_state/doc_mutate.rs +++ b/nodedb/src/engine/crdt/tenant_state/doc_mutate.rs @@ -3,6 +3,7 @@ //! Server-built document-row mutations for `CrdtOp::DocUpsert` / `DocDelete`. use loro::LoroValue; +use nodedb_crdt::state::RowImage; use super::TenantCrdtEngine; @@ -39,4 +40,34 @@ impl TenantCrdtEngine { .delete(collection, row_id) .map_err(crate::Error::Crdt) } + + /// Capture the state of a document row that `doc_upsert` and + /// `doc_set_fields` can change. `None`: the collection has no local + /// state. + pub fn doc_row_image(&self, collection: &str, row_id: &str) -> crate::Result> { + match self.collections.get(collection) { + Some(state) => state + .row_image(collection, row_id) + .map(Some) + .map_err(crate::Error::Crdt), + None => Ok(None), + } + } + + /// Put a document row back to the image `doc_row_image` captured. `None` + /// removes the collection's local state: it had none before the write. + pub fn restore_doc_row( + &mut self, + collection: &str, + row_id: &str, + image: Option<&RowImage>, + ) -> crate::Result<()> { + let Some(image) = image else { + self.collections.remove(collection); + return Ok(()); + }; + self.state_mut(collection)? + .restore_row_image(collection, row_id, image) + .map_err(crate::Error::Crdt) + } } diff --git a/nodedb/src/engine/crdt/tenant_state/mod.rs b/nodedb/src/engine/crdt/tenant_state/mod.rs index 8888141a9..cb89d101c 100644 --- a/nodedb/src/engine/crdt/tenant_state/mod.rs +++ b/nodedb/src/engine/crdt/tenant_state/mod.rs @@ -6,6 +6,7 @@ //! queue for a single tenant. Lives on the Data Plane (one per tenant per core). pub mod apply; +pub mod apply_target; pub mod apply_validated; pub mod constraints; pub mod core; @@ -20,5 +21,6 @@ mod snapshot_io; mod snapshot_restore; pub mod validate; -pub use apply_validated::{ApplyTarget, DeltaSigningAdmission, ValidatedApplyOutcome}; +pub use apply_target::ApplyTarget; +pub use apply_validated::{DeltaSigningAdmission, ValidatedApplyOutcome}; pub use core::TenantCrdtEngine; diff --git a/nodedb/src/engine/crdt/tenant_state/snapshot_io.rs b/nodedb/src/engine/crdt/tenant_state/snapshot_io.rs index e56b624f9..6fae172c1 100644 --- a/nodedb/src/engine/crdt/tenant_state/snapshot_io.rs +++ b/nodedb/src/engine/crdt/tenant_state/snapshot_io.rs @@ -47,6 +47,10 @@ impl TenantCrdtEngine { /// fully applied — a restore that left operations causally pending has NOT /// restored the collection, and reporting success would leave the caller /// unable to tell a complete restore from a partial one. + /// + /// The bytes are a snapshot this node or its backup exported, so they + /// import without the peer byte and operation ceilings + /// ([`super::ApplyTarget::Collection`]). pub fn import_snapshot_bytes(&mut self, collection: &str, bytes: &[u8]) -> crate::Result<()> { match self.apply_committed_delta_validated( collection, @@ -62,6 +66,11 @@ impl TenantCrdtEngine { detail: reason.to_string(), }, )), + // The DLQ refusal is the error: nothing applied and nothing + // records the snapshot. + super::ValidatedApplyOutcome::DeadLetterRefused { error, .. } => { + Err(crate::Error::Crdt(error)) + } super::ValidatedApplyOutcome::Malformed => Err(crate::Error::Crdt( nodedb_crdt::CrdtError::DeltaApplyFailed("malformed snapshot".into()), )), @@ -70,6 +79,12 @@ impl TenantCrdtEngine { "snapshot import left operations causally pending".into(), ), )), + super::ValidatedApplyOutcome::CandidateUnavailable { error } => Err( + crate::Error::Crdt(nodedb_crdt::CrdtError::DeltaApplyFailed(format!( + "no apply candidate for collection {collection}: {error}; nothing was \ + imported" + ))), + ), } } } @@ -131,6 +146,31 @@ mod tests { ); } + /// A checkpoint of a collection past the peer operation ceiling loads: + /// a snapshot import runs without the peer ceilings. + #[test] + fn a_snapshot_past_the_peer_ceilings_imports() { + let doc = loro::LoroDoc::new(); + doc.set_peer_id(7).unwrap(); + let row = doc + .get_map("docs") + .insert_container("r", loro::LoroMap::new()) + .unwrap(); + row.insert_container("body", loro::LoroText::new()) + .unwrap() + .insert( + 0, + &"x".repeat(nodedb_crdt::state::DEFAULT_MAX_IMPORT_OPS + 1), + ) + .unwrap(); + doc.commit(); + let snapshot = doc.export(loro::ExportMode::Snapshot).unwrap(); + + let mut engine = TenantCrdtEngine::new(TenantId::new(1), 0, ConstraintSet::new()).unwrap(); + engine.import_snapshot_bytes("docs", &snapshot).unwrap(); + assert!(engine.row_exists("docs", "r")); + } + /// Every collection's document is constructed with the same peer id, so Loro /// operation ids are unique only *within* a collection. Two collections' /// snapshots therefore carry colliding `(peer, counter)` identities, and a diff --git a/nodedb/src/engine/sparse/btree/crdt_dead_letter.rs b/nodedb/src/engine/sparse/btree/crdt_dead_letter.rs index 2e852e312..a598e028d 100644 --- a/nodedb/src/engine/sparse/btree/crdt_dead_letter.rs +++ b/nodedb/src/engine/sparse/btree/crdt_dead_letter.rs @@ -166,6 +166,37 @@ impl SparseEngine { .map_err(|e| redb_err("commit crdt dead letter purge", e))?; Ok(keys.len()) } + + /// Replace the dead-letter table with a same-named table of another + /// value type, so every later read or write of it fails. + #[cfg(test)] + pub(crate) fn break_crdt_dead_letter_table_for_test(&self) { + let txn = self.db.begin_write().expect("write txn"); + txn.delete_table(CRDT_DEAD_LETTERS).expect("delete table"); + { + let wrong: TableDefinition<&str, u64> = TableDefinition::new("crdt_dead_letters"); + txn.open_table(wrong).expect("create mismatched table"); + } + txn.commit().expect("commit"); + } + + /// Store `bytes` as the entry of the record at `source_lsn`, unchecked. + #[cfg(test)] + pub(crate) fn put_raw_crdt_dead_letter_for_test( + &self, + database_id: u64, + tenant_id: u64, + source_lsn: u64, + bytes: &[u8], + ) { + let key = entry_key(database_id, tenant_id, source_lsn); + let txn = self.db.begin_write().expect("write txn"); + { + let mut table = txn.open_table(CRDT_DEAD_LETTERS).expect("open table"); + table.insert(key.as_str(), bytes).expect("insert"); + } + txn.commit().expect("commit"); + } } #[cfg(test)] @@ -227,6 +258,25 @@ mod tests { ); } + #[test] + fn a_broken_table_fails_the_store_and_the_load() { + let (_dir, engine) = open(); + engine.break_crdt_dead_letter_table_for_test(); + assert!( + engine + .put_crdt_dead_letter(1, 3, 5, &entry("users", 5)) + .is_err() + ); + assert!(engine.load_crdt_dead_letters(1, 3).is_err()); + } + + #[test] + fn an_undecodable_entry_fails_the_load() { + let (_dir, engine) = open(); + engine.put_raw_crdt_dead_letter_for_test(1, 3, 5, b"\xc1"); + assert!(engine.load_crdt_dead_letters(1, 3).is_err()); + } + #[test] fn a_purged_collection_or_tenant_leaves_no_entry() { let (_dir, engine) = open(); From eda4a2b66a89899a0e8eb74f289cd563701ee919 Mon Sep 17 00:00:00 2001 From: Farhan Syah Date: Thu, 8 Oct 2026 16:51:47 +0800 Subject: [PATCH 12/24] fix(event): dead-letter events whose row image does not render A stored strict row that does not decode no longer reaches the Event Plane with a missing or raw image. The write event carries an image fault, and delivery dead-letters it with an audit entry instead of running triggers, change streams, views or CRDT sync. Redo journalling of such a row fail-stops the core. --- .../misc_suite/cases/cluster_triggers.rs | 1 + nodedb/src/control/event_trigger.rs | 1 + .../security/permission_tree/event_handler.rs | 1 + .../src/data/executor/core_loop/deferred.rs | 3 + .../src/data/executor/core_loop/event_emit.rs | 245 ++++---------- .../executor/core_loop/event_emit_engines.rs | 2 + .../data/executor/core_loop/event_image.rs | 306 ++++++++++++++++++ .../src/data/executor/core_loop/redo_image.rs | 41 +-- .../executor/handlers/timeseries/events.rs | 1 + .../handlers/transaction/redo_apply/events.rs | 9 +- .../src/data/executor/strict_format/render.rs | 151 +++++++++ nodedb/src/diag/context/event_image.rs | 85 +++++ nodedb/src/diag/recording/event_image.rs | 44 +++ nodedb/src/event/audit_dml/consumer.rs | 1 + nodedb/src/event/bus.rs | 1 + nodedb/src/event/cdc/router.rs | 1 + nodedb/src/event/consumer/delivery.rs | 1 + nodedb/src/event/consumer/drain.rs | 1 + nodedb/src/event/consumer/pipeline.rs | 82 +++++ nodedb/src/event/crdt_sync/packager.rs | 1 + nodedb/src/event/image_fault.rs | 57 ++++ nodedb/src/event/mod.rs | 1 + nodedb/src/event/plane.rs | 1 + nodedb/src/event/record_numbering.rs | 1 + nodedb/src/event/topic/committed/event.rs | 1 + nodedb/src/event/trigger/lane/held.rs | 4 + nodedb/src/event/types.rs | 7 + nodedb/src/event/wal_replay_kv_shapes.rs | 3 + nodedb/src/event/wal_replay_parse.rs | 10 + nodedb/src/event/wal_replay_timeseries.rs | 1 + nodedb/tests/inproc/cases/bitemporal_cdc.rs | 1 + nodedb/tests/inproc/cases/cdc_arc_fanout.rs | 1 + nodedb/tests/inproc/cases/event_trigger.rs | 1 + .../inproc/cases/shutdown_event_plane.rs | 1 + 34 files changed, 873 insertions(+), 195 deletions(-) create mode 100644 nodedb/src/data/executor/core_loop/event_image.rs create mode 100644 nodedb/src/data/executor/strict_format/render.rs create mode 100644 nodedb/src/diag/context/event_image.rs create mode 100644 nodedb/src/diag/recording/event_image.rs create mode 100644 nodedb/src/event/image_fault.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/src/control/event_trigger.rs b/nodedb/src/control/event_trigger.rs index 0105857e1..ae3dc9371 100644 --- a/nodedb/src/control/event_trigger.rs +++ b/nodedb/src/control/event_trigger.rs @@ -481,6 +481,7 @@ mod tests { user_id: None, statement_digest: None, commit_hlc: Some(crate::event::test_utils::test_commit_hlc()), + image_fault: None, } } diff --git a/nodedb/src/control/security/permission_tree/event_handler.rs b/nodedb/src/control/security/permission_tree/event_handler.rs index 507027420..b6a798c81 100644 --- a/nodedb/src/control/security/permission_tree/event_handler.rs +++ b/nodedb/src/control/security/permission_tree/event_handler.rs @@ -221,6 +221,7 @@ mod tests { user_id: None, statement_digest: None, commit_hlc: Some(crate::event::test_utils::test_commit_hlc()), + image_fault: None, } } diff --git a/nodedb/src/data/executor/core_loop/deferred.rs b/nodedb/src/data/executor/core_loop/deferred.rs index 6c4762441..aa7c2daf0 100644 --- a/nodedb/src/data/executor/core_loop/deferred.rs +++ b/nodedb/src/data/executor/core_loop/deferred.rs @@ -24,6 +24,8 @@ pub(in crate::data::executor) struct DeferredWrite { pub source: EventSource, pub new_value: Option>, pub old_value: Option>, + /// The image that did not render; its slot above is `None`. + pub image_fault: Option, } impl CoreLoop { @@ -69,6 +71,7 @@ impl CoreLoop { user_id: None, statement_digest: None, commit_hlc, + image_fault: write.image_fault, }; self.send_write_event(event); diff --git a/nodedb/src/data/executor/core_loop/event_emit.rs b/nodedb/src/data/executor/core_loop/event_emit.rs index d942d87c7..47ccea942 100644 --- a/nodedb/src/data/executor/core_loop/event_emit.rs +++ b/nodedb/src/data/executor/core_loop/event_emit.rs @@ -3,9 +3,9 @@ use std::borrow::Cow; use std::sync::Arc; -use nodedb_query::msgpack_scan; - use super::CoreLoop; +use super::event_image::{StoredImage, UndecodableImage}; +use crate::event::image_fault::ImageFault; /// Bundled arguments for [`CoreLoop::emit_graph_edge_event`]. pub(in crate::data::executor) struct GraphEdgeEvent<'a> { @@ -27,103 +27,29 @@ pub(in crate::data::executor) struct RowWriteEvent<'a> { pub row_id: crate::event::types::RowId, pub new_value: Option<&'a [u8]>, pub old_value: Option<&'a [u8]>, + /// The image that did not render; its slot above is `None`. + pub image_fault: Option, } impl CoreLoop { - /// Convert stored bytes to msgpack for Event Plane consumption. - /// - /// For strict collections, the stored format is Binary Tuple which the - /// Event Plane cannot decode (it lacks the schema). This method converts - /// Binary Tuple → msgpack so triggers can deserialize the payload. - /// Returns `None` for schemaless collections (already msgpack). - pub(in crate::data::executor) fn resolve_event_payload( - &self, - database_id: u64, - tid: u64, - collection: &str, - stored_bytes: &[u8], - ) -> Option> { - let config_key = ( - crate::types::DatabaseId::new(database_id), - crate::types::TenantId::new(tid), - collection.to_string(), - ); - let config = self.doc_configs.get(&config_key)?; - if let nodedb_physical::physical_plan::StorageMode::Strict { ref schema } = - config.storage_mode - { - crate::data::executor::strict_format::binary_tuple_to_msgpack(stored_bytes, schema) - } else { - None - } - } - - /// A stored document row as the Event Plane reads it. A strict Binary - /// Tuple becomes MessagePack. A schemaless body gains its `id`. - pub(in crate::data::executor) fn stored_event_image( - &self, - database_id: u64, - tid: u64, - collection: &str, - identity: &str, - stored: &[u8], - ) -> Vec { - match self.resolve_event_payload(database_id, tid, collection, stored) { - Some(converted) => converted, - None => self.body_event_image(database_id, tid, collection, identity, stored), - } - } - - /// A MessagePack document body as the Event Plane reads it. A schemaless - /// body gains its `id`, the identity every read injects. - pub(in crate::data::executor) fn body_event_image( - &self, - database_id: u64, - tid: u64, - collection: &str, - identity: &str, - body: &[u8], - ) -> Vec { - if self.is_schemaless_document_collection(database_id, tid, collection) { - msgpack_scan::inject_str_field(body, "id", identity) - } else { - body.to_vec() - } - } - - /// Whether `collection` is a schemaless document collection. - /// - /// A schemaless body carries no storage key of its own, so its `id` field - /// is absent whenever the caller declared no `id` column. A strict row's - /// `id` is a real tuple column, already present after Binary Tuple - /// conversion, so it needs no identity injection. - fn is_schemaless_document_collection( - &self, - database_id: u64, - tid: u64, - collection: &str, - ) -> bool { - let config_key = ( - crate::types::DatabaseId::new(database_id), - crate::types::TenantId::new(tid), - collection.to_string(), - ); - matches!( - self.doc_configs.get(&config_key).map(|c| &c.storage_mode), - Some(nodedb_physical::physical_plan::StorageMode::Schemaless) - ) - } - /// Emit a point write/overwrite/update event derived from the new bytes /// produced by the handler and the prior bytes returned from storage. /// /// Shared by every handler that runs a put-style mutation against a /// document engine: PointPut, Upsert (both branches), batched PointPut, /// columnar-row overwrite. Each of these knows its *new* bytes and - /// receives *prior* bytes from the storage API; the Event Plane - /// payload (`new_value` / `old_value`) is derived from both after - /// applying the strict→msgpack shim, and the `WriteOp` tag is computed - /// from their presence. + /// receives *prior* bytes from the storage API. Both become MessagePack + /// images through [`Self::stored_event_image`], and the `WriteOp` tag is + /// computed from the prior's presence. + /// + /// A schemaless body that lacks its identity column carries its identity + /// only in the storage key. The image gains that identity under the + /// identity column, as every read path injects it via `sparse_row_to_doc`, + /// so a WHEN filter, CDC, or change stream reads the same row a query + /// returns. This identity also becomes the `WriteEvent.row_id`. + /// + /// A stored strict row that does not decode has no image: its slot is + /// `None` and the event names the fault. pub(in crate::data::executor) fn emit_put_event( &mut self, task: &super::super::task::ExecutionTask, @@ -134,73 +60,67 @@ impl CoreLoop { prior_stored: Option<&[u8]>, ) { let database_id = task.request.database_id.as_u64(); - let new_converted = self.resolve_event_payload(database_id, tid, collection, new_stored); - let old_converted = - prior_stored.and_then(|p| self.resolve_event_payload(database_id, tid, collection, p)); - let old_bytes: Option<&[u8]> = match (prior_stored, old_converted.as_deref()) { - (Some(_), Some(c)) => Some(c), - (Some(raw), None) => Some(raw), - (None, _) => None, - }; - let op = if old_bytes.is_some() { + let id = identity.as_str(); + let new_image = self.stored_event_image(database_id, tid, collection, id, new_stored); + let old_image = prior_stored + .map(|prior| self.stored_event_image(database_id, tid, collection, id, prior)); + let op = if prior_stored.is_some() { crate::event::WriteOp::Update } else { crate::event::WriteOp::Insert }; - - // A schemaless body with no declared `id` column carries its identity - // only in the storage key. The caller decided that identity already; - // this injects it, the same string every read path injects via - // `sparse_row_to_doc`, so a WHEN filter, CDC, or change stream reads - // the same `id` a query returns. This identity also becomes the - // `WriteEvent.row_id` that fans out to CDC serialization, streaming - // materialized views, event-trigger SQL generation, CRDT sync - // packaging, and webhook delivery. `inject_str_field` is a no-op - // when the body already carries `id`, so a declared primary key is - // never overwritten. - let doc_id = self - .is_schemaless_document_collection(database_id, tid, collection) - .then_some(identity.as_str()); - let new_final: Cow<[u8]> = match (new_converted.as_deref(), doc_id) { - (Some(c), _) => Cow::Borrowed(c), - (None, Some(id)) => Cow::Owned(msgpack_scan::inject_str_field(new_stored, "id", id)), - (None, None) => Cow::Borrowed(new_stored), - }; - let old_final: Option> = - old_bytes.map(|b| match (old_converted.is_some(), doc_id) { - (false, Some(id)) => Cow::Owned(msgpack_scan::inject_str_field(b, "id", id)), - _ => Cow::Borrowed(b), - }); - - self.emit_write_event( + let image_fault = ImageFault::of( + new_image.is_err(), + matches!(old_image, Some(Err(UndecodableImage))), + ); + self.emit_event_with_row_id_as( task, - collection, - op, - identity, - Some(new_final.as_ref()), - old_final.as_deref(), + task.request.event_source, + RowWriteEvent { + collection, + op, + row_id: crate::event::types::RowId::row(identity), + new_value: new_image.as_deref().ok(), + old_value: old_image.as_ref().and_then(|image| image.as_deref().ok()), + image_fault, + }, ); } /// Emit a document-row DELETE event to the Event Plane. /// - /// `identity` is the client-visible identity of the deleted row. The - /// caller decides the encoding, mirroring [`Self::emit_put_event`] on - /// the put side. + /// `identity` is the client-visible identity of the deleted row. + /// `old_stored` is the row as storage held it. A strict Binary Tuple + /// becomes MessagePack. A stored strict row that does not decode has no + /// image: the slot is `None` and the event names the fault. pub(in crate::data::executor) fn emit_document_delete_event( &mut self, task: &super::super::task::ExecutionTask, + tid: u64, collection: &str, identity: crate::engine::document::store::RowIdentity, - old_value: Option<&[u8]>, + old_stored: Option<&[u8]>, ) { - self.emit_write_event( + let database_id = task.request.database_id.as_u64(); + let old_image = old_stored.map(|stored| { + match self.resolve_event_payload(database_id, tid, collection, stored) { + StoredImage::Converted(converted) => Ok(Cow::Owned(converted)), + StoredImage::AsStored => Ok(Cow::Borrowed(stored)), + StoredImage::Undecodable => Err(UndecodableImage), + } + }); + let image_fault = ImageFault::of(false, matches!(old_image, Some(Err(UndecodableImage)))); + self.emit_event_with_row_id_as( task, - collection, - crate::event::WriteOp::Delete, - identity, - None, - old_value, + task.request.event_source, + RowWriteEvent { + collection, + op: crate::event::WriteOp::Delete, + row_id: crate::event::types::RowId::row(identity), + new_value: None, + old_value: old_image.as_ref().and_then(|image| image.as_deref().ok()), + image_fault, + }, ); } @@ -209,44 +129,13 @@ impl CoreLoop { self.events.producer = Some(producer); } - /// Emit a write event for one row to the Event Plane. - /// - /// Called after a successful write (PointPut, PointDelete, PointUpdate, - /// BatchInsert, BulkDelete, atomic KV ops, etc.). The Data Plane NEVER - /// blocks here — if the ring buffer is full, the event is dropped and - /// the Event Plane will detect the gap via sequence numbers and replay - /// from WAL. - /// - /// Prefer [`CoreLoop::emit_put_event`] for any handler that performs a - /// put-style mutation against a document engine — it derives the - /// Insert/Update tag from the prior bytes returned by storage so the - /// emit site cannot disagree with what the row actually did. This - /// lower-level entry point stays for paths where the op is structurally - /// determined by the operation itself (kv-atomic increment, CAS, plain - /// delete) rather than by inspecting pre/post state. - pub(in crate::data::executor) fn emit_write_event( - &mut self, - task: &super::super::task::ExecutionTask, - collection: &str, - op: crate::event::WriteOp, - identity: crate::engine::document::store::RowIdentity, - new_value: Option<&[u8]>, - old_value: Option<&[u8]>, - ) { - self.emit_event_with_row_id( - task, - collection, - op, - crate::event::types::RowId::row(identity), - new_value, - old_value, - ); - } - /// Emit a write event carrying any [`crate::event::types::RowId`]. /// - /// [`Self::emit_write_event`] is the entry point for single rows. Edge - /// events name an `(src, label, dst)` triple and call this directly. + /// Called after a successful write. The Data Plane never blocks here: if + /// the ring buffer is full, the event is dropped and the Event Plane + /// detects the gap by sequence number and replays from the WAL. Document + /// puts and deletes use [`CoreLoop::emit_put_event`] and + /// [`CoreLoop::emit_document_delete_event`], which render the row image. pub(in crate::data::executor) fn emit_event_with_row_id( &mut self, task: &super::super::task::ExecutionTask, @@ -265,6 +154,7 @@ impl CoreLoop { row_id, new_value, old_value, + image_fault: None, }, ); } @@ -283,6 +173,7 @@ impl CoreLoop { row_id, new_value, old_value, + image_fault, } = event; if self.events.producer.is_none() { return; // Event Plane not configured. @@ -313,6 +204,7 @@ impl CoreLoop { user_id: task.request.user_id.clone(), statement_digest: task.request.statement_digest.clone(), commit_hlc: task.request.commit_hlc, + image_fault, }; // The install pass of a committed-redo apply holds its events until @@ -384,6 +276,7 @@ impl CoreLoop { user_id: None, statement_digest: None, commit_hlc: None, + image_fault: None, }; producer.emit(event); diff --git a/nodedb/src/data/executor/core_loop/event_emit_engines.rs b/nodedb/src/data/executor/core_loop/event_emit_engines.rs index cd49ad24c..9f7697678 100644 --- a/nodedb/src/data/executor/core_loop/event_emit_engines.rs +++ b/nodedb/src/data/executor/core_loop/event_emit_engines.rs @@ -72,6 +72,7 @@ impl CoreLoop { row_id: crate::event::types::RowId::row(identity), new_value: new_row.as_deref(), old_value: old_row.as_deref(), + image_fault: None, }, ); } @@ -137,6 +138,7 @@ impl CoreLoop { row_id: crate::event::types::RowId::row(identity), new_value, old_value, + image_fault: None, }, ); } diff --git a/nodedb/src/data/executor/core_loop/event_image.rs b/nodedb/src/data/executor/core_loop/event_image.rs new file mode 100644 index 000000000..0f2a52480 --- /dev/null +++ b/nodedb/src/data/executor/core_loop/event_image.rs @@ -0,0 +1,306 @@ +// SPDX-License-Identifier: BUSL-1.1 + +//! Stored rows as MessagePack images, for the Event Plane and the write-set +//! journal. +//! +//! The Event Plane has no schema, so it reads every row image as +//! MessagePack. A strict collection stores Binary Tuples, which are decoded +//! here against the collection's schema. A strict row that does not decode +//! is corrupt on disk: it is reported once at detection, and its image is +//! withheld, never handed on as raw bytes no reader can decode. + +use nodedb_query::msgpack_scan; + +use super::CoreLoop; +use crate::data::executor::strict_format::strict_row_to_msgpack; + +/// How a stored row reads as MessagePack. +#[derive(Debug, PartialEq, Eq)] +pub(in crate::data::executor) enum StoredImage { + /// The stored bytes are MessagePack already: a schemaless body, a row of + /// a collection this core holds no strict schema for, or a MessagePack + /// map stored before its collection became strict. + AsStored, + /// A strict Binary Tuple, decoded to MessagePack. + Converted(Vec), + /// A strict row that decodes neither as a Binary Tuple nor as a + /// MessagePack map. The recorder report is already filed. + Undecodable, +} + +/// A stored strict row with no MessagePack image. The recorder report is +/// already filed where the decode failed. +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub(in crate::data::executor) struct UndecodableImage; + +impl CoreLoop { + /// How `stored`, a row of `collection`, reads as MessagePack. + pub(in crate::data::executor) fn resolve_event_payload( + &self, + database_id: u64, + tid: u64, + collection: &str, + stored: &[u8], + ) -> StoredImage { + let config_key = ( + crate::types::DatabaseId::new(database_id), + crate::types::TenantId::new(tid), + collection.to_string(), + ); + let Some(config) = self.doc_configs.get(&config_key) else { + return StoredImage::AsStored; + }; + let nodedb_physical::physical_plan::StorageMode::Strict { ref schema } = + config.storage_mode + else { + return StoredImage::AsStored; + }; + match strict_row_to_msgpack(stored, schema) { + Ok(Some(converted)) => StoredImage::Converted(converted), + Ok(None) => StoredImage::AsStored, + Err(fault) => { + tracing::error!( + core = self.core_id, + collection, + fault = fault.as_str(), + "stored strict row does not decode; its image is withheld" + ); + crate::diag::strict_row_image_unrendered(collection, fault.as_str()); + StoredImage::Undecodable + } + } + } + + /// A stored document row as the Event Plane reads it. A strict Binary + /// Tuple becomes MessagePack. A schemaless body gains its identity under + /// the collection's identity column when it lacks that column. + pub(in crate::data::executor) fn stored_event_image( + &self, + database_id: u64, + tid: u64, + collection: &str, + identity: &str, + stored: &[u8], + ) -> Result, UndecodableImage> { + match self.resolve_event_payload(database_id, tid, collection, stored) { + StoredImage::Converted(converted) => Ok(converted), + StoredImage::AsStored => { + Ok(self.body_event_image(database_id, tid, collection, identity, stored)) + } + StoredImage::Undecodable => Err(UndecodableImage), + } + } + + /// A MessagePack document body as the Event Plane reads it. A schemaless + /// body that lacks its identity column gains the identity there, as every + /// read injects it. A declared-key body holds its key and gains no `id`. + pub(in crate::data::executor) fn body_event_image( + &self, + database_id: u64, + tid: u64, + collection: &str, + identity: &str, + body: &[u8], + ) -> Vec { + if self.is_schemaless_document_collection(database_id, tid, collection) { + let identity_column = self.identity_column(database_id, tid, collection); + msgpack_scan::inject_str_field(body, &identity_column, identity) + } else { + body.to_vec() + } + } + + /// Whether `collection` is a schemaless document collection. + /// + /// A schemaless body carries no storage key of its own, so its `id` field + /// is absent whenever the caller declared no `id` column. A strict row's + /// `id` is a real tuple column, already present after Binary Tuple + /// conversion, so it needs no identity injection. + fn is_schemaless_document_collection( + &self, + database_id: u64, + tid: u64, + collection: &str, + ) -> bool { + let config_key = ( + crate::types::DatabaseId::new(database_id), + crate::types::TenantId::new(tid), + collection.to_string(), + ); + matches!( + self.doc_configs.get(&config_key).map(|c| &c.storage_mode), + Some(nodedb_physical::physical_plan::StorageMode::Schemaless) + ) + } +} + +#[cfg(test)] +mod tests { + use nodedb_physical::physical_plan::StorageMode; + use nodedb_types::columnar::{ColumnDef, ColumnType, StrictSchema}; + use nodedb_types::value::Value; + + use super::*; + use crate::data::executor::core_loop::redo_image::StoredRow; + use crate::data::executor::core_loop::tests::{make_core_with_dir, make_default_task}; + use crate::engine::document::store::{CollectionConfig, RowIdentity}; + use crate::event::bus::{EventConsumerRx, create_event_bus_with_capacity}; + use crate::event::image_fault::ImageFault; + use crate::event::{WriteEvent, WriteOp}; + use crate::types::{DatabaseId, TenantId}; + + const TID: u64 = 1; + const COLL: &str = "strict_images"; + /// Bytes no strict row is stored as: neither a Binary Tuple nor a map. + const CORRUPT: &[u8] = b"corrupt row bytes"; + + fn schema() -> StrictSchema { + StrictSchema::new(vec![ + ColumnDef::required("id", ColumnType::String), + ColumnDef::nullable("name", ColumnType::String), + ]) + .expect("valid schema") + } + + fn strict_core(dir: &std::path::Path) -> (CoreLoop, EventConsumerRx) { + let (mut core, _tx, _rx) = make_core_with_dir(dir); + core.seed_doc_configs(&[( + (DatabaseId::DEFAULT, TenantId::new(TID), COLL.to_string()), + CollectionConfig::new(COLL).with_storage_mode(StorageMode::Strict { schema: schema() }), + )]); + let (mut producers, mut consumers) = create_event_bus_with_capacity(1, 16); + core.set_event_producer(producers.pop().expect("producer")); + (core, consumers.pop().expect("consumer")) + } + + fn tuple(name: &str) -> Vec { + let row = Value::Object( + [ + ("id".to_string(), Value::String("r1".into())), + ("name".to_string(), Value::String(name.into())), + ] + .into_iter() + .collect(), + ); + crate::data::executor::strict_format::value_to_binary_tuple(&row, &schema(), COLL) + .expect("row encodes") + } + + fn drain(events: &mut EventConsumerRx) -> Vec { + std::iter::from_fn(|| events.try_recv()).collect() + } + + fn name_of(image: &[u8]) -> String { + nodedb_types::value_from_msgpack(image) + .expect("image is MessagePack") + .as_object() + .and_then(|o| o.get("name")) + .and_then(|v| v.as_str()) + .expect("name field") + .to_owned() + } + + #[test] + fn a_tuple_converts_and_corrupt_bytes_are_undecodable() { + let dir = tempfile::tempdir().expect("tempdir"); + let (core, _events) = strict_core(dir.path()); + match core.resolve_event_payload(0, TID, COLL, &tuple("alice")) { + StoredImage::Converted(image) => assert_eq!(name_of(&image), "alice"), + other => panic!("a tuple converts, got {other:?}"), + } + assert_eq!( + core.resolve_event_payload(0, TID, COLL, CORRUPT), + StoredImage::Undecodable + ); + // A collection with no strict schema reads its bytes as stored. + assert_eq!( + core.resolve_event_payload(0, TID, "schemaless", CORRUPT), + StoredImage::AsStored + ); + } + + #[test] + fn a_corrupt_prior_row_withholds_the_old_image_and_names_the_fault() { + let dir = tempfile::tempdir().expect("tempdir"); + let (mut core, mut events) = strict_core(dir.path()); + let task = make_default_task(); + core.emit_put_event( + &task, + TID, + COLL, + RowIdentity::from_user_key("r1"), + &tuple("bob"), + Some(CORRUPT), + ); + let emitted = drain(&mut events); + assert_eq!(emitted.len(), 1); + let event = &emitted[0]; + assert_eq!(event.op, WriteOp::Update); + assert_eq!(event.image_fault, Some(ImageFault::Old)); + assert!( + event.old_value.is_none(), + "raw bytes never reach the Event Plane" + ); + assert_eq!( + name_of(event.new_value.as_deref().expect("new image")), + "bob" + ); + } + + #[test] + fn a_corrupt_deleted_row_emits_no_raw_bytes() { + let dir = tempfile::tempdir().expect("tempdir"); + let (mut core, mut events) = strict_core(dir.path()); + let task = make_default_task(); + core.emit_document_delete_event( + &task, + TID, + COLL, + RowIdentity::from_user_key("r1"), + Some(CORRUPT), + ); + let emitted = drain(&mut events); + assert_eq!(emitted.len(), 1); + assert_eq!(emitted[0].op, WriteOp::Delete); + assert_eq!(emitted[0].image_fault, Some(ImageFault::Old)); + assert!(emitted[0].old_value.is_none()); + + // A decodable prior converts and names no fault. + core.emit_document_delete_event( + &task, + TID, + COLL, + RowIdentity::from_user_key("r1"), + Some(&tuple("carol")), + ); + let emitted = drain(&mut events); + assert_eq!(emitted[0].image_fault, None); + assert_eq!( + name_of(emitted[0].old_value.as_deref().expect("old image")), + "carol" + ); + } + + #[test] + fn a_corrupt_post_image_refuses_its_write_set_and_stops_the_core() { + let dir = tempfile::tempdir().expect("tempdir"); + let (mut core, _events) = strict_core(dir.path()); + let row = || StoredRow { + database_id: 0, + tid: TID, + collection: COLL, + surrogate: 7, + identity: RowIdentity::from_user_key("r1"), + }; + assert!(core.stored_row_image(row(), &tuple("dave"), None).is_ok()); + assert!(!core.is_fail_stopped()); + let error = core + .stored_row_image(row(), CORRUPT, None) + .expect_err("a corrupt row journals no redo body"); + assert!( + matches!(error, crate::Error::Serialization { .. }), + "{error:?}" + ); + assert!(core.is_fail_stopped()); + } +} diff --git a/nodedb/src/data/executor/core_loop/redo_image.rs b/nodedb/src/data/executor/core_loop/redo_image.rs index b43b7a760..bc31c16d8 100644 --- a/nodedb/src/data/executor/core_loop/redo_image.rs +++ b/nodedb/src/data/executor/core_loop/redo_image.rs @@ -9,36 +9,41 @@ //! collection carries the version key it landed at. use crate::bridge::envelope::{RowVersion, WriteSetEntry}; +use crate::data::executor::strict_format::undecodable_strict_row; use crate::engine::document::store::RowIdentity; use super::CoreLoop; +use super::event_image::StoredImage; +use super::fail_stop::FailStopCause; impl CoreLoop { - /// The redo body of `stored`, a row as `collection` stores it: MessagePack - /// for a strict collection's Binary Tuple, the bytes themselves otherwise. - pub(in crate::data::executor) fn redo_body( - &self, - database_id: u64, - tid: u64, - collection: &str, - stored: &[u8], - ) -> Vec { - self.resolve_event_payload(database_id, tid, collection, stored) - .unwrap_or_else(|| stored.to_vec()) - } - /// The write-set entry of a row a write stored as `stored`. /// `sys_from_ms` is the version's system time on a `bitemporal=true` /// collection, whose version is valid for all time. + /// + /// The redo body is MessagePack: a strict collection's Binary Tuple is + /// decoded, other bytes are journalled as stored. A strict row that does + /// not decode is an error: replay would encode raw tuple bytes as a + /// MessagePack body and store a different row. The write already landed + /// and its write set cannot be journalled, so the core fail-stops. pub(in crate::data::executor) fn stored_row_image( - &self, + &mut self, row: StoredRow<'_>, stored: &[u8], sys_from_ms: Option, - ) -> WriteSetEntry { - let body = self.redo_body(row.database_id, row.tid, row.collection, stored); - WriteSetEntry::put(row.surrogate, row.identity, body) - .versioned(sys_from_ms.map(RowVersion::open)) + ) -> crate::Result { + let body = + match self.resolve_event_payload(row.database_id, row.tid, row.collection, stored) { + StoredImage::Converted(converted) => converted, + StoredImage::AsStored => stored.to_vec(), + StoredImage::Undecodable => { + let error = undecodable_strict_row(row.collection, row.identity.as_str()); + self.fail_stop_core(FailStopCause::WriteSetUnpersisted, &error.to_string()); + return Err(error); + } + }; + Ok(WriteSetEntry::put(row.surrogate, row.identity, body) + .versioned(sys_from_ms.map(RowVersion::open))) } } diff --git a/nodedb/src/data/executor/handlers/timeseries/events.rs b/nodedb/src/data/executor/handlers/timeseries/events.rs index 33bd5a12c..93dc78693 100644 --- a/nodedb/src/data/executor/handlers/timeseries/events.rs +++ b/nodedb/src/data/executor/handlers/timeseries/events.rs @@ -71,6 +71,7 @@ impl CoreLoop { row_id: RowId::Batch, new_value: Some(&image), old_value: None, + image_fault: None, }, ); } diff --git a/nodedb/src/data/executor/handlers/transaction/redo_apply/events.rs b/nodedb/src/data/executor/handlers/transaction/redo_apply/events.rs index ac7a23295..124be36e4 100644 --- a/nodedb/src/data/executor/handlers/transaction/redo_apply/events.rs +++ b/nodedb/src/data/executor/handlers/transaction/redo_apply/events.rs @@ -32,6 +32,7 @@ use crate::data::executor::core_loop::KvWriteEvent; use crate::data::executor::core_loop::deferred::DeferredWrite; use crate::data::executor::core_loop::event_emit::RowWriteEvent; use crate::data::executor::task::ExecutionTask; +use crate::event::image_fault::ImageFault; use crate::event::types::RowId; use crate::event::{EventSource, WriteOp}; use crate::wal::{RedoPublish, RowSourceIndex}; @@ -189,14 +190,19 @@ impl CoreLoop { let new_value = write.new_body.as_deref().map(|body| { self.body_event_image(database_id, tid, &write.collection, id, body) }); - let old_value = write.old_value.as_deref().map(|stored| { + // A stored strict row that does not decode has no image: the + // event names the fault instead. + let old_image = write.old_value.as_deref().map(|stored| { self.stored_event_image(database_id, tid, &write.collection, id, stored) }); + let image_fault = ImageFault::of(false, matches!(old_image, Some(Err(_)))); + let old_value = old_image.and_then(Result::ok); let source = committed_source(rows, record_source, document_source, &write.collection, id); DeferredWrite { new_value, old_value, + image_fault, source, collection: write.collection, op: write.op, @@ -247,6 +253,7 @@ impl CoreLoop { row_id: RowId::Batch, new_value: Some(&value), old_value: None, + image_fault: None, }, ); } diff --git a/nodedb/src/data/executor/strict_format/render.rs b/nodedb/src/data/executor/strict_format/render.rs new file mode 100644 index 000000000..0579e1e4a --- /dev/null +++ b/nodedb/src/data/executor/strict_format/render.rs @@ -0,0 +1,151 @@ +// SPDX-License-Identifier: BUSL-1.1 + +//! Render a stored strict row as standard MessagePack, naming why it failed. +//! +//! A strict collection stores Binary Tuples. A MessagePack map stored before +//! the collection became strict is read as it is. Any other stored form is +//! corrupt: the image readers (the Event Plane, the write-set journal) never +//! receive it, so each failure carries the class that names its cause. + +use nodedb_query::msgpack_scan::reader::skip_value; +use nodedb_types::columnar::StrictSchema; + +use super::decode::binary_tuple_to_value; + +/// Why a stored strict row did not render as MessagePack. Each class names +/// one root cause, never a row, so a report groups by it. +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub(crate) enum StrictImageFault { + /// The bytes are neither a Binary Tuple header nor one well-formed + /// MessagePack map: a truncated or overwritten body. + NotATuple, + /// The tuple names schema version 0, or a version newer than the schema + /// this core holds for the collection. + SchemaVersion, + /// The tuple header is valid but its column layout does not decode. + Layout, + /// The decoded row did not encode as MessagePack. + Encode, +} + +impl StrictImageFault { + pub(crate) fn as_str(self) -> &'static str { + match self { + Self::NotATuple => "not_a_tuple", + Self::SchemaVersion => "schema_version", + Self::Layout => "layout", + Self::Encode => "encode", + } + } +} + +/// The MessagePack image of `stored`, a row of a strict collection. +/// +/// `Ok(None)` when `stored` is one well-formed MessagePack map already, so +/// the stored bytes are the image. `Ok(Some(_))` for a Binary Tuple decoded +/// against `schema`. +pub(crate) fn strict_row_to_msgpack( + stored: &[u8], + schema: &StrictSchema, +) -> Result>, StrictImageFault> { + let is_map_header = stored + .first() + .is_some_and(|&first| (0x80..=0x8F).contains(&first) || first == 0xDE || first == 0xDF); + if is_map_header { + return match skip_value(stored, 0) { + Some(end) if end == stored.len() => Ok(None), + _ => Err(StrictImageFault::NotATuple), + }; + } + let version = nodedb_strict::TupleDecoder::new(schema) + .schema_version(stored) + .map_err(|_| StrictImageFault::NotATuple)?; + if version == 0 || version > schema.version { + return Err(StrictImageFault::SchemaVersion); + } + let value = binary_tuple_to_value(stored, schema).ok_or(StrictImageFault::Layout)?; + nodedb_types::value_to_msgpack(&value) + .map(Some) + .map_err(|_| StrictImageFault::Encode) +} + +#[cfg(test)] +mod tests { + use nodedb_types::columnar::{ColumnDef, ColumnType}; + use nodedb_types::value::Value; + + use super::*; + use crate::data::executor::strict_format::value_to_binary_tuple; + + fn schema() -> StrictSchema { + StrictSchema::new(vec![ + ColumnDef::required("id", ColumnType::String), + ColumnDef::nullable("name", ColumnType::String), + ]) + .expect("valid schema") + } + + fn tuple() -> Vec { + let row = Value::Object( + [ + ("id".to_string(), Value::String("r1".into())), + ("name".to_string(), Value::String("alice".into())), + ] + .into_iter() + .collect(), + ); + value_to_binary_tuple(&row, &schema(), "c").expect("row encodes") + } + + #[test] + fn tuple_renders_as_msgpack_map() { + let image = strict_row_to_msgpack(&tuple(), &schema()) + .expect("tuple decodes") + .expect("a tuple converts"); + let value = nodedb_types::value_from_msgpack(&image).expect("image is msgpack"); + assert_eq!( + value + .as_object() + .and_then(|o| o.get("name")) + .and_then(|v| v.as_str()), + Some("alice") + ); + } + + #[test] + fn stored_msgpack_map_is_its_own_image() { + let map = nodedb_types::value_to_msgpack(&Value::Object( + [("id".to_string(), Value::String("r1".into()))] + .into_iter() + .collect(), + )) + .expect("encode"); + assert_eq!(strict_row_to_msgpack(&map, &schema()), Ok(None)); + } + + #[test] + fn corrupt_rows_name_their_fault() { + let schema = schema(); + assert_eq!( + strict_row_to_msgpack(b"garbage bytes", &schema), + Err(StrictImageFault::NotATuple) + ); + // A map header whose body is cut short is no map. + assert_eq!( + strict_row_to_msgpack(&[0x81, 0xa2, b'i'], &schema), + Err(StrictImageFault::NotATuple) + ); + let mut newer = tuple(); + newer[5..9].copy_from_slice(&(schema.version + 1).to_le_bytes()); + assert_eq!( + strict_row_to_msgpack(&newer, &schema), + Err(StrictImageFault::SchemaVersion) + ); + let full = tuple(); + let truncated = &full[..full.len() - 3]; + assert_eq!( + strict_row_to_msgpack(truncated, &schema), + Err(StrictImageFault::Layout) + ); + } +} diff --git a/nodedb/src/diag/context/event_image.rs b/nodedb/src/diag/context/event_image.rs new file mode 100644 index 000000000..cc0e07cec --- /dev/null +++ b/nodedb/src/diag/context/event_image.rs @@ -0,0 +1,85 @@ +// SPDX-License-Identifier: BUSL-1.1 + +//! Forensic payloads for a stored row whose MessagePack image does not +//! render or decode. + +use faultbox::DomainContext; +use faultbox::serde_json::{Value, json}; + +/// A stored strict row that decodes neither as a Binary Tuple of its +/// collection's schema nor as a MessagePack map. Its image is withheld from +/// the Event Plane and from the write-set journal. +pub(in crate::diag) struct StrictRowImageUnrendered<'a> { + /// Collection whose stored row did not render. + pub collection: &'a str, + /// Stable class of the decode failure (`not_a_tuple`, `schema_version`, + /// `layout`, `encode`). + pub fault: &'static str, +} + +impl DomainContext for StrictRowImageUnrendered<'_> { + fn domain_kind(&self) -> &'static str { + "nodedb.strict_row_image_unrendered" + } + + fn grouping_key(&self) -> String { + // Collection and fault class name the root cause. The row is the + // occurrence, so a scan over many bad rows files one report. + format!("collection={};fault={}", self.collection, self.fault) + } + + fn to_json(&self) -> Value { + json!({ + "collection": self.collection, + "fault": self.fault, + "why_fatal": "the row's stored bytes do not decode against the schema this core \ + holds. A trigger, change stream, materialized view, or CRDT peer \ + that received the raw bytes would read them as a different row or \ + fail to read them. The event is dead-lettered without the image, \ + and a write whose journal needs the image is refused", + "operator_action": "read the named collection's rows directly. One bad row points \ + at a truncated or overwritten body. Every row failing with \ + `schema_version` points at a core whose schema is older than \ + the rows it stores. Restore the collection from a snapshot or \ + rewrite the affected rows", + }) + } +} + +/// A landed timeseries row whose stored MessagePack image does not decode, +/// so the statement's `RETURNING` set cannot carry it. +pub(in crate::diag) struct TimeseriesRowImageUndecodable<'a> { + /// Collection the row landed in. + pub collection: &'a str, + /// Path that decoded the image. + pub site: &'static str, + /// The error class: the error text before its first colon. + pub error_class: &'a str, +} + +impl DomainContext for TimeseriesRowImageUndecodable<'_> { + fn domain_kind(&self) -> &'static str { + "nodedb.timeseries_row_image_undecodable" + } + + fn grouping_key(&self) -> String { + // The collection names the root cause. The row is the occurrence, + // so a statement over many bad images files one report. + format!("collection={};site={}", self.collection, self.site) + } + + fn to_json(&self) -> Value { + json!({ + "collection": self.collection, + "site": self.site, + "error_class": self.error_class, + "why_fatal": "the image is the row as a scan reads it, built when the ingest \ + resolved. An image that does not decode means the resolved record \ + is damaged. The statement is refused instead of returning fewer \ + rows than it stored", + "operator_action": "read the named collection's rows directly and compare them \ + with the ingest's input. A repeat on every ingest points at \ + the resolver that builds the image", + }) + } +} diff --git a/nodedb/src/diag/recording/event_image.rs b/nodedb/src/diag/recording/event_image.rs new file mode 100644 index 000000000..275e4268e --- /dev/null +++ b/nodedb/src/diag/recording/event_image.rs @@ -0,0 +1,44 @@ +// SPDX-License-Identifier: BUSL-1.1 + +//! Capture sites for a stored row whose MessagePack image does not render +//! or decode. + +use faultbox::{Capture, EventKind, error_chain_of}; + +use super::shared::error_class; +use crate::diag::context; + +/// Report a stored strict row that does not render as MessagePack. Called +/// only from the Data Plane's image resolution, where the decode fails. +/// `fault` is the stable class of the failure. +pub fn strict_row_image_unrendered(collection: &str, fault: &'static str) { + let ctx = context::StrictRowImageUnrendered { collection, fault }; + let _ = Capture::new( + EventKind::Corruption, + "stored strict row did not render as MessagePack: its image is withheld", + ) + .domain(&ctx) + .with_backtrace() + .emit(); +} + +/// Report a landed timeseries row whose stored image does not decode as +/// MessagePack. Called only from the resolved timeseries ingest, where the +/// `RETURNING` decode fails. The caller returns the error alongside this +/// report. +pub fn timeseries_row_image_undecodable(err: &crate::Error, collection: &str, site: &'static str) { + let class = error_class(err); + let ctx = context::TimeseriesRowImageUndecodable { + collection, + site, + error_class: &class, + }; + let _ = Capture::new( + EventKind::Corruption, + "landed timeseries row image does not decode, so the RETURNING statement is refused", + ) + .error_chain(error_chain_of(err)) + .domain(&ctx) + .with_backtrace() + .emit(); +} diff --git a/nodedb/src/event/audit_dml/consumer.rs b/nodedb/src/event/audit_dml/consumer.rs index b216e7b3d..c166832a9 100644 --- a/nodedb/src/event/audit_dml/consumer.rs +++ b/nodedb/src/event/audit_dml/consumer.rs @@ -136,6 +136,7 @@ mod tests { user_id: Some(Arc::from("alice")), statement_digest: None, commit_hlc: Some(crate::event::test_utils::test_commit_hlc()), + image_fault: None, } } diff --git a/nodedb/src/event/bus.rs b/nodedb/src/event/bus.rs index 59e802415..a66633712 100644 --- a/nodedb/src/event/bus.rs +++ b/nodedb/src/event/bus.rs @@ -281,6 +281,7 @@ mod tests { user_id: None, statement_digest: None, commit_hlc: Some(crate::event::test_utils::test_commit_hlc()), + image_fault: None, } } diff --git a/nodedb/src/event/cdc/router.rs b/nodedb/src/event/cdc/router.rs index 8a6bb07c5..685dff8c3 100644 --- a/nodedb/src/event/cdc/router.rs +++ b/nodedb/src/event/cdc/router.rs @@ -480,6 +480,7 @@ mod tests { user_id: None, statement_digest: None, commit_hlc: Some(crate::event::test_utils::test_commit_hlc()), + image_fault: None, } } diff --git a/nodedb/src/event/consumer/delivery.rs b/nodedb/src/event/consumer/delivery.rs index e5c2acdff..f9671d7bd 100644 --- a/nodedb/src/event/consumer/delivery.rs +++ b/nodedb/src/event/consumer/delivery.rs @@ -116,6 +116,7 @@ mod tests { user_id: None, statement_digest: None, commit_hlc: Some(crate::event::test_utils::test_commit_hlc()), + image_fault: None, } } diff --git a/nodedb/src/event/consumer/drain.rs b/nodedb/src/event/consumer/drain.rs index ff34aeaa5..63e9a5891 100644 --- a/nodedb/src/event/consumer/drain.rs +++ b/nodedb/src/event/consumer/drain.rs @@ -156,6 +156,7 @@ mod tests { user_id: None, statement_digest: None, commit_hlc: Some(crate::event::test_utils::test_commit_hlc()), + image_fault: None, } } diff --git a/nodedb/src/event/consumer/pipeline.rs b/nodedb/src/event/consumer/pipeline.rs index 2f9fd3de4..80a2679c6 100644 --- a/nodedb/src/event/consumer/pipeline.rs +++ b/nodedb/src/event/consumer/pipeline.rs @@ -72,6 +72,13 @@ async fn deliver_event( .advance_lsn_only(event.vshard_id.as_u32(), event.lsn.as_u64()); return; } + if let Some(fault) = event.image_fault { + dead_letter_unrendered_image(event, fault, shared_state); + shared_state + .watermark_tracker + .advance_lsn_only(event.vshard_id.as_u32(), event.lsn.as_u64()); + return; + } // The key the non-idempotent sinks remember the event by, so an event // delivered again after a restart reaches each of them once. let key = SinkEventKey::of(core_id, event); @@ -87,6 +94,42 @@ async fn deliver_event( accumulate_data_event(event, key.as_ref(), shared_state, cdc_router); } +/// Dead-letter an event whose row image the Data Plane could not render. +/// +/// The stored row is corrupt, so no side effect can act on the write: a +/// trigger would bind a missing NEW or OLD row, a change stream would carry +/// a row without its image, and a view or CRDT peer would apply it as one. +/// The event runs none of them. A durable audit entry names the collection, +/// row, LSN and image, and the Data Plane already filed the recorder report. +fn dead_letter_unrendered_image( + event: &WriteEvent, + fault: crate::event::image_fault::ImageFault, + shared_state: &SharedState, +) { + let detail = format!( + "event dead-lettered: the {} image of row '{}' in '{}' at LSN {} did not render; \ + no trigger, change stream, view or CRDT peer received the write", + fault.as_str(), + event.row_id.as_str(), + event.collection, + event.lsn.as_u64(), + ); + tracing::error!( + collection = %event.collection, + row = event.row_id.as_str(), + lsn = event.lsn.as_u64(), + image = fault.as_str(), + "{detail}" + ); + shared_state.audit_record_with_db( + crate::control::security::audit::AuditEvent::AdminAction, + Some(event.tenant_id), + Some(event.database_id), + "event_plane", + &detail, + ); +} + /// Whether `event` is a row write that triggers and event actions act on. /// /// A graph edge write reaches the stream of the edge's collection for its @@ -178,6 +221,7 @@ mod tests { user_id: None, statement_digest: None, commit_hlc: Some(crate::event::test_utils::test_commit_hlc()), + image_fault: None, } } @@ -190,6 +234,44 @@ mod tests { assert!(!event_actions_required(&event(WriteOp::Heartbeat))); } + #[tokio::test] + async fn an_event_with_an_unrendered_image_is_dead_lettered() { + let dir = tempfile::tempdir().expect("tempdir"); + let (_wal, _watermarks, shared_state, _dlq, cdc_router) = + crate::event::test_utils::event_test_deps(&dir); + let mut faulted = event(WriteOp::Update); + faulted.lsn = Lsn::new(9); + faulted.new_value = Some(Arc::from(&[0x80u8][..])); + faulted.image_fault = Some(crate::event::image_fault::ImageFault::Old); + let mut guard = super::super::delivery::DeliveryGuard::new(Lsn::new(0)); + + let delivered = super::deliver_events( + 0, + std::slice::from_ref(&faulted), + &mut guard, + &shared_state, + &cdc_router, + ) + .await; + + assert_eq!(delivered, 1); + let audit = shared_state.audit.lock().expect("read audit log"); + let entry = audit + .all() + .iter() + .find(|entry| entry.detail.contains("event dead-lettered")) + .expect("the dead-lettered event is audited"); + assert!( + entry + .detail + .contains("old image of row 'row-1' in 'events' at LSN 9") + ); + assert!( + cdc_router.stream_buffers().is_empty(), + "no change stream receives the write" + ); + } + #[test] fn an_edge_write_runs_no_row_actions() { let mut edge = event(WriteOp::Insert); diff --git a/nodedb/src/event/crdt_sync/packager.rs b/nodedb/src/event/crdt_sync/packager.rs index b797bd2ed..982ef4c27 100644 --- a/nodedb/src/event/crdt_sync/packager.rs +++ b/nodedb/src/event/crdt_sync/packager.rs @@ -216,6 +216,7 @@ mod tests { user_id: None, statement_digest: None, commit_hlc: Some(crate::event::test_utils::test_commit_hlc()), + image_fault: None, } } diff --git a/nodedb/src/event/image_fault.rs b/nodedb/src/event/image_fault.rs new file mode 100644 index 000000000..dea330b63 --- /dev/null +++ b/nodedb/src/event/image_fault.rs @@ -0,0 +1,57 @@ +// SPDX-License-Identifier: BUSL-1.1 + +//! A row image the Data Plane could not render for the Event Plane. +//! +//! The Event Plane reads row images as MessagePack. A strict collection +//! stores a Binary Tuple, which the Data Plane decodes against the +//! collection's schema before emitting. A stored row that does not decode is +//! corrupt on disk. Its image slot stays empty and the event names the +//! fault, so no consumer reads an absent image as a row without one. The +//! delivery pipeline dead-letters such an event instead of running its side +//! effects. + +/// Which row image of a write event did not render. +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum ImageFault { + /// The post-image (`new_value`). + New, + /// The pre-image (`old_value`). + Old, + /// Both images. + Both, +} + +impl ImageFault { + /// The fault of an event whose post-image failed when `new` is true and + /// whose pre-image failed when `old` is true. `None` when neither failed. + pub fn of(new: bool, old: bool) -> Option { + match (new, old) { + (true, true) => Some(Self::Both), + (true, false) => Some(Self::New), + (false, true) => Some(Self::Old), + (false, false) => None, + } + } + + /// Stable label for logs and dead-letter records. + pub fn as_str(self) -> &'static str { + match self { + Self::New => "new", + Self::Old => "old", + Self::Both => "new_and_old", + } + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn fault_names_the_failed_images() { + assert_eq!(ImageFault::of(false, false), None); + assert_eq!(ImageFault::of(true, false), Some(ImageFault::New)); + assert_eq!(ImageFault::of(false, true), Some(ImageFault::Old)); + assert_eq!(ImageFault::of(true, true), Some(ImageFault::Both)); + } +} diff --git a/nodedb/src/event/mod.rs b/nodedb/src/event/mod.rs index d3a62a049..a7451c239 100644 --- a/nodedb/src/event/mod.rs +++ b/nodedb/src/event/mod.rs @@ -14,6 +14,7 @@ pub mod crdt_sync; pub mod cross_shard; pub mod field_diff; pub mod graph_cdc; +pub mod image_fault; pub mod interest; pub mod kafka; pub mod metrics; diff --git a/nodedb/src/event/plane.rs b/nodedb/src/event/plane.rs index b0a164bc5..553bd77b2 100644 --- a/nodedb/src/event/plane.rs +++ b/nodedb/src/event/plane.rs @@ -319,6 +319,7 @@ mod tests { user_id: None, statement_digest: None, commit_hlc: Some(crate::event::test_utils::test_commit_hlc()), + image_fault: None, } } diff --git a/nodedb/src/event/record_numbering.rs b/nodedb/src/event/record_numbering.rs index 957ab0022..576d51365 100644 --- a/nodedb/src/event/record_numbering.rs +++ b/nodedb/src/event/record_numbering.rs @@ -109,6 +109,7 @@ mod tests { user_id: None, statement_digest: None, commit_hlc: Some(crate::event::test_utils::test_commit_hlc()), + image_fault: None, } } diff --git a/nodedb/src/event/topic/committed/event.rs b/nodedb/src/event/topic/committed/event.rs index 1a5021e40..fd52f2e25 100644 --- a/nodedb/src/event/topic/committed/event.rs +++ b/nodedb/src/event/topic/committed/event.rs @@ -89,6 +89,7 @@ pub(crate) fn replayed_publish_events( user_id: None, statement_digest: None, commit_hlc: scope.commit_hlc, + image_fault: None, }); } events diff --git a/nodedb/src/event/trigger/lane/held.rs b/nodedb/src/event/trigger/lane/held.rs index 0087ac103..b353141a3 100644 --- a/nodedb/src/event/trigger/lane/held.rs +++ b/nodedb/src/event/trigger/lane/held.rs @@ -153,6 +153,8 @@ impl HeldAction { user_id: None, statement_digest: None, commit_hlc: self.commit_hlc, + // Delivery holds no event whose image did not render. + image_fault: None, }) } } @@ -181,6 +183,7 @@ mod tests { user_id: None, statement_digest: None, commit_hlc: Some(5), + image_fault: None, }; let held = HeldAction::of(&event).expect("a row write is held"); let back = HeldAction::from_bytes(&held.to_bytes().expect("encode")).expect("decode"); @@ -215,6 +218,7 @@ mod tests { user_id: None, statement_digest: None, commit_hlc: None, + image_fault: None, }; assert!(HeldAction::of(&event).is_none()); event.op = WriteOp::Insert; diff --git a/nodedb/src/event/types.rs b/nodedb/src/event/types.rs index f6ac31570..830538943 100644 --- a/nodedb/src/event/types.rs +++ b/nodedb/src/event/types.rs @@ -224,6 +224,12 @@ pub struct WriteEvent { /// event rebuilt from a WAL record, which the record's commit HLC dates, /// and for a heartbeat. pub commit_hlc: Option, + + /// The row image the Data Plane could not render: a stored strict row + /// that does not decode. The named image slot is `None`. Delivery + /// dead-letters the event instead of running its side effects. `None` + /// for every event whose images rendered. + pub image_fault: Option, } /// Where an event sits in the WAL record it reproduces. @@ -613,6 +619,7 @@ mod tests { user_id: None, statement_digest: None, commit_hlc: Some(crate::event::test_utils::test_commit_hlc()), + image_fault: None, }; assert_eq!(event.sequence, 1); assert_eq!(event.op, WriteOp::Insert); diff --git a/nodedb/src/event/wal_replay_kv_shapes.rs b/nodedb/src/event/wal_replay_kv_shapes.rs index 3281e5c11..a2b4f76cc 100644 --- a/nodedb/src/event/wal_replay_kv_shapes.rs +++ b/nodedb/src/event/wal_replay_kv_shapes.rs @@ -98,6 +98,7 @@ pub(super) fn parse_kv_put_family( user_id: None, statement_digest: None, commit_hlc, + image_fault: None, }); } @@ -127,6 +128,7 @@ pub(super) fn parse_kv_put_family( user_id: None, statement_digest: None, commit_hlc, + image_fault: None, }); } @@ -154,6 +156,7 @@ pub(super) fn parse_kv_put_family( user_id: None, statement_digest: None, commit_hlc, + image_fault: None, }); } None diff --git a/nodedb/src/event/wal_replay_parse.rs b/nodedb/src/event/wal_replay_parse.rs index 6bc985997..30dc73c2d 100644 --- a/nodedb/src/event/wal_replay_parse.rs +++ b/nodedb/src/event/wal_replay_parse.rs @@ -87,6 +87,7 @@ pub(super) fn parse_put_record( user_id: None, statement_digest: None, commit_hlc, + image_fault: None, }); } @@ -115,6 +116,7 @@ pub(super) fn parse_put_record( user_id: None, statement_digest: None, commit_hlc, + image_fault: None, }); } @@ -146,6 +148,7 @@ pub(super) fn parse_put_record( user_id: None, statement_digest: None, commit_hlc, + image_fault: None, }); } @@ -197,6 +200,7 @@ pub(super) fn parse_put_record( user_id: None, statement_digest: None, commit_hlc, + image_fault: None, }); } @@ -269,6 +273,7 @@ pub(super) fn parse_graph_node_label_record( user_id: None, statement_digest: None, commit_hlc, + image_fault: None, }) } @@ -312,6 +317,7 @@ pub(super) fn parse_delete_record( user_id: None, statement_digest: None, commit_hlc, + image_fault: None, }); } @@ -338,6 +344,7 @@ pub(super) fn parse_delete_record( user_id: None, statement_digest: None, commit_hlc, + image_fault: None, }); } @@ -364,6 +371,7 @@ pub(super) fn parse_delete_record( user_id: None, statement_digest: None, commit_hlc, + image_fault: None, }); } @@ -388,6 +396,7 @@ pub(super) fn parse_delete_record( user_id: None, statement_digest: None, commit_hlc, + image_fault: None, }); } @@ -434,6 +443,7 @@ pub(super) fn parse_delete_record( user_id: None, statement_digest: None, commit_hlc, + image_fault: None, }); } diff --git a/nodedb/src/event/wal_replay_timeseries.rs b/nodedb/src/event/wal_replay_timeseries.rs index 98937decc..9cce66c01 100644 --- a/nodedb/src/event/wal_replay_timeseries.rs +++ b/nodedb/src/event/wal_replay_timeseries.rs @@ -101,6 +101,7 @@ pub(crate) fn replayed_timeseries_events( user_id: None, statement_digest: None, commit_hlc: scope.commit_hlc, + image_fault: None, }); } events diff --git a/nodedb/tests/inproc/cases/bitemporal_cdc.rs b/nodedb/tests/inproc/cases/bitemporal_cdc.rs index 110879910..99f7ea191 100644 --- a/nodedb/tests/inproc/cases/bitemporal_cdc.rs +++ b/nodedb/tests/inproc/cases/bitemporal_cdc.rs @@ -59,6 +59,7 @@ fn write_event(seq: u64, op: WriteOp, payload_bytes: Vec, is_delete: bool) - .duration_since(std::time::UNIX_EPOCH) .ok() .and_then(|elapsed| u64::try_from(elapsed.as_nanos()).ok()), + image_fault: None, } } diff --git a/nodedb/tests/inproc/cases/cdc_arc_fanout.rs b/nodedb/tests/inproc/cases/cdc_arc_fanout.rs index d3cdc677b..03223b63a 100644 --- a/nodedb/tests/inproc/cases/cdc_arc_fanout.rs +++ b/nodedb/tests/inproc/cases/cdc_arc_fanout.rs @@ -75,6 +75,7 @@ fn write_event(seq: u64) -> WriteEvent { .duration_since(std::time::UNIX_EPOCH) .ok() .and_then(|elapsed| u64::try_from(elapsed.as_nanos()).ok()), + image_fault: None, } } diff --git a/nodedb/tests/inproc/cases/event_trigger.rs b/nodedb/tests/inproc/cases/event_trigger.rs index 37ba7e772..fb298c270 100644 --- a/nodedb/tests/inproc/cases/event_trigger.rs +++ b/nodedb/tests/inproc/cases/event_trigger.rs @@ -46,6 +46,7 @@ fn make_event(source: EventSource, op: WriteOp, collection: &str) -> WriteEvent .duration_since(std::time::UNIX_EPOCH) .ok() .and_then(|elapsed| u64::try_from(elapsed.as_nanos()).ok()), + image_fault: None, } } diff --git a/nodedb/tests/inproc/cases/shutdown_event_plane.rs b/nodedb/tests/inproc/cases/shutdown_event_plane.rs index 830ce0f78..56b71d617 100644 --- a/nodedb/tests/inproc/cases/shutdown_event_plane.rs +++ b/nodedb/tests/inproc/cases/shutdown_event_plane.rs @@ -49,6 +49,7 @@ fn make_write_event(seq: u64, lsn_val: u64) -> WriteEvent { .duration_since(std::time::UNIX_EPOCH) .ok() .and_then(|elapsed| u64::try_from(elapsed.as_nanos()).ok()), + image_fault: None, } } From 67ccb39c50b434bb76ddd7e8a001636c5ad09a84 Mon Sep 17 00:00:00 2001 From: Farhan Syah Date: Thu, 8 Oct 2026 16:51:47 +0800 Subject: [PATCH 13/24] feat(columnar): match flushed rows and evaluate expressions in DML Columnar UPDATE and DELETE match current rows in flushed segments as well as the memtable, on the autocommit, resolve and transaction staging paths. An UPDATE assignment can be an expression evaluated against each row's pre-image. Assignments travel as UpdateValue through the plan, WAL and replication, and a filter that does not decode refuses the statement. --- nodedb-physical/src/physical_plan/columnar.rs | 8 +- nodedb-types/src/columnar/dml_wal_record.rs | 5 +- .../control/server/wal_dispatch/timeseries.rs | 4 +- .../wal_replication/decode/columnar.rs | 4 +- nodedb/src/control/write_resolve/columnar.rs | 2 +- .../executor/handlers/columnar_assignments.rs | 113 +++++++ .../executor/handlers/columnar_resolve.rs | 320 +++++++++++++----- .../executor/handlers/columnar_resolve_dml.rs | 30 +- .../handlers/transaction/resolve/entry.rs | 11 +- .../stage_write/stage_columnar_dml.rs | 111 +++--- .../data/executor/wal_replay_columnar_dml.rs | 66 +++- nodedb/src/wal/columnar_dml_updates.rs | 83 +++++ nodedb/src/wal/mod.rs | 2 + .../cases/columnar_predicate_dml_flushed.rs | 224 ++++++++++++ 14 files changed, 803 insertions(+), 180 deletions(-) create mode 100644 nodedb/src/data/executor/handlers/columnar_assignments.rs create mode 100644 nodedb/src/wal/columnar_dml_updates.rs create mode 100644 nodedb/tests/wire/cases/columnar_predicate_dml_flushed.rs 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-types/src/columnar/dml_wal_record.rs b/nodedb-types/src/columnar/dml_wal_record.rs index 09b06a9f9..02c3731fa 100644 --- a/nodedb-types/src/columnar/dml_wal_record.rs +++ b/nodedb-types/src/columnar/dml_wal_record.rs @@ -42,8 +42,9 @@ pub struct ColumnarDmlWalRecord { pub is_update: bool, /// Serialized `Vec` (MessagePack). pub filters: Vec, - /// Field assignments for `Update`: `(column_name, msgpack_value_bytes)`. - /// Always empty for `Delete`. + /// Field assignments for `Update`: `(column_name, value_bytes)`. Each + /// value is a MessagePack-encoded `UpdateValue`, a literal or an + /// expression over the row. Always empty for `Delete`. pub updates: Vec<(String, Vec)>, } diff --git a/nodedb/src/control/server/wal_dispatch/timeseries.rs b/nodedb/src/control/server/wal_dispatch/timeseries.rs index 37ce39044..443c1c91f 100644 --- a/nodedb/src/control/server/wal_dispatch/timeseries.rs +++ b/nodedb/src/control/server/wal_dispatch/timeseries.rs @@ -233,14 +233,14 @@ pub(crate) fn encode_columnar_dml_payload( collection: &str, is_update: bool, filters: &[u8], - updates: &[(String, Vec)], + updates: &[(String, nodedb_physical::physical_plan::UpdateValue)], ) -> crate::Result> { let record = nodedb_types::columnar::ColumnarDmlWalRecord { kind: "columnar_dml".to_string(), collection: collection.to_string(), is_update, filters: filters.to_vec(), - updates: updates.to_vec(), + updates: crate::wal::encode_columnar_dml_updates(updates)?, }; zerompk::to_msgpack_vec(&record).map_err(|e| crate::Error::Serialization { format: "msgpack".into(), diff --git a/nodedb/src/control/wal_replication/decode/columnar.rs b/nodedb/src/control/wal_replication/decode/columnar.rs index 81250aff3..e24987784 100644 --- a/nodedb/src/control/wal_replication/decode/columnar.rs +++ b/nodedb/src/control/wal_replication/decode/columnar.rs @@ -3,7 +3,7 @@ //! Decode `ReplicatedWrite` variants that produce `PhysicalPlan::Columnar`. use crate::bridge::envelope::PhysicalPlan; -use nodedb_physical::physical_plan::{ColumnarOp, TimeseriesOp}; +use nodedb_physical::physical_plan::{ColumnarOp, TimeseriesOp, UpdateValue}; use nodedb_types::RlsWriteCheck; /// Reconstruct a `ColumnarOp::Truncate` plan. Same idempotent-replay @@ -39,7 +39,7 @@ pub(super) fn bulk_dml( collection: &str, filters: &[u8], is_update: bool, - updates: &[(String, Vec)], + updates: &[(String, UpdateValue)], ) -> PhysicalPlan { if is_update { PhysicalPlan::Columnar(ColumnarOp::Update { diff --git a/nodedb/src/control/write_resolve/columnar.rs b/nodedb/src/control/write_resolve/columnar.rs index 3ce2572ed..e6d147d05 100644 --- a/nodedb/src/control/write_resolve/columnar.rs +++ b/nodedb/src/control/write_resolve/columnar.rs @@ -23,7 +23,7 @@ pub struct ColumnarWriteResolver { /// 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, nodedb_physical::physical_plan::UpdateValue)>, is_update: bool, rls_write_check: RlsWriteCheck, } diff --git a/nodedb/src/data/executor/handlers/columnar_assignments.rs b/nodedb/src/data/executor/handlers/columnar_assignments.rs new file mode 100644 index 000000000..5e9706657 --- /dev/null +++ b/nodedb/src/data/executor/handlers/columnar_assignments.rs @@ -0,0 +1,113 @@ +// SPDX-License-Identifier: BUSL-1.1 + +//! The SET-list of a columnar UPDATE, bound to schema column indexes. +//! +//! Every columnar UPDATE path builds post-images here: the autocommit and +//! COMMIT-replay handler, the write-policy resolve, and the in-transaction +//! staging. One binding decides what a row becomes on every path. +//! +//! A literal decodes once per statement. An expression evaluates once per +//! matched row with the shared `SqlExpr` evaluator. Each expression reads +//! the row's pre-image, so one assignment never observes another. This is +//! PostgreSQL's rule, and the document UPDATE paths follow it too. +//! +//! The post-image then meets the declared column rule. A value past a +//! declared width is an error, so the caller refuses the statement before +//! the first row changes. + +use nodedb_physical::physical_plan::UpdateValue; +use nodedb_query::expr::SqlExpr; +use nodedb_types::Value; +use nodedb_types::columnar::ColumnarSchema; + +use crate::data::executor::handlers::columnar_read::convert::row_to_projected_value; +use crate::data::executor::handlers::columnar_write::coerce_columnar_row; + +/// The value one assignment writes. +enum AssignedValue { + /// A constant, decoded once for the statement. + Literal(Value), + /// An expression over the row's pre-image. + Expr(SqlExpr), +} + +/// A columnar UPDATE's SET-list, bound to the schema it writes. +pub(in crate::data::executor) struct ColumnarAssignments { + /// `(column index, value)` per assignment, in SET-list order. + entries: Vec<(usize, AssignedValue)>, +} + +impl ColumnarAssignments { + /// Bind `updates` to `schema`. An assignment to a column the schema + /// does not hold writes nothing. A literal that does not decode is an + /// `Internal` error. + pub(in crate::data::executor) fn bind( + schema: &ColumnarSchema, + updates: &[(String, UpdateValue)], + ) -> crate::Result { + let mut entries = Vec::with_capacity(updates.len()); + for (field_name, update) in updates { + let Some(col_idx) = schema.columns.iter().position(|c| c.name == *field_name) else { + continue; + }; + let value = match update { + UpdateValue::Literal(bytes) => { + let value = nodedb_types::value_from_msgpack(bytes).map_err(|e| { + crate::Error::Internal { + detail: format!( + "failed to decode update value for field '{field_name}': {e}" + ), + } + })?; + AssignedValue::Literal(value) + } + UpdateValue::Expr(expr) => AssignedValue::Expr(expr.clone()), + }; + entries.push((col_idx, value)); + } + Ok(Self { entries }) + } + + /// Build the post-image of `row`, a schema-ordered pre-image. + /// + /// `Err` when an expression fails to evaluate, such as a division by + /// zero, or when a cell does not meet its declared column, such as a + /// value past a `SMALLINT` width (SQLSTATE 22003). + pub(in crate::data::executor) fn apply( + &self, + schema: &ColumnarSchema, + row: Vec, + ) -> crate::Result> { + // Every value is computed from the pre-image before any cell + // changes. The evaluation context is built on the first expression. + let mut context: Option = None; + let mut values = Vec::with_capacity(self.entries.len()); + for (col_idx, assigned) in &self.entries { + let value = match assigned { + AssignedValue::Literal(value) => value.clone(), + AssignedValue::Expr(expr) => { + let context = match context { + Some(ref context) => context, + None => { + context.insert(row_to_projected_value(&row, schema, &[], &[], false)?) + } + }; + expr.eval(context)? + } + }; + values.push((*col_idx, value)); + } + let mut new_row = row; + for (col_idx, value) in values { + let cells = new_row.len(); + let Some(cell) = new_row.get_mut(col_idx) else { + return Err(crate::Error::Internal { + detail: format!("columnar UPDATE: row holds {cells} cells, no cell {col_idx}"), + }); + }; + *cell = value; + } + coerce_columnar_row(schema, &mut new_row)?; + Ok(new_row) + } +} diff --git a/nodedb/src/data/executor/handlers/columnar_resolve.rs b/nodedb/src/data/executor/handlers/columnar_resolve.rs index 7ee552d42..9a501415e 100644 --- a/nodedb/src/data/executor/handlers/columnar_resolve.rs +++ b/nodedb/src/data/executor/handlers/columnar_resolve.rs @@ -1,22 +1,32 @@ // SPDX-License-Identifier: BUSL-1.1 -//! Shared row-selection and assignment logic for a columnar UPDATE/DELETE: -//! which memtable rows match the WHERE filters, and — for an UPDATE — what -//! their post-image is once the assignments are applied. +//! Shared row selection for a columnar UPDATE/DELETE: which current rows +//! match the WHERE filters, and for an UPDATE, what each match becomes once +//! the assignments apply. +//! +//! A current row is a row the delete bitmaps do not tombstone and the +//! primary-key index binds. It lives in a flushed segment or in the +//! memtable. Flushed rows read through the shared flushed-segment reader. //! //! `execute_columnar_update` / `execute_columnar_delete` -//! (`columnar_mutation.rs`) call this to select the rows they then mutate. -//! `execute_columnar_resolve_dml` (`columnar_resolve_dml.rs`, backing -//! `ColumnarOp::ResolveDml`) calls the same functions to report the identical -//! selection to the Control Plane without mutating anything. One -//! implementation of "which rows match and what do they become" decides both -//! what a predicate DML writes and what it is reported to write — so the two -//! can never drift into deciding the write policy against different images. - -use nodedb_types::Value; +//! (`columnar_mutation.rs`) apply the selection. +//! `execute_columnar_resolve_dml` (`columnar_resolve_dml.rs`) reports the +//! same selection to the Control Plane without applying it. The transaction +//! staging path (`stage_columnar_dml.rs`) reads the same current rows. One +//! selection decides what a predicate DML writes and what it reports. + +use nodedb_columnar::MutationEngine; +use nodedb_columnar::pk_index::RowLocation; +use nodedb_physical::physical_plan::UpdateValue; use nodedb_types::columnar::ColumnarSchema; +use nodedb_types::{Surrogate, Value}; use crate::bridge::scan_filter::ScanFilter; +use crate::data::executor::core_loop::CoreLoop; +use crate::data::executor::handlers::columnar_assignments::ColumnarAssignments; +use crate::data::executor::handlers::columnar_mutation_apply::{ + ColumnarEngineKey, flushed_row_surrogate, +}; use crate::data::executor::handlers::columnar_read::filter::row_matches_filters; use crate::data::executor::handlers::rls_write_gate::admit_columnar_row; @@ -36,96 +46,242 @@ pub(in crate::data::executor) fn require_pk_column_index( }) } -/// Bundled arguments for [`resolve_update_rows`]. +/// One current row of a columnar collection. +pub(in crate::data::executor) struct CurrentColumnarRow { + /// The cross-engine surrogate the row carries, when one was recorded. + pub surrogate: Option, + /// The row's cells in schema order. + pub values: Vec, +} + +/// Bundled arguments for [`CoreLoop::resolve_columnar_update_rows`]. pub(in crate::data::executor) struct ResolveUpdateRowsParams<'a> { - pub engine: &'a nodedb_columnar::MutationEngine, + pub key: &'a ColumnarEngineKey, pub schema: &'a ColumnarSchema, pub pk_col_idx: usize, pub filter_predicates: &'a [ScanFilter], - pub updates: &'a [(String, Vec)], + pub updates: &'a [(String, UpdateValue)], pub rls_write_check: &'a nodedb_types::RlsWriteCheck, pub tid: u64, pub collection: &'a str, } -/// Match memtable rows against `filter_predicates`, apply `updates` to build -/// each match's post-image, and decide every post-image against -/// `rls_write_check` — fail-fast: the first rejected row is the whole -/// statement's error, before the remaining rows are even resolved. -/// -/// Mutates nothing. Returns `(old_primary_key, post_image)` pairs in match -/// order. `old_primary_key` is the row's PK column value BEFORE `updates` is -/// applied — the value that identifies the row to remove even when the -/// update assigns the PK column a new value, exactly as -/// `execute_columnar_update` extracts it today. -pub(in crate::data::executor) fn resolve_update_rows( - params: ResolveUpdateRowsParams<'_>, -) -> crate::Result)>> { - let ResolveUpdateRowsParams { - engine, - schema, - pk_col_idx, - filter_predicates, - updates, - rls_write_check, - tid, - collection, - } = params; - let mut resolved = Vec::new(); - for row in engine.scan_memtable_rows() { - if !filter_predicates.is_empty() { - match row_matches_filters(&row, schema, filter_predicates) { - Ok(true) => {} - Ok(false) => continue, - Err(e) => return Err(crate::Error::from(e)), +/// Bundled arguments for [`CoreLoop::resolve_columnar_delete_rows`]. +pub(in crate::data::executor) struct ResolveDeleteRowsParams<'a> { + pub key: &'a ColumnarEngineKey, + pub schema: &'a ColumnarSchema, + pub pk_col_idx: usize, + pub filter_predicates: &'a [ScanFilter], + pub rls_write_check: &'a nodedb_types::RlsWriteCheck, + pub tid: u64, + pub collection: &'a str, +} + +impl CoreLoop { + /// Visit every current row of the collection at `key`: each flushed + /// segment in segment order, then the memtable. + /// + /// A row its segment's delete bitmap marks is not visited. A live row the + /// primary-key index does not bind is a superseded version, which only a + /// bitemporal collection keeps. It is not visited. In any other collection + /// an unbound live row is an `Internal` error. + /// + /// `Err` when a segment or a memtable cell does not read, or when `visit` + /// fails. The shared flushed reader files the corruption report, with + /// `site` naming the read path. + pub(in crate::data::executor) fn for_each_current_columnar_row( + &self, + key: &ColumnarEngineKey, + site: &'static str, + mut visit: impl FnMut(CurrentColumnarRow) -> crate::Result<()>, + ) -> crate::Result<()> { + let Some(engine) = self.columnar_engines.get(key) else { + return Ok(()); + }; + let schema = engine.schema(); + let collection = key.2.as_str(); + + let segments = self + .columnar_flushed_segments + .get(key) + .map(Vec::as_slice) + .unwrap_or_default(); + for (seg_idx, seg_bytes) in segments.iter().enumerate() { + // Segment ids are 1-based. Id 0 names the memtable. + let segment_id = seg_idx as u64 + 1; + let segment = self.decode_flushed_segment( + collection, + segment_id, + seg_bytes, + schema.columns.len(), + site, + )?; + let deletes = engine.delete_bitmap(segment_id); + for row_idx in 0..segment.row_count() { + let location = RowLocation { + segment_id, + row_index: row_index_u32(collection, segment_id, row_idx)?, + }; + if deletes.is_some_and(|bm| bm.is_deleted(location.row_index)) { + continue; + } + let values = segment.row(schema, row_idx)?; + if !row_is_bound(engine, &values, location, collection)? { + continue; + } + let surrogate = + flushed_row_surrogate(&self.columnar_flushed_surrogates, key, location); + visit(CurrentColumnarRow { surrogate, values })?; } } - let old_pk = row[pk_col_idx].clone(); - let mut new_row = row; - for (field_name, value_bytes) in updates { - if let Some(col_idx) = schema.columns.iter().position(|c| c.name == *field_name) { - let typed_val = nodedb_types::value_from_msgpack(value_bytes).map_err(|e| { - crate::Error::Internal { - detail: format!( - "failed to decode update value for field '{field_name}': {e}" - ), - } - })?; - new_row[col_idx] = typed_val; + let memtable_segment = engine.memtable_segment_id(); + let surrogates = engine.memtable_surrogates(); + for row_idx in 0..engine.memtable().row_count() { + // `None` is a row the memtable delete bitmap marks. + let Some(values) = engine.get_memtable_row(row_idx)? else { + continue; + }; + let location = RowLocation { + segment_id: memtable_segment, + row_index: row_index_u32(collection, memtable_segment, row_idx)?, + }; + if !row_is_bound(engine, &values, location, collection)? { + continue; } + let surrogate = surrogates.get(row_idx).copied().flatten(); + visit(CurrentColumnarRow { surrogate, values })?; } + Ok(()) + } + + /// Match current rows against `filter_predicates`, apply `updates` to + /// build each match's post-image, and decide every post-image against + /// `rls_write_check`. An expression assignment evaluates against the + /// matched row's pre-image. Each post-image meets the declared column + /// rule. The first row that fails any step is the statement's error. + /// + /// Mutates nothing. Returns `(old_primary_key, post_image)` pairs in match + /// order. `old_primary_key` is the row's PK value before `updates` + /// applies. It identifies the row to remove when the update assigns the + /// PK column a new value. + pub(in crate::data::executor) fn resolve_columnar_update_rows( + &self, + params: ResolveUpdateRowsParams<'_>, + ) -> crate::Result)>> { + let ResolveUpdateRowsParams { + key, + schema, + pk_col_idx, + filter_predicates, + updates, + rls_write_check, + tid, + collection, + } = params; + let assignments = ColumnarAssignments::bind(schema, updates)?; + let mut resolved = Vec::new(); + self.for_each_current_columnar_row(key, "columnar_dml_resolve", |row| { + let row = row.values; + if !row_matches(&row, schema, filter_predicates)? { + return Ok(()); + } + let old_pk = pk_value(&row, pk_col_idx, collection)?; + let new_row = assignments.apply(schema, row)?; + admit_columnar_row(rls_write_check, &new_row, schema, tid, collection)?; + resolved.push((old_pk, new_row)); + Ok(()) + })?; + Ok(resolved) + } - admit_columnar_row(rls_write_check, &new_row, schema, tid, collection)?; - resolved.push((old_pk, new_row)); + /// Match current rows against `filter_predicates` and decide each matched + /// row's pre-image, the image a DELETE removes, against + /// `rls_write_check`. Mutates nothing. Returns matched primary-key values + /// in match order. + pub(in crate::data::executor) fn resolve_columnar_delete_rows( + &self, + params: ResolveDeleteRowsParams<'_>, + ) -> crate::Result> { + let ResolveDeleteRowsParams { + key, + schema, + pk_col_idx, + filter_predicates, + rls_write_check, + tid, + collection, + } = params; + let mut pks = Vec::new(); + self.for_each_current_columnar_row(key, "columnar_dml_resolve", |row| { + let row = row.values; + if !row_matches(&row, schema, filter_predicates)? { + return Ok(()); + } + admit_columnar_row(rls_write_check, &row, schema, tid, collection)?; + pks.push(pk_value(&row, pk_col_idx, collection)?); + Ok(()) + })?; + Ok(pks) } - Ok(resolved) } -/// Match memtable rows against `filter_predicates` and decide each matched -/// row's pre-image — the image a DELETE removes — against -/// `rls_write_check`. Mutates nothing. Returns matched primary-key values in -/// match order. -pub(in crate::data::executor) fn resolve_delete_rows( - engine: &nodedb_columnar::MutationEngine, +/// Whether `row` passes every WHERE predicate. No predicate passes. +fn row_matches( + row: &[Value], schema: &ColumnarSchema, - pk_col_idx: usize, filter_predicates: &[ScanFilter], - rls_write_check: &nodedb_types::RlsWriteCheck, - tid: u64, +) -> crate::Result { + if filter_predicates.is_empty() { + return Ok(true); + } + row_matches_filters(row, schema, filter_predicates).map_err(crate::Error::from) +} + +/// The primary-key cell of `row`. +fn pk_value(row: &[Value], pk_col_idx: usize, collection: &str) -> crate::Result { + row.get(pk_col_idx) + .cloned() + .ok_or_else(|| crate::Error::Internal { + detail: format!( + "columnar '{collection}': row holds {} cells, no primary-key cell {pk_col_idx}", + row.len() + ), + }) +} + +/// Whether the primary-key index binds `values` to `location`. +/// +/// `Ok(false)` for a superseded bitemporal version. `Err` for an unbound live +/// row of a collection that is not bitemporal: the index lost a binding. +fn row_is_bound( + engine: &MutationEngine, + values: &[Value], + location: RowLocation, collection: &str, -) -> crate::Result> { - let mut pks = Vec::new(); - for row in engine.scan_memtable_rows() { - if !filter_predicates.is_empty() { - match row_matches_filters(&row, schema, filter_predicates) { - Ok(true) => {} - Ok(false) => continue, - Err(e) => return Err(crate::Error::from(e)), - } - } - admit_columnar_row(rls_write_check, &row, schema, tid, collection)?; - pks.push(row[pk_col_idx].clone()); +) -> crate::Result { + let pk_bytes = engine.encode_pk_from_row(values)?; + if engine.pk_index().get(&pk_bytes) == Some(&location) { + return Ok(true); + } + if engine.schema().is_bitemporal() { + return Ok(false); } - Ok(pks) + Err(crate::Error::Internal { + detail: format!( + "columnar '{collection}': live row {} of segment {} is not bound by the \ + primary-key index", + location.row_index, location.segment_id + ), + }) +} + +/// `row_idx` as the `u32` row index a `RowLocation` carries. +fn row_index_u32(collection: &str, segment_id: u64, row_idx: usize) -> crate::Result { + u32::try_from(row_idx).map_err(|_| crate::Error::Internal { + detail: format!( + "columnar '{collection}': row {row_idx} of segment {segment_id} is past the \ + u32 row index range" + ), + }) } diff --git a/nodedb/src/data/executor/handlers/columnar_resolve_dml.rs b/nodedb/src/data/executor/handlers/columnar_resolve_dml.rs index 52e205bc0..a340153a9 100644 --- a/nodedb/src/data/executor/handlers/columnar_resolve_dml.rs +++ b/nodedb/src/data/executor/handlers/columnar_resolve_dml.rs @@ -17,13 +17,14 @@ //! statement, exactly as `execute_columnar_update` / `execute_columnar_delete` //! do — the caller never sees a partial resolved set. +use nodedb_physical::physical_plan::UpdateValue; use tracing::debug; use crate::bridge::envelope::{ErrorCode, Response}; -use crate::bridge::scan_filter::ScanFilter; +use crate::bridge::scan_filter::decode_scan_filters; use crate::data::executor::core_loop::CoreLoop; use crate::data::executor::handlers::columnar_resolve::{ - ResolveUpdateRowsParams, require_pk_column_index, resolve_delete_rows, resolve_update_rows, + ResolveDeleteRowsParams, ResolveUpdateRowsParams, require_pk_column_index, }; use crate::data::executor::response_codec; use crate::data::executor::task::ExecutionTask; @@ -37,7 +38,7 @@ impl CoreLoop { task: &ExecutionTask, collection: &str, filters: &[u8], - updates: &[(String, Vec)], + updates: &[(String, UpdateValue)], is_update: bool, rls_write_check: &nodedb_types::RlsWriteCheck, ) -> Response { @@ -67,16 +68,17 @@ impl CoreLoop { Err(e) => return self.response_error(task, e), }; - let filter_predicates: Vec = if !filters.is_empty() { - zerompk::from_msgpack(filters).unwrap_or_default() - } else { - Vec::new() + // A filter that does not decode refuses the statement. Read as no + // filter, it would resolve every row. + let filter_predicates = match decode_scan_filters(filters, "columnar DML filter") { + Ok(filters) => filters, + Err(e) => return self.response_error(task, e), }; let tid = task.request.tenant_id.as_u64(); let payload = if is_update { - let rows = match resolve_update_rows(ResolveUpdateRowsParams { - engine, + let rows = match self.resolve_columnar_update_rows(ResolveUpdateRowsParams { + key: &key, schema: &schema, pk_col_idx, filter_predicates: &filter_predicates, @@ -90,15 +92,15 @@ impl CoreLoop { }; response_codec::encode(&rows) } else { - let pks = match resolve_delete_rows( - engine, - &schema, + let pks = match self.resolve_columnar_delete_rows(ResolveDeleteRowsParams { + key: &key, + schema: &schema, pk_col_idx, - &filter_predicates, + filter_predicates: &filter_predicates, rls_write_check, tid, collection, - ) { + }) { Ok(pks) => pks, Err(e) => return self.response_error(task, e), }; diff --git a/nodedb/src/data/executor/handlers/transaction/resolve/entry.rs b/nodedb/src/data/executor/handlers/transaction/resolve/entry.rs index b09338bd3..241f5a608 100644 --- a/nodedb/src/data/executor/handlers/transaction/resolve/entry.rs +++ b/nodedb/src/data/executor/handlers/transaction/resolve/entry.rs @@ -2872,6 +2872,7 @@ mod tests { .map(|engine| { engine .scan_memtable_rows_with_surrogates() + .map(|scanned| scanned.expect("read")) .filter_map(|(s, row)| s.map(|s| (s.as_u32(), row))) .collect() }) @@ -2988,8 +2989,10 @@ mod tests { filters: pk_filter("a"), updates: vec![( "id".to_string(), - nodedb_types::value_to_msgpack(&nodedb_types::Value::String("z".into())) - .expect("encode assignment"), + UpdateValue::Literal( + nodedb_types::value_to_msgpack(&nodedb_types::Value::String("z".into())) + .expect("encode assignment"), + ), )], rls_write_check: nodedb_types::RlsWriteCheck::NoPolicyApplies, }); @@ -3195,7 +3198,9 @@ mod tests { let (mut src, _src_dir) = make_core(); let task = make_task(); let txn = TxnId::new(44); - let line = r"cpu\,load,host\ name=west\,1 count=18446744073709551615u 1700000000000000001"; + // The unsigned field is the largest an `Int64` column holds: a larger + // one refuses the line at ingest rather than wrapping negative. + let line = r"cpu\,load,host\ name=west\,1 count=9223372036854775807u 1700000000000000001"; let payload = zerompk::to_msgpack_vec(&vec![line.to_string()]).expect("encode canonical ILP lines"); let plan = PhysicalPlan::Timeseries(TimeseriesOp::Ingest { diff --git a/nodedb/src/data/executor/handlers/transaction/stage_write/stage_columnar_dml.rs b/nodedb/src/data/executor/handlers/transaction/stage_write/stage_columnar_dml.rs index ad66f9d13..5d50ba192 100644 --- a/nodedb/src/data/executor/handlers/transaction/stage_write/stage_columnar_dml.rs +++ b/nodedb/src/data/executor/handlers/transaction/stage_write/stage_columnar_dml.rs @@ -39,10 +39,12 @@ //! transaction's redo record (`resolve::columnar_image`). The first statement //! that stages a base row records that row's primary key //! (`stage_columnar_base_key`), so the redo names the base row a key-changing -//! UPDATE or a DELETE removes. The staged set is resolved from the live -//! memtable (plus overlay), the scope the autocommit handlers -//! (`execute_columnar_delete` / `execute_columnar_update`) match against. +//! UPDATE or a DELETE removes. The staged set is resolved from the current +//! rows, flushed and in the memtable (plus overlay), the scope the +//! autocommit handlers (`execute_columnar_delete` / `execute_columnar_update`) +//! match against. +use nodedb_physical::physical_plan::UpdateValue; use nodedb_types::columnar::ColumnarSchema; use nodedb_types::value::Value; use nodedb_types::{RowIdentity, Surrogate, value_to_pk_string}; @@ -50,8 +52,10 @@ use nodedb_types::{RowIdentity, Surrogate, value_to_pk_string}; use crate::bridge::envelope::{ErrorCode, Response}; use crate::bridge::scan_filter::ScanFilter; use crate::data::executor::core_loop::CoreLoop; +use crate::data::executor::handlers::columnar_assignments::ColumnarAssignments; use crate::data::executor::handlers::columnar_read::convert::row_to_projected_value; use crate::data::executor::handlers::columnar_read::filter::row_matches_filters; +use crate::data::executor::handlers::columnar_resolve::CurrentColumnarRow; use crate::data::executor::handlers::transaction::overlay::{ ColumnarMatchedRow, ColumnarOverlayMergeParams, }; @@ -95,9 +99,9 @@ pub(in crate::data::executor) struct StageColumnarUpdateParams<'a> { pub txn_id: TxnId, pub collection: &'a str, pub filter_bytes: &'a [u8], - /// Field assignments: `(column_name, msgpack_value_bytes)`, the same shape - /// `execute_columnar_update` applies on the durable path. - pub updates: &'a [(String, Vec)], + /// Field assignments, the same shape `execute_columnar_update` applies + /// on the durable path. + pub updates: &'a [(String, UpdateValue)], /// Compiled row-level-security WRITE predicate carried by the plan, /// decided against each row's post-image once the assignments are applied. pub rls_write_check: &'a nodedb_types::RlsWriteCheck, @@ -214,15 +218,21 @@ impl CoreLoop { // it only exists once the assignments are applied. A refusal partway // through would leave the rows ahead of it staged and visible to this // transaction's own reads. + // + // An expression assignment evaluates against each row's pre-image, and + // each post-image meets the declared column rule, as on the durable + // path. + let assignments = match ColumnarAssignments::bind(&schema, updates) { + Ok(a) => a, + Err(e) => return self.response_error(task, e), + }; let affected = affected_rows.len(); let mut new_rows: Vec<(u32, Vec)> = Vec::with_capacity(affected); let mut base_rows: Vec<(u32, Vec)> = Vec::with_capacity(affected); for (surrogate, row) in affected_rows { - match apply_columnar_updates(&schema, row.clone(), updates) { + match assignments.apply(&schema, row.clone()) { Ok(r) => new_rows.push((surrogate, r)), - Err(detail) => { - return self.response_error(task, ErrorCode::Internal { detail }); - } + Err(e) => return self.response_error(task, e), } base_rows.push((surrogate, row)); } @@ -270,15 +280,16 @@ impl CoreLoop { } /// Resolve the CURRENT in-transaction matching set for a columnar - /// predicate DELETE/UPDATE: committed memtable rows matching the WHERE + /// predicate DELETE/UPDATE: committed current rows matching the WHERE /// predicate, folded with this transaction's own staged overlay /// (tombstones dropped, staged puts/inserts surfaced) via /// [`Self::merge_overlay_into_columnar_scan`]. Returns each affected row's /// `(surrogate, schema-ordered values)`. /// /// Mirrors `execute_columnar_delete` / `execute_columnar_update`: the base - /// set is the live memtable (the same scope the durable handlers mutate), - /// so the staged affected set matches exactly what COMMIT replay applies. + /// set is every current row, flushed and in the memtable (the same scope + /// the durable handlers mutate), so the staged affected set matches + /// exactly what COMMIT replay applies. pub(super) fn columnar_txn_matching_rows( &self, task: &ExecutionTask, @@ -311,32 +322,30 @@ impl CoreLoop { collection.to_string(), ); - // BASE: live memtable rows matching the predicate, carried as the - // shared `ColumnarMatchedRow` tuple the overlay merge consumes. A - // missing engine means the only affected rows are overlay-only staged - // inserts, which the merge appends below. + // BASE: current rows, flushed and in the memtable, matching the + // predicate, carried as the shared `ColumnarMatchedRow` tuple the + // overlay merge consumes. A missing engine visits no row, so the only + // affected rows are overlay-only staged inserts, which the merge + // appends below. let mut matched: Vec = Vec::new(); - if let Some(engine) = self.columnar_engines.get(&coll_key) { - for (surrogate, row) in engine.scan_memtable_rows_with_surrogates() { - if !filter_predicates.is_empty() { - match row_matches_filters(&row, &schema, &filter_predicates) { - Ok(true) => {} - Ok(false) => continue, - Err(e) => { - return Err(self.response_error(task, ErrorCode::from(e))); - } - } + let base = + self.for_each_current_columnar_row(&coll_key, "columnar_txn_dml_resolve", |row| { + let CurrentColumnarRow { surrogate, values } = row; + if !filter_predicates.is_empty() + && !row_matches_filters(&values, &schema, &filter_predicates)? + { + return Ok(()); } - // No computed columns on this path (`&[]` below), so this - // cannot raise an evaluation error today. It is handled like - // every other `row_to_projected_value` caller instead of - // assuming that invariant with an `unwrap`. - let obj = match row_to_projected_value(&row, &schema, &[], &[], false) { - Ok(v) => v, - Err(e) => return Err(self.response_error(task, e)), - }; - matched.push((surrogate, row, obj)); - } + // No computed columns on this path (`&[]` below), so this cannot + // raise an evaluation error today. It is handled like every other + // `row_to_projected_value` caller instead of assuming that + // invariant with an `unwrap`. + let obj = row_to_projected_value(&values, &schema, &[], &[], false)?; + matched.push((surrogate, values, obj)); + Ok(()) + }); + if let Err(e) = base { + return Err(self.response_error(task, e)); } // Fold the transaction's own staged writes into the base set: drops @@ -412,33 +421,7 @@ impl CoreLoop { ) -> Response { match response_codec::encode_json_as_msgpack(&serde_json::json!({ "affected": affected })) { Ok(payload) => self.response_with_payload(task, payload), - Err(e) => self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ), + Err(e) => self.response_error(task, ErrorCode::from(e)), } } } - -/// Apply the columnar UPDATE SET-list to one schema-ordered row, mirroring -/// `execute_columnar_update`'s per-field application: each `(field, bytes)` -/// pair overwrites the row's value at the field's schema column index, with -/// the value decoded from MessagePack. Unknown fields are ignored (same as the -/// durable path). Returns the new row, or a decode-error detail string. -fn apply_columnar_updates( - schema: &ColumnarSchema, - mut row: Vec, - updates: &[(String, Vec)], -) -> Result, String> { - for (field_name, value_bytes) in updates { - let Some(col_idx) = schema.columns.iter().position(|c| c.name == *field_name) else { - continue; - }; - let typed_val = nodedb_types::value_from_msgpack(value_bytes) - .map_err(|e| format!("failed to decode update value for field '{field_name}': {e}"))?; - row[col_idx] = typed_val; - } - Ok(row) -} diff --git a/nodedb/src/data/executor/wal_replay_columnar_dml.rs b/nodedb/src/data/executor/wal_replay_columnar_dml.rs index 1267c8000..66e43be72 100644 --- a/nodedb/src/data/executor/wal_replay_columnar_dml.rs +++ b/nodedb/src/data/executor/wal_replay_columnar_dml.rs @@ -109,6 +109,23 @@ impl CoreLoop { else { return Some(0); }; + // An assignment that does not decode rejects the record, the same way + // a malformed resolved row does below. + let updates = match crate::wal::decode_columnar_dml_updates(&record.updates) { + Ok(updates) => updates, + Err(e) => { + self.replay_record_rejected( + "columnar", + record_lsn, + Some(Box::new(crate::bridge::envelope::ErrorCode::from(e))), + &format!( + "columnar predicate DML on '{}': malformed assignment", + record.collection + ), + ); + return Some(0); + } + }; // The task carries the real predicate even though today's handlers read // only `task.request.{database_id, tenant_id}`. A placeholder plan would @@ -126,7 +143,7 @@ impl CoreLoop { record.collection.clone(), ), filters: record.filters.clone(), - updates: record.updates.clone(), + updates: updates.clone(), rls_write_check: RlsWriteCheck::already_decided_elsewhere(), }) } else { @@ -164,7 +181,7 @@ impl CoreLoop { &task, &record.collection, &record.filters, - &record.updates, + &updates, &replay_check, recording.then_some(&mut undo), ) @@ -398,7 +415,7 @@ mod tests { use crate::control::server::wal_dispatch::wal_append_if_write; use crate::types::{DatabaseId, TenantId, VShardId}; use crate::wal::manager::WalManager; - use nodedb_physical::physical_plan::{ColumnarInsertIntent, ColumnarOp}; + use nodedb_physical::physical_plan::{ColumnarInsertIntent, ColumnarOp, UpdateValue}; use nodedb_query::scan_filter::{FilterOp, ScanFilter}; use nodedb_types::{QualifiedCollection, RlsWriteCheck, Value}; use nodedb_wal::TombstoneSet; @@ -528,6 +545,7 @@ mod tests { engine .scan_memtable_rows() .map(|row| { + let row = row.expect("read"); let id = match &row[id_idx] { Value::Integer(n) => *n, other => panic!("expected integer id, got {other:?}"), @@ -574,11 +592,13 @@ mod tests { filters: eq_filter_bytes("id", Value::Integer(1)), updates: vec![( "v".to_string(), - // `execute_columnar_update` decodes each update value with the + // `execute_columnar_update` decodes each literal with the // plain reader (`value_from_msgpack`), matching the planner's // `sql_value_to_msgpack` output — not the tagged `Value` enum // encoding `zerompk::to_msgpack_vec(&Value)` would emit. - nodedb_types::value_to_msgpack(&Value::Integer(999)).expect("encode"), + UpdateValue::Literal( + nodedb_types::value_to_msgpack(&Value::Integer(999)).expect("encode"), + ), )], rls_write_check: RlsWriteCheck::already_decided_elsewhere(), }); @@ -633,7 +653,9 @@ mod tests { // (`sql_value_to_msgpack`), not the tagged `Value` enum encoding. updates: vec![( "v".to_string(), - nodedb_types::value_to_msgpack(&Value::Integer(999)).expect("encode"), + UpdateValue::Literal( + nodedb_types::value_to_msgpack(&Value::Integer(999)).expect("encode"), + ), )], rls_write_check: RlsWriteCheck::already_decided_elsewhere(), }); @@ -654,4 +676,36 @@ mod tests { a non-idempotent double-apply would instead duplicate the row)" ); } + + #[test] + fn computed_update_replays_the_expression_against_each_row() { + use nodedb_query::expr::{BinaryOp, SqlExpr}; + + let insert = insert_plan(vec![row(1, 10), row(2, 20)]); + // `UPDATE t SET v = v + 5`: the record carries the expression, and + // replay evaluates it against each matched row's pre-image. + let update = PhysicalPlan::Columnar(ColumnarOp::Update { + collection: QualifiedCollection::new(DatabaseId::DEFAULT, COLLECTION), + filters: Vec::new(), + updates: vec![( + "v".to_string(), + UpdateValue::Expr(SqlExpr::BinaryOp { + left: Box::new(SqlExpr::Column("v".to_string())), + op: BinaryOp::Add, + right: Box::new(SqlExpr::Literal(Value::Integer(5))), + }), + )], + rls_write_check: RlsWriteCheck::already_decided_elsewhere(), + }); + + let records = append_via_autocommit(&[insert, update]); + + let mut h = make_core(); + h.core + .replay_timeseries_wal(&records, 1, &TombstoneSet::new()); + + let mut rows = scan_ids(&mut h.core); + rows.sort(); + assert_eq!(rows, vec![(1, 15), (2, 25)]); + } } diff --git a/nodedb/src/wal/columnar_dml_updates.rs b/nodedb/src/wal/columnar_dml_updates.rs new file mode 100644 index 000000000..508fd88f7 --- /dev/null +++ b/nodedb/src/wal/columnar_dml_updates.rs @@ -0,0 +1,83 @@ +// SPDX-License-Identifier: BUSL-1.1 + +//! The SET-list a columnar predicate-UPDATE WAL record carries. +//! +//! `ColumnarDmlWalRecord` lives in `nodedb-types`, which cannot name +//! `UpdateValue`. The record stores each assignment's value as the +//! MessagePack encoding of its `UpdateValue`. A literal and an expression +//! both survive the log, so replay re-executes the same computed UPDATE. + +use nodedb_physical::physical_plan::UpdateValue; + +/// Encode `updates` into the record's `(column, value bytes)` shape. +pub(crate) fn encode_columnar_dml_updates( + updates: &[(String, UpdateValue)], +) -> crate::Result)>> { + updates + .iter() + .map(|(field, value)| { + let bytes = + zerompk::to_msgpack_vec(value).map_err(|e| crate::Error::Serialization { + format: "msgpack".into(), + detail: format!("wal columnar dml: assignment to '{field}': {e}"), + })?; + Ok((field.clone(), bytes)) + }) + .collect() +} + +/// Decode the record's `(column, value bytes)` shape. A value that does not +/// decode as an `UpdateValue` is a `Serialization` error. +pub(crate) fn decode_columnar_dml_updates( + updates: &[(String, Vec)], +) -> crate::Result> { + updates + .iter() + .map(|(field, bytes)| { + let value = zerompk::from_msgpack::(bytes).map_err(|e| { + crate::Error::Serialization { + format: "msgpack".into(), + detail: format!("wal columnar dml: assignment to '{field}': {e}"), + } + })?; + Ok((field.clone(), value)) + }) + .collect() +} + +#[cfg(test)] +mod tests { + use super::*; + use nodedb_query::expr::{BinaryOp, SqlExpr}; + + #[test] + fn literal_and_expression_round_trip() { + let updates = vec![ + ( + "label".to_string(), + UpdateValue::Literal( + nodedb_types::value_to_msgpack(&nodedb_types::Value::String("x".into())) + .expect("encode literal"), + ), + ), + ( + "v".to_string(), + UpdateValue::Expr(SqlExpr::BinaryOp { + left: Box::new(SqlExpr::Column("v".to_string())), + op: BinaryOp::Add, + right: Box::new(SqlExpr::Literal(nodedb_types::Value::Integer(1))), + }), + ), + ]; + let encoded = encode_columnar_dml_updates(&updates).expect("encode"); + let decoded = decode_columnar_dml_updates(&encoded).expect("decode"); + assert_eq!(decoded, updates); + } + + #[test] + fn undecodable_value_is_an_error() { + let err = decode_columnar_dml_updates(&[("v".to_string(), vec![0xc1])]) + .expect_err("0xc1 is not a MessagePack value"); + assert!(matches!(err, crate::Error::Serialization { .. }), "{err:?}"); + } +} diff --git a/nodedb/src/wal/mod.rs b/nodedb/src/wal/mod.rs index 62515d578..229e71169 100644 --- a/nodedb/src/wal/mod.rs +++ b/nodedb/src/wal/mod.rs @@ -3,6 +3,7 @@ pub mod archiver; pub mod audit_archive; pub mod audit_segment; +pub mod columnar_dml_updates; pub mod crdt_doc_payload; pub mod crdt_list_payload; pub mod crdt_payload; @@ -12,6 +13,7 @@ pub mod replay; pub mod timeseries_batch_payload; pub use audit_segment::AuditWalSegment; +pub(crate) use columnar_dml_updates::{decode_columnar_dml_updates, encode_columnar_dml_updates}; pub(crate) use crdt_doc_payload::CrdtDocOpWalRecord; pub(crate) use crdt_list_payload::CrdtListOpWalRecord; pub(crate) use crdt_payload::{ diff --git a/nodedb/tests/wire/cases/columnar_predicate_dml_flushed.rs b/nodedb/tests/wire/cases/columnar_predicate_dml_flushed.rs new file mode 100644 index 000000000..993c900ca --- /dev/null +++ b/nodedb/tests/wire/cases/columnar_predicate_dml_flushed.rs @@ -0,0 +1,224 @@ +// SPDX-License-Identifier: BUSL-1.1 + +//! Predicate `UPDATE` and `DELETE` on a columnar collection match rows in +//! flushed segments and rows in the memtable alike. +//! +//! The server flushes the memtable once it holds two rows. The seed writes +//! four rows in one statement, which flush to a segment, then one row that +//! stays in the memtable. Each predicate matches rows in both places. +//! +//! A computed `UPDATE` (`SET x = x + 10`) evaluates against each matched +//! row's pre-image. A result past a declared width refuses the statement, +//! and no row changes. + +use crate::harness::TestServer; +use tokio_postgres::SimpleQueryMessage; + +/// The flush threshold every test here starts the server with. +const FLUSH_THRESHOLD: usize = 2; + +/// Rows `1..=4` are flushed. Row `5` is in the memtable. +async fn seed(server: &TestServer, name: &str) { + server + .exec(&format!( + "CREATE COLLECTION {name} (id INT PRIMARY KEY, x INT) WITH (engine='columnar')" + )) + .await + .unwrap_or_else(|e| panic!("create {name}: {e}")); + server + .exec(&format!( + "INSERT INTO {name} (id, x) VALUES (1, 1), (2, 6), (3, 7), (4, 2)" + )) + .await + .unwrap_or_else(|e| panic!("seed flushed rows of {name}: {e}")); + server + .exec(&format!("INSERT INTO {name} (id, x) VALUES (5, 8)")) + .await + .unwrap_or_else(|e| panic!("seed memtable row of {name}: {e}")); +} + +/// The row count `sql` reports in its command tag. +async fn affected(server: &TestServer, sql: &str) -> u64 { + let messages = server + .client + .simple_query(sql) + .await + .unwrap_or_else(|e| panic!("{sql}: {e}")); + messages + .iter() + .find_map(|m| match m { + SimpleQueryMessage::CommandComplete(n) => Some(*n), + _ => None, + }) + .unwrap_or_else(|| panic!("{sql} reported no command tag")) +} + +/// `id` and `x` of every row joined by a tab, ordered by `id`. +async fn rows(server: &TestServer, name: &str) -> Vec { + server + .query_text_joined(&format!("SELECT id, x FROM {name} ORDER BY id")) + .await + .unwrap_or_else(|e| panic!("select {name}: {e}")) +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn a_predicate_delete_removes_flushed_and_memtable_rows() { + let server = TestServer::start_with_columnar_flush_threshold(FLUSH_THRESHOLD).await; + seed(&server, "cpd_delete").await; + + let count = affected(&server, "DELETE FROM cpd_delete WHERE x > 5").await; + + assert_eq!( + count, 3, + "rows 2 and 3 are flushed, row 5 is in the memtable" + ); + assert_eq!(rows(&server, "cpd_delete").await, vec!["1\t1", "4\t2"]); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn a_predicate_update_changes_flushed_and_memtable_rows() { + let server = TestServer::start_with_columnar_flush_threshold(FLUSH_THRESHOLD).await; + seed(&server, "cpd_update").await; + + let count = affected(&server, "UPDATE cpd_update SET x = 0 WHERE x > 5").await; + + assert_eq!( + count, 3, + "rows 2 and 3 are flushed, row 5 is in the memtable" + ); + assert_eq!( + rows(&server, "cpd_update").await, + vec!["1\t1", "2\t0", "3\t0", "4\t2", "5\t0"] + ); + // The updated flushed rows read once each: the originals are tombstoned. + assert_eq!( + affected(&server, "DELETE FROM cpd_update WHERE x = 0").await, + 3 + ); + assert_eq!(rows(&server, "cpd_update").await, vec!["1\t1", "4\t2"]); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn a_predicate_delete_in_a_transaction_removes_flushed_and_memtable_rows() { + let server = TestServer::start_with_columnar_flush_threshold(FLUSH_THRESHOLD).await; + seed(&server, "cpd_txn").await; + + server.exec("BEGIN").await.expect("begin"); + let count = affected(&server, "DELETE FROM cpd_txn WHERE x > 5").await; + assert_eq!( + count, 3, + "rows 2 and 3 are flushed, row 5 is in the memtable" + ); + server.exec("COMMIT").await.expect("commit"); + + assert_eq!(rows(&server, "cpd_txn").await, vec!["1\t1", "4\t2"]); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn a_computed_update_evaluates_each_flushed_and_memtable_row() { + let server = TestServer::start_with_columnar_flush_threshold(FLUSH_THRESHOLD).await; + seed(&server, "cpd_computed").await; + + let count = affected(&server, "UPDATE cpd_computed SET x = x + 10 WHERE x > 5").await; + + assert_eq!( + count, 3, + "rows 2 and 3 are flushed, row 5 is in the memtable" + ); + assert_eq!( + rows(&server, "cpd_computed").await, + vec!["1\t1", "2\t16", "3\t17", "4\t2", "5\t18"] + ); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn a_computed_update_in_a_transaction_reads_its_own_writes_and_commits() { + let server = TestServer::start_with_columnar_flush_threshold(FLUSH_THRESHOLD).await; + seed(&server, "cpd_computed_txn").await; + + server.exec("BEGIN").await.expect("begin"); + let count = affected(&server, "UPDATE cpd_computed_txn SET x = x * 2 WHERE x > 5").await; + assert_eq!(count, 3); + assert_eq!( + rows(&server, "cpd_computed_txn").await, + vec!["1\t1", "2\t12", "3\t14", "4\t2", "5\t16"], + "the transaction reads its own computed post-images" + ); + server.exec("COMMIT").await.expect("commit"); + + assert_eq!( + rows(&server, "cpd_computed_txn").await, + vec!["1\t1", "2\t12", "3\t14", "4\t2", "5\t16"] + ); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn a_rolled_back_computed_update_changes_no_row() { + let server = TestServer::start_with_columnar_flush_threshold(FLUSH_THRESHOLD).await; + seed(&server, "cpd_computed_rb").await; + + server.exec("BEGIN").await.expect("begin"); + affected(&server, "UPDATE cpd_computed_rb SET x = x - 1").await; + server.exec("ROLLBACK").await.expect("rollback"); + + assert_eq!( + rows(&server, "cpd_computed_rb").await, + vec!["1\t1", "2\t6", "3\t7", "4\t2", "5\t8"] + ); +} + +/// `x * 4096` fits `SMALLINT` for every row but row 5, the last row the +/// update visits. The statement is refused whole: no flushed row changes. +async fn assert_overflowing_computed_update_is_refused(server: &TestServer, name: &str) { + server + .expect_error(&format!("UPDATE {name} SET x = x * 4096"), "SQLSTATE 22003") + .await; +} + +/// Rows `1..=4` are flushed. Row `5` is in the memtable. `x` is `SMALLINT`. +async fn seed_smallint(server: &TestServer, name: &str) { + server + .exec(&format!( + "CREATE COLLECTION {name} (id INT PRIMARY KEY, x SMALLINT) WITH (engine='columnar')" + )) + .await + .unwrap_or_else(|e| panic!("create {name}: {e}")); + server + .exec(&format!( + "INSERT INTO {name} (id, x) VALUES (1, 1), (2, 6), (3, 7), (4, 2)" + )) + .await + .unwrap_or_else(|e| panic!("seed flushed rows of {name}: {e}")); + server + .exec(&format!("INSERT INTO {name} (id, x) VALUES (5, 8)")) + .await + .unwrap_or_else(|e| panic!("seed memtable row of {name}: {e}")); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn a_computed_update_past_a_declared_width_changes_no_row() { + let server = TestServer::start_with_columnar_flush_threshold(FLUSH_THRESHOLD).await; + seed_smallint(&server, "cpd_width").await; + + assert_overflowing_computed_update_is_refused(&server, "cpd_width").await; + + assert_eq!( + rows(&server, "cpd_width").await, + vec!["1\t1", "2\t6", "3\t7", "4\t2", "5\t8"] + ); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn a_computed_update_past_a_declared_width_in_a_transaction_changes_no_row() { + let server = TestServer::start_with_columnar_flush_threshold(FLUSH_THRESHOLD).await; + seed_smallint(&server, "cpd_width_txn").await; + + server.exec("BEGIN").await.expect("begin"); + assert_overflowing_computed_update_is_refused(&server, "cpd_width_txn").await; + server.exec("ROLLBACK").await.expect("rollback"); + + assert_eq!( + rows(&server, "cpd_width_txn").await, + vec!["1\t1", "2\t6", "3\t7", "4\t2", "5\t8"] + ); +} From c6b9f6f2cc4ae9afd88ceafd70086de3d73ec797 Mon Sep 17 00:00:00 2001 From: Farhan Syah Date: Thu, 8 Oct 2026 16:51:47 +0800 Subject: [PATCH 14/24] feat(fts): index each top-level string field in its own scope Every full-text table key gains a field component: a collection has a whole-document index plus one index per top-level string field. A document is indexed as (field, text) pairs through the WAL, sync and replication records, and a re-index retracts the fields the new body no longer fills. A segment that fails validation or is listed but missing is an error. Deletes are tracked per segment so merges drop dead postings. Bulk DELETE and TRUNCATE remove text in batched transactions, and a full TRUNCATE clears a collection in one purge. The writer and memtable split into module directories. --- nodedb-fts/src/backend/memory.rs | 455 ++++++------ nodedb-fts/src/backend/traits.rs | 80 ++- nodedb-fts/src/document_text.rs | 145 ++++ nodedb-fts/src/index/analyzer_config.rs | 32 +- nodedb-fts/src/index/error.rs | 28 + nodedb-fts/src/index/fieldnorm.rs | 16 +- nodedb-fts/src/index/stats.rs | 11 +- nodedb-fts/src/index/writer.rs | 489 ------------- nodedb-fts/src/index/writer/document.rs | 528 ++++++++++++++ nodedb-fts/src/index/writer/maintenance.rs | 349 ++++++++++ nodedb-fts/src/index/writer/mod.rs | 9 + nodedb-fts/src/index/writer/state.rs | 97 +++ nodedb-fts/src/lib.rs | 7 + nodedb-fts/src/lsm/compaction.rs | 86 ++- nodedb-fts/src/lsm/memtable.rs | 386 ----------- nodedb-fts/src/lsm/memtable/config.rs | 26 + nodedb-fts/src/lsm/memtable/mod.rs | 8 + nodedb-fts/src/lsm/memtable/scope_state.rs | 117 ++++ nodedb-fts/src/lsm/memtable/table.rs | 649 ++++++++++++++++++ nodedb-fts/src/lsm/merge.rs | 83 ++- nodedb-fts/src/lsm/mod.rs | 1 + nodedb-fts/src/lsm/parallel_build.rs | 73 +- nodedb-fts/src/lsm/query.rs | 264 +++++-- nodedb-fts/src/lsm/segment/error.rs | 5 + nodedb-fts/src/lsm/segment/reader.rs | 107 +-- nodedb-fts/src/lsm/segment/writer.rs | 20 +- nodedb-fts/src/lsm/segment_deletes.rs | 150 ++++ nodedb-fts/src/scope.rs | 98 +++ nodedb-types/src/sync/wire/fts.rs | 14 +- nodedb-wal/src/record/fts_spatial.rs | 77 ++- nodedb/src/control/server/sync/fts_handler.rs | 45 +- nodedb/src/control/server/sync/fts_session.rs | 18 +- .../src/control/server/wal_dispatch/core.rs | 7 +- .../src/control/server/wal_dispatch/text.rs | 10 +- .../decode/entry_columnar_family.rs | 4 +- .../wal_replication/decode_sync_engines.rs | 17 +- .../wal_replication/encode/columnar.rs | 6 +- .../encode/entry_columnar_family.rs | 4 +- .../wal_replication/types/replicated_write.rs | 8 +- .../src/data/executor/core_loop/accessors.rs | 14 +- nodedb/src/data/executor/fts_text.rs | 59 +- .../data/executor/handlers/bulk_dml/delete.rs | 93 ++- .../handlers/bulk_dml/delete_cascade.rs | 219 ++++-- .../executor/handlers/columnar_mutation.rs | 100 +-- .../src/data/executor/handlers/compact/fts.rs | 73 +- .../control/checkpoint_durable_lsn.rs | 2 +- nodedb/src/data/executor/handlers/fts_sync.rs | 85 ++- .../handlers/point/update_reindex_text.rs | 11 +- .../handlers/snapshot/restore/text.rs | 4 +- .../src/data/executor/handlers/text_index.rs | 99 +++ .../redo_apply/install_refusal_tests.rs | 2 +- .../handlers/transaction/resolve/text.rs | 4 +- .../transaction/undo/columnar_insert.rs | 7 +- nodedb/src/data/executor/handlers/truncate.rs | 160 +++-- nodedb/src/data/executor/wal_replay_fts.rs | 48 +- .../engine/sparse/fts_redb/backend/core.rs | 116 ++-- .../sparse/fts_redb/backend/doc_lengths.rs | 88 ++- .../engine/sparse/fts_redb/backend/meta.rs | 14 +- .../sparse/fts_redb/backend/postings.rs | 38 +- .../engine/sparse/fts_redb/backend/purge.rs | 294 ++++---- .../sparse/fts_redb/backend/segments.rs | 80 +-- .../engine/sparse/fts_redb/backend/shared.rs | 14 +- .../engine/sparse/fts_redb/backend/stats.rs | 46 +- nodedb/src/engine/sparse/fts_redb/keys.rs | 239 +++++++ nodedb/src/engine/sparse/fts_redb/mod.rs | 2 + nodedb/src/engine/sparse/fts_redb/scan.rs | 205 ++++++ nodedb/src/engine/sparse/fts_redb/tables.rs | 49 +- .../engine/sparse/inverted/batch_removal.rs | 206 ++++++ .../src/engine/sparse/inverted/compaction.rs | 22 +- nodedb/src/engine/sparse/inverted/core.rs | 70 +- .../engine/sparse/inverted/corpus_stats.rs | 102 ++- .../src/engine/sparse/inverted/doc_fields.rs | 84 +++ .../src/engine/sparse/inverted/doc_image.rs | 164 +++-- .../src/engine/sparse/inverted/doc_terms.rs | 203 +++--- nodedb/src/engine/sparse/inverted/document.rs | 275 ++++++++ nodedb/src/engine/sparse/inverted/errors.rs | 43 +- nodedb/src/engine/sparse/inverted/indexing.rs | 238 ++++--- nodedb/src/engine/sparse/inverted/mod.rs | 14 +- .../engine/sparse/inverted/rebuild_install.rs | 273 ++++++-- .../engine/sparse/inverted/rebuild_journal.rs | 21 +- .../sparse/inverted/rebuild_snapshot.rs | 178 +++-- nodedb/src/engine/sparse/inverted/removal.rs | 83 ++- nodedb/src/engine/sparse/inverted/search.rs | 87 ++- .../engine/sparse/inverted/test_support.rs | 19 + nodedb/src/event/wal_replay.rs | 11 +- nodedb/src/wal/manager/append.rs | 7 +- .../cases/collection_purge_persistent.rs | 7 +- .../tests/inproc/cases/fts_update_reindex.rs | 39 ++ nodedb/tests/inproc/cases/mod.rs | 1 + .../inproc/cases/reindex_field_scopes.rs | 66 ++ nodedb/tests/inproc/cases/sync_compat.rs | 48 ++ 91 files changed, 6439 insertions(+), 2516 deletions(-) create mode 100644 nodedb-fts/src/document_text.rs delete mode 100644 nodedb-fts/src/index/writer.rs create mode 100644 nodedb-fts/src/index/writer/document.rs create mode 100644 nodedb-fts/src/index/writer/maintenance.rs create mode 100644 nodedb-fts/src/index/writer/mod.rs create mode 100644 nodedb-fts/src/index/writer/state.rs delete mode 100644 nodedb-fts/src/lsm/memtable.rs create mode 100644 nodedb-fts/src/lsm/memtable/config.rs create mode 100644 nodedb-fts/src/lsm/memtable/mod.rs create mode 100644 nodedb-fts/src/lsm/memtable/scope_state.rs create mode 100644 nodedb-fts/src/lsm/memtable/table.rs create mode 100644 nodedb-fts/src/lsm/segment_deletes.rs create mode 100644 nodedb-fts/src/scope.rs create mode 100644 nodedb/src/data/executor/handlers/text_index.rs create mode 100644 nodedb/src/engine/sparse/fts_redb/keys.rs create mode 100644 nodedb/src/engine/sparse/fts_redb/scan.rs create mode 100644 nodedb/src/engine/sparse/inverted/batch_removal.rs create mode 100644 nodedb/src/engine/sparse/inverted/doc_fields.rs create mode 100644 nodedb/src/engine/sparse/inverted/document.rs create mode 100644 nodedb/src/engine/sparse/inverted/test_support.rs create mode 100644 nodedb/tests/inproc/cases/reindex_field_scopes.rs 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/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/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-types/src/sync/wire/fts.rs b/nodedb-types/src/sync/wire/fts.rs index 8c2ce897e..c17d3811d 100644 --- a/nodedb-types/src/sync/wire/fts.rs +++ b/nodedb-types/src/sync/wire/fts.rs @@ -2,7 +2,7 @@ //! FTS index/delete sync messages (client → server / server → client). //! -//! `FtsIndexMsg` carries one document's text content from a Lite client to +//! `FtsIndexMsg` carries one document's string fields from a Lite client to //! Origin for full-text indexing. `FtsDeleteMsg` removes a document from //! Origin's inverted index. //! @@ -18,9 +18,11 @@ use crate::sync::wire::ack_status::AckStatus; /// FTS index request (Lite → Origin, 0xA6). /// -/// Requests that Origin index the concatenated text of a document into its -/// inverted BM25 index. Origin allocates a surrogate for `(collection, doc_id)` -/// and calls `InvertedIndex::index_document` on the Data Plane. +/// Requests that Origin index a document's top-level string fields into its +/// inverted BM25 indexes: the whole-document index and one index per field. +/// Origin allocates a surrogate for `(collection, doc_id)` and calls +/// `InvertedIndex::index_document` on the Data Plane. Empty `fields` remove +/// the document from every index. #[derive( Debug, Clone, Serialize, Deserialize, zerompk::ToMessagePack, zerompk::FromMessagePack, )] @@ -31,8 +33,8 @@ pub struct FtsIndexMsg { pub collection: String, /// External document identifier. pub doc_id: String, - /// Concatenated text to index (all string fields joined by space). - pub text: String, + /// `(field, text)` per top-level string field of the document. + pub fields: Vec<(String, String)>, /// Monotonic batch ID (Lite-assigned, per-document). Used for ACK correlation. pub batch_id: u64, /// Stable identity of the originating producer. 0 for legacy clients. diff --git a/nodedb-wal/src/record/fts_spatial.rs b/nodedb-wal/src/record/fts_spatial.rs index 599a570d3..67f278c76 100644 --- a/nodedb-wal/src/record/fts_spatial.rs +++ b/nodedb-wal/src/record/fts_spatial.rs @@ -14,7 +14,8 @@ //! │producer_id u64│ epoch u64│stream_id u64│ seq u64 │ name_len u32 │ collection │ id_len u32 + id... │ //! └──────────────┴──────────┴───────────┴──────────┴──────────────┴──────────────┴────────────────────┘ //! ``` -//! FtsIndex additionally carries `text_len u32 + text bytes`. +//! FtsIndex additionally carries `field_count u32` followed by +//! `field_len u32 + field bytes + text_len u32 + text bytes` per field. //! SpatialPut additionally carries `field_len u32 + field bytes + geometry_len u32 + geometry bytes`. //! SpatialDelete additionally carries `field_len u32 + field bytes`. //! FtsDelete and SpatialDelete prefix the id with a presence tag `u8`: @@ -163,10 +164,9 @@ fn push_provenance(buf: &mut Vec, prov: &SyncProvenance) { /// WAL payload for `RecordType::FtsIndex`. /// /// Carries the minimum fields needed for Data-Plane replay: the collection -/// name, document identifier, the text to index, and producer provenance for -/// idempotency checks. Additional fields (field weights, analyzer config) can -/// be appended when the handler is wired — the length-prefixed layout is -/// forward-compatible. +/// name, document identifier, the document's `(field, text)` pairs, and +/// producer provenance for idempotency checks. Empty `fields` remove the +/// document from every index. #[derive(Debug, Clone, PartialEq, Eq)] pub struct FtsIndexPayload { /// Producer provenance for idempotency checks. @@ -175,22 +175,26 @@ pub struct FtsIndexPayload { pub collection: String, /// External document identifier. pub doc_id: String, - /// Concatenated text to index. - pub text: String, + /// `(field, text)` per top-level string field. + pub fields: Vec<(String, String)>, } +/// Smallest encoding of one `(field, text)` pair: two empty length-prefixed +/// strings. +const MIN_FIELD_PAIR_BYTES: usize = 8; + impl FtsIndexPayload { pub fn new( provenance: SyncProvenance, collection: impl Into, doc_id: impl Into, - text: impl Into, + fields: Vec<(String, String)>, ) -> Self { Self { provenance, collection: collection.into(), doc_id: doc_id.into(), - text: text.into(), + fields, } } @@ -199,7 +203,14 @@ impl FtsIndexPayload { push_provenance(&mut buf, &self.provenance); push_str_field(&mut buf, &self.collection)?; push_str_field(&mut buf, &self.doc_id)?; - push_str_field(&mut buf, &self.text)?; + let count = u32::try_from(self.fields.len()).map_err(|_| WalError::InvalidPayload { + detail: format!("too many FTS fields: {}", self.fields.len()), + })?; + push_u32_le(&mut buf, count); + for (field, text) in &self.fields { + push_str_field(&mut buf, field)?; + push_str_field(&mut buf, text)?; + } Ok(buf) } @@ -209,12 +220,28 @@ impl FtsIndexPayload { off = next; let (doc_id, next) = read_utf8_field(buf, off)?; off = next; - let (text, _) = read_utf8_field(buf, off)?; + let count = read_u32_le(buf, off)? as usize; + off += 4; + let remaining = buf.len().saturating_sub(off); + if count > remaining / MIN_FIELD_PAIR_BYTES { + return Err(WalError::InvalidPayload { + detail: format!( + "FTS field count {count} exceeds what {remaining} remaining bytes can hold" + ), + }); + } + let mut fields = Vec::with_capacity(count); + for _ in 0..count { + let (field, next) = read_utf8_field(buf, off)?; + let (text, next) = read_utf8_field(buf, next)?; + off = next; + fields.push((field, text)); + } Ok(Self { provenance, collection, doc_id, - text, + fields, }) } } @@ -392,21 +419,28 @@ mod tests { } } + fn pairs(fields: &[(&str, &str)]) -> Vec<(String, String)> { + fields + .iter() + .map(|(f, t)| ((*f).to_string(), (*t).to_string())) + .collect() + } + #[test] fn fts_index_roundtrip() { let p = FtsIndexPayload::new( prov(0xCAFE_BABE, 3, 7, 42), "articles", "doc-1", - "hello world", + pairs(&[("body", "hello world"), ("title", "Rust"), ("empty", "")]), ); let bytes = p.to_bytes().unwrap(); assert_eq!(FtsIndexPayload::from_bytes(&bytes).unwrap(), p); } #[test] - fn fts_index_empty_text_roundtrip() { - let p = FtsIndexPayload::new(prov(0, 0, 0, 0), "c", "d", ""); + fn fts_index_no_fields_roundtrip() { + let p = FtsIndexPayload::new(prov(0, 0, 0, 0), "c", "d", Vec::new()); assert_eq!( FtsIndexPayload::from_bytes(&p.to_bytes().unwrap()).unwrap(), p @@ -480,9 +514,20 @@ mod tests { #[test] fn truncated_buf_rejected() { - let p = FtsIndexPayload::new(prov(1, 2, 3, 4), "col", "id", "text"); + let p = FtsIndexPayload::new(prov(1, 2, 3, 4), "col", "id", pairs(&[("body", "text")])); let bytes = p.to_bytes().unwrap(); // Truncated to just provenance — should fail on collection field. assert!(FtsIndexPayload::from_bytes(&bytes[..32]).is_err()); + // Truncated inside the last field's text. + assert!(FtsIndexPayload::from_bytes(&bytes[..bytes.len() - 1]).is_err()); + } + + #[test] + fn fts_index_field_count_past_the_buffer_is_rejected() { + let p = FtsIndexPayload::new(prov(1, 2, 3, 4), "col", "id", Vec::new()); + let mut bytes = p.to_bytes().unwrap(); + let count_at = bytes.len() - 4; + bytes[count_at..].copy_from_slice(&u32::MAX.to_le_bytes()); + assert!(FtsIndexPayload::from_bytes(&bytes).is_err()); } } diff --git a/nodedb/src/control/server/sync/fts_handler.rs b/nodedb/src/control/server/sync/fts_handler.rs index 56673727e..73f180aec 100644 --- a/nodedb/src/control/server/sync/fts_handler.rs +++ b/nodedb/src/control/server/sync/fts_handler.rs @@ -28,14 +28,15 @@ use crate::types::{DatabaseId, TenantId, VShardId}; /// [`SyncAckResult`] for gate status propagation to the Lite client. #[async_trait] pub trait FtsDispatcher: Send + Sync { - /// Index a document's text on the Data Plane. + /// Index a document's `(field, text)` pairs on the Data Plane. Empty + /// `fields` remove the document from every index. async fn dispatch_index( &self, tenant_id: TenantId, vshard: VShardId, collection: String, surrogate: Surrogate, - text: String, + fields: Vec<(String, String)>, provenance: nodedb_types::sync::wire::SyncProvenance, ) -> crate::Result>; @@ -88,7 +89,7 @@ impl<'a> FtsDispatcher for SharedStateFtsDispatcher<'a> { vshard: VShardId, collection: String, surrogate: Surrogate, - text: String, + fields: Vec<(String, String)>, provenance: nodedb_types::sync::wire::SyncProvenance, ) -> crate::Result> { use crate::bridge::envelope::PhysicalPlan; @@ -112,7 +113,7 @@ impl<'a> FtsDispatcher for SharedStateFtsDispatcher<'a> { let plan = PhysicalPlan::Text(TextOp::FtsIndexDoc { collection: nodedb_types::QualifiedCollection::new(database_id, &collection), surrogate, - text, + fields, provenance: Some(prov), }); @@ -202,7 +203,7 @@ impl FtsDispatcher for NoOpFtsDispatcher { _vshard: VShardId, _collection: String, _surrogate: Surrogate, - _text: String, + _fields: Vec<(String, String)>, _provenance: nodedb_types::sync::wire::SyncProvenance, ) -> crate::Result> { Err(super::raft_dispatch::noop_dispatch_error("FTS index")) @@ -254,7 +255,7 @@ mod tests { use super::super::wire::*; use super::*; - type MockCallLog = Arc>>; + type MockCallLog = Arc)>>>; struct MockDispatcher { index_calls: MockCallLog, @@ -308,14 +309,14 @@ mod tests { _vshard: VShardId, collection: String, _surrogate: Surrogate, - _text: String, + fields: Vec<(String, String)>, provenance: nodedb_types::sync::wire::SyncProvenance, ) -> crate::Result> { let seq = provenance.seq; self.index_calls .lock() .unwrap() - .push((tenant_id, collection, String::new())); + .push((tenant_id, collection, fields)); super::super::test_support::mock_applied_ack(&self.result, seq) } @@ -332,7 +333,7 @@ mod tests { self.delete_calls .lock() .unwrap() - .push((tenant_id, collection, String::new())); + .push((tenant_id, collection, Vec::new())); super::super::test_support::mock_applied_ack(&self.result, seq) } @@ -362,12 +363,19 @@ mod tests { SyncSession::new("test-fts-session".to_string()) } + /// An index message whose only field is `body`; empty `text` carries no + /// fields at all. fn make_index_msg(collection: &str, doc_id: &str, text: &str) -> FtsIndexMsg { + let fields = if text.is_empty() { + Vec::new() + } else { + vec![("body".to_string(), text.to_string())] + }; FtsIndexMsg { lite_id: "lite-test".to_string(), collection: collection.to_string(), doc_id: doc_id.to_string(), - text: text.to_string(), + fields, batch_id: 1, producer_id: 0, epoch: 0, @@ -412,11 +420,19 @@ mod tests { let ack: FtsIndexAckMsg = frame.unwrap().decode_body().unwrap(); assert!(ack.accepted); assert_eq!(ack.doc_id, "d1"); - assert_eq!(indexes.lock().unwrap().len(), 1); + let calls = indexes.lock().unwrap(); + assert_eq!(calls.len(), 1); + assert_eq!( + calls[0].2, + vec![("body".to_string(), "hello world".to_string())], + "the message's fields reach the Data Plane unchanged" + ); } + /// Empty fields are an update that stripped every string field: they + /// dispatch, so the Data Plane removes the document's prior postings. #[tokio::test] - async fn empty_text_acks_without_dispatch() { + async fn empty_fields_dispatch_as_a_removal() { let mut session = make_session(); session.authenticated = true; let (mock, indexes, _) = MockDispatcher::ok(); @@ -425,7 +441,10 @@ mod tests { let frame = session.handle_fts_index(&msg, &mock).await; let ack: FtsIndexAckMsg = frame.unwrap().decode_body().unwrap(); assert!(ack.accepted); - assert!(indexes.lock().unwrap().is_empty()); + let calls = indexes.lock().unwrap(); + assert_eq!(calls.len(), 1, "empty fields must reach the Data Plane"); + assert!(calls[0].2.is_empty()); + assert_eq!(*mock.assigned.lock().unwrap(), vec!["d1".to_string()]); } #[tokio::test] diff --git a/nodedb/src/control/server/sync/fts_session.rs b/nodedb/src/control/server/sync/fts_session.rs index ee6e16fdc..49b5dcd32 100644 --- a/nodedb/src/control/server/sync/fts_session.rs +++ b/nodedb/src/control/server/sync/fts_session.rs @@ -38,20 +38,8 @@ impl SyncSession { return SyncFrame::try_encode(SyncMessageType::FtsIndexAck, &ack); } - if msg.text.is_empty() { - // Empty text — nothing to index; ACK immediately. - let ack = FtsIndexAckMsg { - collection: msg.collection.clone(), - doc_id: msg.doc_id.clone(), - batch_id: msg.batch_id, - accepted: true, - reject_reason: None, - applied_seq: msg.seq, - status: AckStatus::Applied, - }; - return SyncFrame::try_encode(SyncMessageType::FtsIndexAck, &ack); - } - + // Empty fields dispatch like any other text: indexing no text removes + // the document's prior postings from every index. let surrogate = match dispatcher .assign_surrogate( self.database_id(), @@ -105,7 +93,7 @@ impl SyncSession { vshard, msg.collection.clone(), surrogate, - msg.text.clone(), + msg.fields.clone(), nodedb_types::sync::wire::SyncProvenance { producer_id: self.producer_id, epoch: self.accepted_epoch, diff --git a/nodedb/src/control/server/wal_dispatch/core.rs b/nodedb/src/control/server/wal_dispatch/core.rs index 8e2afc044..b7ef91567 100644 --- a/nodedb/src/control/server/wal_dispatch/core.rs +++ b/nodedb/src/control/server/wal_dispatch/core.rs @@ -283,7 +283,7 @@ mod tests { let plan = PhysicalPlan::Text(TextOp::FtsIndexDoc { collection: QualifiedCollection::new(DatabaseId::DEFAULT, "docs"), surrogate: Surrogate::new(7), - text: "hello world".to_string(), + fields: vec![("body".to_string(), "hello world".to_string())], provenance: None, }); @@ -304,7 +304,10 @@ mod tests { let decoded = nodedb_wal::record::FtsIndexPayload::from_bytes(&record.payload).expect("decode"); assert_eq!(decoded.collection, "docs"); - assert_eq!(decoded.text, "hello world"); + assert_eq!( + decoded.fields, + vec![("body".to_string(), "hello world".to_string())] + ); assert_eq!( decoded.doc_id, crate::engine::document::store::StorageKey::for_surrogate(Surrogate::new(7)) diff --git a/nodedb/src/control/server/wal_dispatch/text.rs b/nodedb/src/control/server/wal_dispatch/text.rs index 84a43923b..6a113d9ef 100644 --- a/nodedb/src/control/server/wal_dispatch/text.rs +++ b/nodedb/src/control/server/wal_dispatch/text.rs @@ -45,14 +45,18 @@ pub(crate) fn encode_text_op_record(op: &TextOp) -> crate::Result { let doc_id = crate::engine::document::store::StorageKey::for_surrogate(*surrogate).to_string(); let prov = provenance.clone().unwrap_or_default(); - let payload = - nodedb_wal::record::FtsIndexPayload::new(prov, collection.as_str(), &doc_id, text); + let payload = nodedb_wal::record::FtsIndexPayload::new( + prov, + collection.as_str(), + &doc_id, + fields.clone(), + ); Some(( RecordType::FtsIndex, payload.to_bytes().map_err(crate::Error::Wal)?, diff --git a/nodedb/src/control/wal_replication/decode/entry_columnar_family.rs b/nodedb/src/control/wal_replication/decode/entry_columnar_family.rs index 172b497c1..3c1678ead 100644 --- a/nodedb/src/control/wal_replication/decode/entry_columnar_family.rs +++ b/nodedb/src/control/wal_replication/decode/entry_columnar_family.rs @@ -67,9 +67,9 @@ pub(super) fn decode_arm(write: &ReplicatedWrite) -> crate::Result ReplicatedWrite::FtsIndex { collection, surrogate, - text, + fields, provenance, - } => decode_sync_engines::fts_index(collection, *surrogate, text, provenance), + } => decode_sync_engines::fts_index(collection, *surrogate, fields, provenance), ReplicatedWrite::FtsDelete { collection, surrogate, diff --git a/nodedb/src/control/wal_replication/decode_sync_engines.rs b/nodedb/src/control/wal_replication/decode_sync_engines.rs index ee049c89f..f746e7ffc 100644 --- a/nodedb/src/control/wal_replication/decode_sync_engines.rs +++ b/nodedb/src/control/wal_replication/decode_sync_engines.rs @@ -120,14 +120,14 @@ pub fn timeseries_ingest( pub fn fts_index( collection: &str, surrogate: u32, - text: &str, + fields: &[(String, String)], prov_bytes: &Option>, ) -> crate::Result { let provenance = decode_provenance(prov_bytes)?; Ok(PhysicalPlan::Text(TextOp::FtsIndexDoc { collection: nodedb_types::QualifiedCollection::from_stored(collection.to_owned()), surrogate: Surrogate::new(surrogate), - text: text.to_owned(), + fields: fields.to_vec(), provenance, })) } @@ -603,7 +603,10 @@ mod tests { let plan = PhysicalPlan::Text(TextOp::FtsIndexDoc { collection: QualifiedCollection::new(DatabaseId::DEFAULT, "articles"), surrogate: nodedb_types::Surrogate::new(500), - text: "hello world".into(), + fields: vec![ + ("body".to_string(), "hello world".to_string()), + ("title".to_string(), "greeting".to_string()), + ], provenance: Some(prov.clone()), }); let entry = to_replicated_entry(tenant, DatabaseId::DEFAULT, vshard, &plan) @@ -616,10 +619,18 @@ mod tests { match decoded_plan { PhysicalPlan::Text(TextOp::FtsIndexDoc { surrogate, + fields, provenance, .. }) => { assert_eq!(surrogate, nodedb_types::Surrogate::new(500)); + assert_eq!( + fields, + vec![ + ("body".to_string(), "hello world".to_string()), + ("title".to_string(), "greeting".to_string()), + ] + ); assert_eq!(provenance, Some(prov)); } other => panic!("expected Text(FtsIndexDoc), got {other:?}"), diff --git a/nodedb/src/control/wal_replication/encode/columnar.rs b/nodedb/src/control/wal_replication/encode/columnar.rs index 516bd1997..42dbe7a86 100644 --- a/nodedb/src/control/wal_replication/encode/columnar.rs +++ b/nodedb/src/control/wal_replication/encode/columnar.rs @@ -67,13 +67,13 @@ pub(super) fn timeseries_ingest(fields: TimeseriesIngestFields<'_>) -> Replicate pub(super) fn fts_index( collection: &str, surrogate: u32, - text: &str, + fields: &[(String, String)], provenance: Option>, ) -> ReplicatedWrite { ReplicatedWrite::FtsIndex { collection: collection.to_owned(), surrogate, - text: text.to_owned(), + fields: fields.to_vec(), provenance, } } @@ -144,7 +144,7 @@ pub(super) fn bulk_delete(collection: &str, filters: &[u8]) -> ReplicatedWrite { pub(super) fn bulk_update( collection: &str, filters: &[u8], - updates: &[(String, Vec)], + updates: &[(String, nodedb_physical::physical_plan::UpdateValue)], ) -> ReplicatedWrite { ReplicatedWrite::ColumnarBulkDml { collection: collection.to_owned(), diff --git a/nodedb/src/control/wal_replication/encode/entry_columnar_family.rs b/nodedb/src/control/wal_replication/encode/entry_columnar_family.rs index fd0652704..f14547df0 100644 --- a/nodedb/src/control/wal_replication/encode/entry_columnar_family.rs +++ b/nodedb/src/control/wal_replication/encode/entry_columnar_family.rs @@ -183,12 +183,12 @@ pub(super) fn text_write(op: &TextOp) -> Option { TextOp::FtsIndexDoc { collection, surrogate, - text, + fields, provenance, } => columnar::fts_index( collection.as_str(), surrogate.as_u32(), - text, + fields, encode_provenance(provenance), ), TextOp::FtsDeleteDoc { diff --git a/nodedb/src/control/wal_replication/types/replicated_write.rs b/nodedb/src/control/wal_replication/types/replicated_write.rs index eab6d546e..814b58832 100644 --- a/nodedb/src/control/wal_replication/types/replicated_write.rs +++ b/nodedb/src/control/wal_replication/types/replicated_write.rs @@ -323,7 +323,8 @@ pub enum ReplicatedWrite { collection: String, /// Leader-assigned global surrogate for the document. surrogate: u32, - text: String, + /// `(field, text)` per top-level string field. + fields: Vec<(String, String)>, /// Sync provenance encoded as zerompk bytes. #[serde(default)] provenance: Option>, @@ -579,7 +580,8 @@ pub enum ReplicatedWrite { source_key: Vec, dest_key: Vec, field: String, - amount: f64, + /// The amount, typed by the field it moves. + amount: nodedb_physical::physical_plan::TransferAmount, debit_surrogate: u32, credit_surrogate: u32, }, @@ -637,7 +639,7 @@ pub enum ReplicatedWrite { collection: String, filters: Vec, is_update: bool, - updates: Vec<(String, Vec)>, + updates: Vec<(String, nodedb_physical::physical_plan::UpdateValue)>, }, InsertSelect { target_collection: String, diff --git a/nodedb/src/data/executor/core_loop/accessors.rs b/nodedb/src/data/executor/core_loop/accessors.rs index 70f23be4d..9a50285c3 100644 --- a/nodedb/src/data/executor/core_loop/accessors.rs +++ b/nodedb/src/data/executor/core_loop/accessors.rs @@ -266,8 +266,8 @@ impl CoreLoop { .unwrap_or_else(crate::engine::kv::current_ms) } - /// Write a raw segment blob directly into the FTS LSM segment store for - /// a given `(tenant, collection)`. + /// Write a raw segment blob directly into the FTS LSM segment store of a + /// collection's whole-document index. /// /// This bypasses the memtable flush path and is intended for maintenance /// tests and bootstrapping code that need to pre-populate a known number @@ -286,7 +286,7 @@ impl CoreLoop { self.inverted.backend().write_segment( database_id, tenant.as_u64(), - collection, + nodedb_fts::IndexScope::document(collection), segment_id, data, ) @@ -304,8 +304,10 @@ impl CoreLoop { collection: &str, ) -> crate::Result> { use nodedb_fts::backend::FtsBackend; - self.inverted - .backend() - .list_segments(database_id, tenant.as_u64(), collection) + self.inverted.backend().list_segments( + database_id, + tenant.as_u64(), + nodedb_fts::IndexScope::document(collection), + ) } } diff --git a/nodedb/src/data/executor/fts_text.rs b/nodedb/src/data/executor/fts_text.rs index 0195f987c..65aabb1a5 100644 --- a/nodedb/src/data/executor/fts_text.rs +++ b/nodedb/src/data/executor/fts_text.rs @@ -1,22 +1,57 @@ // SPDX-License-Identifier: BUSL-1.1 -//! Shared full-text extraction: concatenate a document's string field values -//! into the single text blob the inverted index analyzes. +//! Shared full-text extraction: a document's top-level string fields, the +//! text the inverted index analyzes per field and as a whole. -/// Concatenate all top-level string-valued fields of a document object into -/// the text the full-text inverted index indexes. +use nodedb_fts::DocumentText; + +/// Collect every top-level string-valued field of a document object as the +/// text the full-text inverted index indexes. /// /// Used by the forward PUT indexing path AND by DELETE-rollback re-indexing so -/// both produce byte-identical text (and therefore identical postings, since +/// both produce identical text (and therefore identical postings, since /// `nodedb_fts::analyze` is deterministic). Non-object values and non-string /// fields contribute nothing. -pub(in crate::data::executor) fn extract_fts_text(doc: &serde_json::Value) -> String { +pub(in crate::data::executor) fn extract_fts_fields(doc: &serde_json::Value) -> DocumentText { match doc.as_object() { - Some(obj) => obj - .values() - .filter_map(|v| v.as_str()) - .collect::>() - .join(" "), - None => String::new(), + Some(obj) => DocumentText::from_fields( + obj.iter() + .filter_map(|(k, v)| v.as_str().map(|s| (k.clone(), s.to_owned()))), + ), + None => DocumentText::default(), + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn only_top_level_strings_are_collected_in_field_order() { + let doc = serde_json::json!({ + "title": "Rust", + "body": "fearless", + "n": 3, + "nested": {"inner": "skip"}, + }); + let text = extract_fts_fields(&doc); + assert_eq!( + text.fields(), + &[ + ("body".to_string(), "fearless".to_string()), + ("title".to_string(), "Rust".to_string()), + ] + ); + assert_eq!(text.whole(), "fearless Rust"); + } + + #[test] + fn a_non_object_has_no_text() { + assert!(extract_fts_fields(&serde_json::json!("plain")).is_empty()); + assert!( + extract_fts_fields(&serde_json::json!(null)) + .fields() + .is_empty() + ); } } diff --git a/nodedb/src/data/executor/handlers/bulk_dml/delete.rs b/nodedb/src/data/executor/handlers/bulk_dml/delete.rs index 05fd05b8e..fccc93f5e 100644 --- a/nodedb/src/data/executor/handlers/bulk_dml/delete.rs +++ b/nodedb/src/data/executor/handlers/bulk_dml/delete.rs @@ -16,7 +16,7 @@ use nodedb_physical::physical_plan::{ OllpPredictedEdge, ResolvedSumTarget, ReturningSpec, StorageMode, }; -use super::delete_cascade::BulkDeleteRowCascade; +use super::delete_cascade::{BulkDeleteRowCascade, TextFlush}; /// OLLP prediction inputs threaded to `execute_bulk_delete`: the predicted /// matched-doc surrogate set and the predicted implicit-edge set. Both are @@ -103,12 +103,7 @@ impl CoreLoop { ) { Ok(ids) => ids, Err(e) => { - return self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ); + return self.response_error(task, ErrorCode::from(e)); } }; @@ -179,6 +174,7 @@ impl CoreLoop { // The period lock is judged here too, on the same image: each row // below commits in its own transaction, so a lock judged there // refuses after the rows ahead of it were removed. + let identity_column = self.identity_column(database_id, tid, collection); let gate_policy = !matches!( rls_write_check.decision(), nodedb_types::WriteGateDecision::AdmitAll @@ -200,6 +196,7 @@ impl CoreLoop { &stored, &key.to_identity(), strict_schema.as_ref(), + &identity_column, tid, collection, ) @@ -251,9 +248,10 @@ impl CoreLoop { } else { Vec::new() }; + // Removed rows whose text has not left the inverted index yet. It + // leaves in batches, one write transaction each. + let mut pending_text: Vec = Vec::new(); for storage_key in &apply_ids { - let doc_id = storage_key.to_string(); - // Capture pre-deletion snapshot if RETURNING was requested, or if // the collection is indexed (needed to recompute the removed // secondary-index tuples below — the delete cascade's prefix scan @@ -281,11 +279,19 @@ impl CoreLoop { &bytes, &identity, strict_schema.as_ref(), + &identity_column, ) { Ok(doc) => Some(doc), Err(e) => { let code = refusal_after_rows(affected, e); - return self.refusal_with_landed_rows(task, code, write_set); + return self.bulk_delete_refusal( + task, + tid, + collection, + code, + &mut pending_text, + write_set, + ); } } } @@ -305,7 +311,14 @@ impl CoreLoop { Ok(txn) => txn, Err(e) => { let code = refusal_after_rows(affected, e); - return self.refusal_with_landed_rows(task, code, write_set); + return self.bulk_delete_refusal( + task, + tid, + collection, + code, + &mut pending_text, + write_set, + ); } }; let deleted_bytes = self @@ -344,7 +357,14 @@ impl CoreLoop { // and every target it had already debited. Err(e) => { let code = refusal_after_rows(affected, e); - return self.refusal_with_landed_rows(task, code, write_set); + return self.bulk_delete_refusal( + task, + tid, + collection, + code, + &mut pending_text, + write_set, + ); } } } @@ -355,7 +375,14 @@ impl CoreLoop { detail: format!("bulk delete commit: {e}"), }, ); - return self.refusal_with_landed_rows(task, code, write_set); + return self.bulk_delete_refusal( + task, + tid, + collection, + code, + &mut pending_text, + write_set, + ); } // One durable redo entry per debited target row, naming the TARGET // collection: this statement's own redo describes the removed source @@ -363,13 +390,12 @@ impl CoreLoop { // as it stood before the delete. write_set.extend(write_hook::target_write_set(&target_writes)); if let Some(bytes) = deleted_bytes.as_deref() { - self.bulk_delete_row_cascade( + let cascaded = self.bulk_delete_row_cascade( BulkDeleteRowCascade { task, database_id, tid, collection, - doc_id: doc_id.as_str(), storage_key: *storage_key, deleted_bytes: bytes, strict_schema: strict_schema.as_ref(), @@ -383,8 +409,31 @@ impl CoreLoop { &mut returned_docs, ); affected += 1; + pending_text.push(storage_key.surrogate()); + // The row is removed and journalled: its index cleanup error + // fails the statement after its pending text leaves the + // inverted index. + if let Err(e) = cascaded { + let code = refusal_after_rows(affected, ErrorCode::from(e)); + return self.bulk_delete_refusal( + task, + tid, + collection, + code, + &mut pending_text, + write_set, + ); + } + let flush = TextFlush::full_batch(tid, collection, affected); + if let Err(code) = self.flush_deleted_text_at(task, flush, &mut pending_text) { + return self.refusal_with_landed_rows(task, code, write_set); + } } } + let flush = TextFlush::remainder(tid, collection, affected); + if let Err(code) = self.flush_deleted_text_at(task, flush, &mut pending_text) { + return self.refusal_with_landed_rows(task, code, write_set); + } // Invalidate aggregate cache — a delete changes count(*) for this // collection. Only needed when at least one row was actually removed. @@ -403,23 +452,13 @@ impl CoreLoop { let mut response = if let Some(spec) = returning { match returning_rows::build_rows_payload(spec, rls_filters, &returned_docs) { Ok(payload) => self.response_with_payload(task, payload), - Err(e) => self.response_error( - task, - ErrorCode::Internal { - detail: format!("RETURNING encode: {e}"), - }, - ), + Err(e) => self.response_error(task, ErrorCode::from(e)), } } else { let result = serde_json::json!({ "affected": affected }); match response_codec::encode_json_as_msgpack(&result) { Ok(payload) => self.response_with_payload(task, payload), - Err(e) => self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ), + Err(e) => self.response_error(task, ErrorCode::from(e)), } }; response.write_set = write_set; diff --git a/nodedb/src/data/executor/handlers/bulk_dml/delete_cascade.rs b/nodedb/src/data/executor/handlers/bulk_dml/delete_cascade.rs index 401b4d0bf..0b9f625bc 100644 --- a/nodedb/src/data/executor/handlers/bulk_dml/delete_cascade.rs +++ b/nodedb/src/data/executor/handlers/bulk_dml/delete_cascade.rs @@ -1,22 +1,78 @@ // SPDX-License-Identifier: BUSL-1.1 -//! Post-commit cascade for one bulk-deleted row: inverted index, secondary -//! indexes, deleted-node bookkeeping, vector index, doc cache, write-version -//! tracking, and the Event Plane emit. +//! Post-commit cascade for one bulk-deleted row, step by step: +//! +//! - Secondary indexes: removes the row's entries. An error fails the +//! statement after the row's other steps run (see below). +//! - Deleted-node bookkeeping, vector index, R-tree, sparse-vector index, doc +//! cache, write-version and index write-value tracking: in-memory updates +//! that cannot fail. +//! - Journal: the row's delete entry joins the statement's write set. +//! - Event Plane: one delete event per row. +//! - Inverted index: the delete loop collects removed surrogates and removes +//! their text in one write transaction per batch. An error fails the +//! statement. //! //! Runs AFTER the row's own transaction committed, so none of this reverses -//! on failure — each step logs and continues rather than aborting a -//! statement it cannot undo. +//! on failure. A failed step fails the statement with every removed row +//! journalled, its event emitted, and its pending text removed first. +use nodedb_types::Surrogate; use nodedb_types::columnar::StrictSchema; -use tracing::warn; -use crate::bridge::envelope::WriteSetEntry; +use crate::bridge::envelope::{ErrorCode, Response, WriteSetEntry}; use crate::data::executor::core_loop::CoreLoop; +use crate::data::executor::handlers::partial_refusal::refusal_after_rows; +use crate::data::executor::handlers::point::apply_put::SpatialEntryId; use crate::data::executor::handlers::transaction::stage_write::stored_row_identity; use crate::data::executor::task::ExecutionTask; use crate::engine::document::store::{IndexPath, StorageKey}; +/// Removed rows whose text leaves the inverted index in one write +/// transaction. +const BULK_DELETE_TEXT_BATCH: usize = 512; + +/// When a bulk delete flushes its pending text, and the statement it +/// belongs to. +pub(in crate::data::executor) struct TextFlush<'a> { + tid: u64, + collection: &'a str, + /// Flush once the pending rows reach this count. + at_least: usize, + /// Rows the statement removed so far. + affected: u64, +} + +impl<'a> TextFlush<'a> { + /// Flush once a full batch is pending. + pub(in crate::data::executor) fn full_batch( + tid: u64, + collection: &'a str, + affected: u64, + ) -> Self { + Self { + tid, + collection, + at_least: BULK_DELETE_TEXT_BATCH, + affected, + } + } + + /// Flush whatever is pending: the statement's last batch. + pub(in crate::data::executor) fn remainder( + tid: u64, + collection: &'a str, + affected: u64, + ) -> Self { + Self { + tid, + collection, + at_least: 1, + affected, + } + } +} + /// Borrowed + owned inputs for [`CoreLoop::bulk_delete_row_cascade`], grouped /// so the call stays within the argument-count budget. pub(in crate::data::executor) struct BulkDeleteRowCascade<'a> { @@ -24,7 +80,6 @@ pub(in crate::data::executor) struct BulkDeleteRowCascade<'a> { pub database_id: u64, pub tid: u64, pub collection: &'a str, - pub doc_id: &'a str, pub storage_key: StorageKey, /// The row's pre-deletion bytes, as `sparse.delete` returned them — /// never re-read. @@ -47,18 +102,21 @@ impl CoreLoop { /// Cascade one committed row removal into every secondary structure a /// bulk delete must also clean up, and account it into `write_set` / /// `returned_docs`. + /// + /// Every step runs for the committed row, so its journal entry and event + /// are never lost. `Err` is the secondary-index removal error: the + /// statement fails with this row counted as removed. pub(in crate::data::executor) fn bulk_delete_row_cascade( &mut self, cascade: BulkDeleteRowCascade<'_>, write_set: &mut Vec, returned_docs: &mut Vec, - ) { + ) -> crate::Result<()> { let BulkDeleteRowCascade { task, database_id, tid, collection, - doc_id, storage_key, deleted_bytes, strict_schema, @@ -79,29 +137,17 @@ impl CoreLoop { storage_key, ); - // Cascade: inverted index. let row_surrogate = storage_key.surrogate(); - if let Err(e) = self.inverted.remove_document( - task.request.database_id.as_u64(), - crate::types::TenantId::new(tid), - collection, - row_surrogate, - ) { - // Recorded here, at the detection site: the row's own - // transaction has already committed, so this cleanup - // failure cannot roll it back. - crate::diag::orphaned_index_entry_after_delete(&e, collection, "inverted"); - warn!(core = self.core_id, %collection, %doc_id, error = %e, "bulk delete: inverted index removal failed"); - } - // Cascade: secondary indexes. - if let Err(e) = self.sparse.delete_indexes_for_document( + // Cascade: secondary indexes. An error is recorded at its detection + // site and returned once the row's remaining steps ran. + let secondary = self.sparse.delete_indexes_for_document( task.request.database_id.as_u64(), tid, collection, &storage_key, - ) { - crate::diag::orphaned_index_entry_after_delete(&e, collection, "secondary"); - warn!(core = self.core_id, %collection, %doc_id, error = %e, "bulk delete: secondary index cascade failed"); + ); + if let Err(e) = &secondary { + crate::diag::orphaned_index_entry_after_delete(e, collection, "secondary"); } // The row's graph node keeps its edges here: the delete's own // transaction tombstones them with `EdgeDelete` tasks. The node is @@ -114,6 +160,22 @@ impl CoreLoop { if has_vectors { self.remove_document_vector_indexes(database_id, tid, collection, storage_key); } + // Cascade: the row's R-tree entries and sparse-vector postings, keyed + // as the put path indexed them. The row is committed, so the undo + // entries go unused. + self.remove_document_spatial_indexes( + database_id, + tid, + collection, + SpatialEntryId::from_storage_key(storage_key), + ); + self.remove_document_sparse_indexes( + database_id, + tid, + collection, + storage_key, + &mut Vec::new(), + ); self.doc_cache.invalidate( task.request.database_id.as_u64(), tid, @@ -147,25 +209,98 @@ impl CoreLoop { // each row a bulk DELETE removed — mirroring // `execute_point_delete`'s single-row emit. `deleted_bytes` is // the prior stored bytes `sparse.delete` returned above (no - // second read needed); `resolve_event_payload` handles the - // strict->msgpack conversion for triggers. Emitted per row + // second read needed); the emit converts a strict row to + // MessagePack for triggers. Emitted per row // (not a `WriteOp::BulkDelete` summary) — the Event Plane's // WAL-replay bulk variant is aggregate metadata reconstructed // only when the live per-row events were lost. - let old_converted = self.resolve_event_payload( - task.request.database_id.as_u64(), - tid, - collection, - deleted_bytes, - ); - self.emit_document_delete_event( - task, - collection, - row_identity, - Some(old_converted.as_deref().unwrap_or(deleted_bytes)), - ); + self.emit_document_delete_event(task, tid, collection, row_identity, Some(deleted_bytes)); if returning && let Some(doc) = pre_delete_doc { returned_docs.push(nodedb_types::Value::from(doc)); } + secondary + } + + /// Remove the text of the rows in `pending` from the inverted index in + /// one write transaction, then clear `pending`. + pub(in crate::data::executor) fn flush_deleted_text( + &self, + database_id: u64, + tid: u64, + collection: &str, + pending: &mut Vec, + ) -> crate::Result<()> { + if let Err(e) = self.inverted.remove_documents( + database_id, + crate::types::TenantId::new(tid), + collection, + pending, + ) { + // Recorded at the detection site: the rows' own transactions have + // committed, so the statement fails with their entries journalled. + crate::diag::orphaned_index_entry_after_delete(&e, collection, "inverted"); + return Err(e); + } + pending.clear(); + Ok(()) + } + + /// Flush `pending` per `flush`. The refusal code when the removal fails. + pub(in crate::data::executor) fn flush_deleted_text_at( + &self, + task: &ExecutionTask, + flush: TextFlush<'_>, + pending: &mut Vec, + ) -> Result<(), ErrorCode> { + if pending.len() < flush.at_least { + return Ok(()); + } + self.flush_deleted_text( + task.request.database_id.as_u64(), + flush.tid, + flush.collection, + pending, + ) + .map_err(|e| { + refusal_after_rows( + flush.affected, + ErrorCode::Internal { + detail: format!("bulk delete: removing deleted rows' text failed: {e}"), + }, + ) + }) + } + + /// The refusal of a bulk delete that stopped part-way. The rows removed + /// so far stay removed, so their pending text leaves the inverted index + /// first. A failure of that removal is the refusal. + pub(in crate::data::executor) fn bulk_delete_refusal( + &self, + task: &ExecutionTask, + tid: u64, + collection: &str, + code: ErrorCode, + pending: &mut Vec, + write_set: Vec, + ) -> Response { + let removed = pending.len() as u64; + let code = match self.flush_deleted_text( + task.request.database_id.as_u64(), + tid, + collection, + pending, + ) { + Ok(()) => code, + Err(e) => refusal_after_rows( + removed, + ErrorCode::Internal { + detail: format!( + "{code:?}; removing the deleted rows' text from the inverted index \ + failed: {e}" + ), + }, + ), + }; + self.refusal_with_landed_rows(task, code, write_set) } } diff --git a/nodedb/src/data/executor/handlers/columnar_mutation.rs b/nodedb/src/data/executor/handlers/columnar_mutation.rs index aa8b39d28..72f240929 100644 --- a/nodedb/src/data/executor/handlers/columnar_mutation.rs +++ b/nodedb/src/data/executor/handlers/columnar_mutation.rs @@ -7,21 +7,21 @@ //! the R-tree cascade for spatial collections, is shared with the //! resolved-row-set handlers through `columnar_mutation_apply.rs`. +use nodedb_physical::physical_plan::UpdateValue; use tracing::debug; use crate::bridge::envelope::{ErrorCode, Response}; -use crate::bridge::scan_filter::ScanFilter; +use crate::bridge::scan_filter::decode_scan_filters; use crate::data::executor::core_loop::CoreLoop; use crate::data::executor::handlers::columnar_resolve::{ - ResolveUpdateRowsParams, require_pk_column_index, resolve_delete_rows, resolve_update_rows, + ResolveDeleteRowsParams, ResolveUpdateRowsParams, require_pk_column_index, }; use crate::data::executor::handlers::transaction::undo::UndoEntry; use crate::data::executor::task::ExecutionTask; impl CoreLoop { - /// Handle columnar UPDATE: scan memtable for matching rows, apply field updates. - /// - /// Currently operates on in-memory memtable rows only. + /// Handle columnar UPDATE: match the current rows, flushed and in the + /// memtable, and apply the field updates to each match. /// Returns `{"affected": N}` as JSON payload. /// /// When `undo_log` is `Some` (the durable COMMIT-replay path inside a @@ -34,7 +34,7 @@ impl CoreLoop { task: &ExecutionTask, collection: &str, filter_bytes: &[u8], - updates: &[(String, Vec)], + updates: &[(String, UpdateValue)], rls_write_check: &nodedb_types::RlsWriteCheck, undo_log: Option<&mut Vec>, ) -> Response { @@ -57,18 +57,20 @@ impl CoreLoop { } }; - // Columnar UPDATE: scan memtable rows matching filter predicates, - // then apply updates via PK-based MutationEngine (delete + re-insert). + // Match the current rows against the filter predicates, then apply + // the updates through the PK-based MutationEngine (delete + re-insert). let schema = engine.schema().clone(); + let row_count_before = engine.memtable().row_count(); let pk_col_idx = match require_pk_column_index(&schema, "UPDATE") { Ok(idx) => idx, Err(e) => return self.response_error(task, e), }; - let filter_predicates: Vec = if !filter_bytes.is_empty() { - zerompk::from_msgpack(filter_bytes).unwrap_or_default() - } else { - Vec::new() + // A filter that does not decode refuses the statement. Read as no + // filter, it would update every row. + let filter_predicates = match decode_scan_filters(filter_bytes, "columnar UPDATE filter") { + Ok(filters) => filters, + Err(e) => return self.response_error(task, e), }; // Resolve every matching row's post-image, and let the write policy @@ -80,8 +82,13 @@ impl CoreLoop { // no way for the caller to see or undo that. Shared with // `execute_columnar_resolve_dml`, which reports this same selection // instead of applying it. - let pending = match resolve_update_rows(ResolveUpdateRowsParams { - engine, + // + // An expression assignment evaluates against each row's pre-image, and + // each post-image meets the declared column rule inside the resolve. A + // failed evaluation or a value past a declared width refuses the + // statement whole. + let pending = match self.resolve_columnar_update_rows(ResolveUpdateRowsParams { + key: &key, schema: &schema, pk_col_idx, filter_predicates: &filter_predicates, @@ -93,15 +100,21 @@ impl CoreLoop { Ok(rows) => rows, Err(e) => return self.response_error(task, e), }; - // Undo capture (only on the durable COMMIT-replay path). `row_count_before` // is the memtable size before any replacement row is appended, so the // undo can truncate back to it; `inserted_pks`/`displaced` reverse the // insert half, `restored` re-materializes each tombstoned original. - let row_count_before = engine.memtable().row_count(); let mut undo_log = undo_log; - let outcome = - self.apply_columnar_update_rows(task, &key, &schema, &pending, undo_log.as_deref_mut()); + let outcome = match self.apply_columnar_update_rows( + task, + &key, + &schema, + &pending, + undo_log.as_deref_mut(), + ) { + Ok(outcome) => outcome, + Err(e) => return self.response_error(task, e), + }; let affected = outcome.affected; if let Some(log) = undo_log { @@ -137,18 +150,12 @@ impl CoreLoop { let result = serde_json::json!({ "affected": affected }); match super::super::response_codec::encode_json_as_msgpack(&result) { Ok(payload) => self.response_with_payload(task, payload), - Err(e) => self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ), + Err(e) => self.response_error(task, ErrorCode::from(e)), } } - /// Handle columnar DELETE: scan memtable for matching rows, delete them. - /// - /// Currently operates on in-memory memtable rows only. + /// Handle columnar DELETE: match the current rows, flushed and in the + /// memtable, and delete each match. /// Returns `{"affected": N}` as JSON payload. /// /// When `undo_log` is `Some` (the durable COMMIT-replay path inside a @@ -189,10 +196,11 @@ impl CoreLoop { Err(e) => return self.response_error(task, e), }; - let filter_predicates: Vec = if !filter_bytes.is_empty() { - zerompk::from_msgpack(filter_bytes).unwrap_or_default() - } else { - Vec::new() + // A filter that does not decode refuses the statement. Read as no + // filter, it would delete every row. + let filter_predicates = match decode_scan_filters(filter_bytes, "columnar DELETE filter") { + Ok(filters) => filters, + Err(e) => return self.response_error(task, e), }; // The image a delete is governed by is the row it removes. Every @@ -200,15 +208,15 @@ impl CoreLoop { // rejection removes nothing at all rather than leaving the rows ahead // of it already tombstoned. Shared with `execute_columnar_resolve_dml`, // which reports this same selection instead of applying it. - let pk_values = match resolve_delete_rows( - engine, - &schema, + let pk_values = match self.resolve_columnar_delete_rows(ResolveDeleteRowsParams { + key: &key, + schema: &schema, pk_col_idx, - &filter_predicates, + filter_predicates: &filter_predicates, rls_write_check, - task.request.tenant_id.as_u64(), + tid: task.request.tenant_id.as_u64(), collection, - ) { + }) { Ok(pks) => pks, Err(e) => return self.response_error(task, e), }; @@ -217,8 +225,15 @@ impl CoreLoop { // and PK bytes of each tombstoned row, so the undo can clear its // delete-bitmap bit and re-bind the PK index. let mut undo_log = undo_log; - let outcome = - self.apply_columnar_delete_pks(&key, &schema, &pk_values, undo_log.as_deref_mut()); + let outcome = match self.apply_columnar_delete_pks( + &key, + &schema, + &pk_values, + undo_log.as_deref_mut(), + ) { + Ok(outcome) => outcome, + Err(e) => return self.response_error(task, e), + }; let affected = outcome.affected; if let Some(log) = undo_log { @@ -242,12 +257,7 @@ impl CoreLoop { let result = serde_json::json!({ "affected": affected }); match super::super::response_codec::encode_json_as_msgpack(&result) { Ok(payload) => self.response_with_payload(task, payload), - Err(e) => self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ), + Err(e) => self.response_error(task, ErrorCode::from(e)), } } } diff --git a/nodedb/src/data/executor/handlers/compact/fts.rs b/nodedb/src/data/executor/handlers/compact/fts.rs index d36c893a6..00dbe0a19 100644 --- a/nodedb/src/data/executor/handlers/compact/fts.rs +++ b/nodedb/src/data/executor/handlers/compact/fts.rs @@ -40,6 +40,7 @@ use std::sync::atomic::{AtomicU64, Ordering}; +use nodedb_fts::IndexScope; use nodedb_fts::backend::FtsBackend as _; use nodedb_fts::lsm::compaction::{ CompactError, CompactLevelParams, CompactionConfig, SegmentMeta, compact_level, @@ -64,7 +65,7 @@ pub(super) struct FtsCompactionOutcome { pub enumeration_failed: bool, } -/// Identifies one `(database, tenant, collection)` FTS compaction unit. +/// Identifies one `(database, tenant, index)` FTS compaction unit. /// /// Bundled so `compact_one_fts_collection` stays within the argument budget; /// `force`, the shared `config`, and the mutable `outcome` accumulator are @@ -72,7 +73,7 @@ pub(super) struct FtsCompactionOutcome { struct FtsCompactionTarget<'a> { database_id: u64, tid: TenantId, - collection: &'a str, + index: IndexScope<'a>, } impl CoreLoop { @@ -94,7 +95,7 @@ impl CoreLoop { // segment. This queries the FTS subsystem for what it owns — no // external registry. A failure here is distinct from per-collection // deferral and is surfaced via `enumeration_failed`. - let collections = match self.inverted.list_all_fts_collections() { + let indexes = match self.inverted.list_all_fts_indexes() { Ok(c) => c, Err(e) => { tracing::warn!( @@ -107,12 +108,12 @@ impl CoreLoop { } }; - for (database_id, tid, collection) in collections { + for (database_id, tid, collection, field) in &indexes { self.compact_one_fts_collection( FtsCompactionTarget { - database_id, - tid, - collection: &collection, + database_id: *database_id, + tid: *tid, + index: IndexScope::from_key(collection, field), }, force, &config, @@ -136,32 +137,32 @@ impl CoreLoop { let FtsCompactionTarget { database_id, tid, - collection, + index, } = target; + let collection = index.collection(); let db = nodedb_types::DatabaseId::new(database_id); let tid_u64 = tid.as_u64(); // Resolve segment list before acquiring the lease so we only hold // the lease across actual merge work, not the read path. - let segment_ids = - match self - .inverted - .backend() - .list_segments(database_id, tid_u64, collection) - { - Ok(ids) => ids, - Err(e) => { - tracing::warn!( - core = self.core_id, - tid = tid_u64, - collection = %collection, - error = %e, - "FTS compaction: failed to list segments — deferred to next cycle" - ); - outcome.deferred += 1; - return; - } - }; + let segment_ids = match self + .inverted + .backend() + .list_segments(database_id, tid_u64, index) + { + Ok(ids) => ids, + Err(e) => { + tracing::warn!( + core = self.core_id, + tid = tid_u64, + collection = %collection, + error = %e, + "FTS compaction: failed to list segments — deferred to next cycle" + ); + outcome.deferred += 1; + return; + } + }; if segment_ids.is_empty() { return; @@ -205,7 +206,7 @@ impl CoreLoop { backend: self.inverted.backend(), database_id, tid: tid_u64, - collection, + index, segments: &segments, level, governor: &self.governor, @@ -223,7 +224,7 @@ impl CoreLoop { .compact_commit(crate::engine::sparse::fts_redb::CompactCommit { database_id, tid: tid_u64, - collection, + index, new_segment_id: &new_seg_id, new_segment_data: &new_bytes, merged_ids: &merged_ids, @@ -280,6 +281,20 @@ impl CoreLoop { "FTS compaction: segment read failed, deferred" ); } + // A missing or corrupt source segment, or delete sets that do + // not decode. No segment is replaced. The backend's segment read + // records a corrupt segment in the quarantine registry. + Err(CompactError::Index(e)) => { + outcome.deferred += 1; + tracing::error!( + core = self.core_id, + tid = tid_u64, + collection = %collection, + level, + error = %e, + "FTS compaction: index state unreadable, deferred" + ); + } } // `_lease` drops here, recording elapsed wall-clock into the budget window. } diff --git a/nodedb/src/data/executor/handlers/control/checkpoint_durable_lsn.rs b/nodedb/src/data/executor/handlers/control/checkpoint_durable_lsn.rs index 04b4302d6..63d603939 100644 --- a/nodedb/src/data/executor/handlers/control/checkpoint_durable_lsn.rs +++ b/nodedb/src/data/executor/handlers/control/checkpoint_durable_lsn.rs @@ -879,7 +879,7 @@ mod tests { tid, "docs", nodedb_types::Surrogate::new(7), - "quick brown fox", + &crate::engine::sparse::inverted::test_support::body("quick brown fox"), ) .expect("index a document"); } diff --git a/nodedb/src/data/executor/handlers/fts_sync.rs b/nodedb/src/data/executor/handlers/fts_sync.rs index 8cf417c9d..927cb84b7 100644 --- a/nodedb/src/data/executor/handlers/fts_sync.rs +++ b/nodedb/src/data/executor/handlers/fts_sync.rs @@ -19,8 +19,10 @@ use nodedb_types::Surrogate; use nodedb_types::sync::wire::{AckStatus, SyncProvenance}; impl CoreLoop { - /// Index a document's text into the inverted BM25 index, optionally gating - /// on the `SyncProvenance` for idempotent replay. + /// Index a document's `(field, text)` pairs into the inverted BM25 + /// indexes (whole-document and per field), optionally gating on the + /// `SyncProvenance` for idempotent replay. Empty `fields` remove the + /// document from every index. /// /// Without provenance behaves identically to the pre-gate implementation. /// With provenance: runs the idempotency gate (`sync_admit`) before @@ -32,7 +34,7 @@ impl CoreLoop { tid: u64, collection: &str, surrogate: Surrogate, - text: &str, + fields: &[(String, String)], provenance: Option<&SyncProvenance>, ) -> Response { if let Some(refusal) = @@ -60,9 +62,10 @@ impl CoreLoop { // ── Engine write ──────────────────────────────────────────────────── let tenant_id = nodedb_types::TenantId::new(tid); let database_id = task.request.database_id.as_u64(); + let text = nodedb_fts::DocumentText::from_fields(fields.iter().cloned()); match self .inverted - .index_document(database_id, tenant_id, collection, surrogate, text) + .index_document(database_id, tenant_id, collection, surrogate, &text) { Ok(()) => { // Advance the collection floor for this committed FTS write. @@ -177,15 +180,83 @@ mod tests { assert_eq!(ack.applied_seq, seq); } + /// A message body whose only string field is `body`. + fn body(text: &str) -> Vec<(String, String)> { + vec![("body".to_string(), text.to_string())] + } + /// Two documents under bound surrogates. fn index_documents(core: &mut CoreLoop, task: &ExecutionTask) { for surrogate in [Surrogate::new(5), Surrogate::new(6)] { - let response = - core.execute_fts_index_doc(task, TID, "notes", surrogate, "hello world", None); + let response = core.execute_fts_index_doc( + task, + TID, + "notes", + surrogate, + &body("hello world"), + None, + ); assert_eq!(response.status, Status::Ok); } } + /// Ids of `notes` documents matching `query` in one index. + fn hits(core: &CoreLoop, index: nodedb_fts::IndexScope<'_>, query: &str) -> Vec { + let mut ids: Vec = core + .inverted + .search( + nodedb_types::DatabaseId::DEFAULT.as_u64(), + nodedb_types::TenantId::new(TID), + index, + nodedb_fts::FtsSearchParams { + query, + top_k: 10, + fuzzy_enabled: false, + mode: nodedb_fts::QueryMode::And, + prefilter: None, + }, + ) + .expect("search") + .into_iter() + .map(|r| r.doc_id.as_u32()) + .collect(); + ids.sort_unstable(); + ids + } + + /// A synced document is searchable per field, and a later message with no + /// fields removes it from every index. + #[test] + fn synced_fields_index_per_field_and_empty_fields_remove() { + let dir = tempfile::tempdir().expect("tempdir"); + let (mut core, _req, _resp) = make_core_with_dir(dir.path()); + let task = make_default_task(); + let title = nodedb_fts::IndexScope::field("notes", "title").expect("field scope"); + let body_index = nodedb_fts::IndexScope::field("notes", "body").expect("field scope"); + let doc = vec![ + ("title".to_string(), "rust".to_string()), + ("body".to_string(), "guide".to_string()), + ]; + let response = + core.execute_fts_index_doc(&task, TID, "notes", Surrogate::new(5), &doc, None); + assert_eq!(response.status, Status::Ok); + let response = + core.execute_fts_index_doc(&task, TID, "notes", Surrogate::new(6), &body("rust"), None); + assert_eq!(response.status, Status::Ok); + + assert_eq!(hits(&core, title, "rust"), vec![5]); + assert_eq!(hits(&core, body_index, "rust"), vec![6]); + assert_eq!(hits(&core, "notes".into(), "rust"), vec![5, 6]); + + let response = + core.execute_fts_index_doc(&task, TID, "notes", Surrogate::new(5), &[], Some(&prov(1))); + assert_applied(&response, 1); + assert!(hits(&core, title, "rust").is_empty()); + assert!(hits(&core, body_index, "guide").is_empty()); + assert_eq!(hits(&core, "notes".into(), "rust"), vec![6]); + assert_eq!(doc_count(&core), 1); + } + /// The indexed document count of `notes`. fn doc_count(core: &CoreLoop) -> u32 { core.inverted @@ -205,7 +276,7 @@ mod tests { let task = make_default_task(); let response = - core.execute_fts_index_doc(&task, TID, "notes", Surrogate::ZERO, "hello", None); + core.execute_fts_index_doc(&task, TID, "notes", Surrogate::ZERO, &body("hello"), None); assert!(matches!( response.error_code.as_deref(), Some(crate::bridge::envelope::ErrorCode::RejectedPrevalidation { .. }) diff --git a/nodedb/src/data/executor/handlers/point/update_reindex_text.rs b/nodedb/src/data/executor/handlers/point/update_reindex_text.rs index 82a2c052e..277004208 100644 --- a/nodedb/src/data/executor/handlers/point/update_reindex_text.rs +++ b/nodedb/src/data/executor/handlers/point/update_reindex_text.rs @@ -8,17 +8,18 @@ //! body held and misses the words its new body holds. //! //! The text is extracted exactly as the insert path extracts it -//! (`fts_text::extract_fts_text`), so an updated row indexes the same way a +//! (`fts_text::extract_fts_fields`), so an updated row indexes the same way a //! freshly inserted row with the same body does. `index_document_in_txn` -//! retracts the terms the new text no longer contains, and removes the row -//! from the index when the new text has no indexable word. +//! retracts the terms the new text no longer contains, retracts the row from +//! every field index the new body no longer fills, and removes the row from +//! an index whose new text has no indexable word. use redb::WriteTransaction; use nodedb_types::Surrogate; use crate::data::executor::core_loop::CoreLoop; -use crate::data::executor::fts_text::extract_fts_text; +use crate::data::executor::fts_text::extract_fts_fields; use crate::engine::sparse::inverted::IndexDocScope; use crate::types::TenantId; @@ -41,7 +42,7 @@ impl CoreLoop { txn: &WriteTransaction, p: UpdateTextReindex<'_>, ) -> crate::Result<()> { - let text = extract_fts_text(p.new_doc); + let text = extract_fts_fields(p.new_doc); self.inverted .index_document_in_txn( txn, diff --git a/nodedb/src/data/executor/handlers/snapshot/restore/text.rs b/nodedb/src/data/executor/handlers/snapshot/restore/text.rs index f1097f3d4..1ec71a73b 100644 --- a/nodedb/src/data/executor/handlers/snapshot/restore/text.rs +++ b/nodedb/src/data/executor/handlers/snapshot/restore/text.rs @@ -206,7 +206,9 @@ mod tests { collection: COLL, surrogate, }, - "alpha original", + &crate::data::executor::fts_text::extract_fts_fields( + &serde_json::json!({ "body": "alpha original" }), + ), ) .expect("index old text"); txn.commit().expect("commit old index"); diff --git a/nodedb/src/data/executor/handlers/text_index.rs b/nodedb/src/data/executor/handlers/text_index.rs new file mode 100644 index 000000000..e530221ea --- /dev/null +++ b/nodedb/src/data/executor/handlers/text_index.rs @@ -0,0 +1,99 @@ +// SPDX-License-Identifier: BUSL-1.1 + +//! Resolution of the inverted index a full-text read of one field runs +//! against. Shared by text search, phrase search, score columns, and the +//! text leg of hybrid search. + +use nodedb_fts::IndexScope; +use nodedb_types::text_search::TextColumnFault; + +use crate::data::executor::core_loop::CoreLoop; +use crate::data::executor::fts_text::extract_fts_fields; +use crate::data::executor::handlers::transaction::overlay::Staged; +use crate::data::executor::task::ExecutionTask; +use crate::types::TenantId; + +impl CoreLoop { + /// The index a read of `field` runs against. `None` field reads the + /// whole-document index. + /// + /// `Ok(None)`: the collection holds no text at all, so the read matches + /// nothing. A field that no document holds as text, while other text + /// exists, is `TextColumn { NotIndexed }`, the same as Lite. A column the + /// strict schema declares is never that error: it holds no text yet. + pub(in crate::data::executor) fn text_index<'a>( + &self, + task: &ExecutionTask, + tid: u64, + collection: &'a str, + field: Option<&'a str>, + ) -> crate::Result>> { + let database_id = task.request.database_id; + let tenant = TenantId::new(tid); + let document = IndexScope::document(collection); + let Some(name) = field else { + return Ok(Some(document)); + }; + let column_error = |fault| crate::Error::TextColumn { + collection: collection.to_owned(), + column: name.to_owned(), + fault, + }; + let index = IndexScope::field(collection, name) + .ok_or_else(|| column_error(TextColumnFault::NotAColumn))?; + let declared = self + .strict_schema_for(database_id, tenant, collection) + .is_some_and(|schema| schema.columns.iter().any(|c| c.name == name)); + if declared + || self + .inverted + .has_text(database_id.as_u64(), tenant, index)? + || self.staged_index_text(task, tid, collection, index)? + { + return Ok(Some(index)); + } + if self + .inverted + .has_text(database_id.as_u64(), tenant, document)? + || self.staged_index_text(task, tid, collection, document)? + { + return Err(column_error(TextColumnFault::NotIndexed)); + } + Ok(None) + } + + /// Whether a staged put of the issuing transaction holds text in `index`. + /// A field first written inside the transaction is thereby searchable + /// before COMMIT. + fn staged_index_text( + &self, + task: &ExecutionTask, + tid: u64, + collection: &str, + index: IndexScope<'_>, + ) -> crate::Result { + let Some(txn_id) = task.request.txn_id else { + return Ok(false); + }; + let Some(overlay) = self.txn_overlays.get(&txn_id) else { + return Ok(false); + }; + let config_key = ( + task.request.database_id, + TenantId::new(tid), + collection.to_string(), + ); + for (_, staged) in overlay.iter_for_collection(&config_key) { + let Staged::Put(body) = staged else { + continue; + }; + let Some(doc) = self.decode_indexed_body(&config_key, body)? else { + continue; + }; + if !extract_fts_fields(&doc).text_of(index).is_empty() { + return Ok(true); + } + } + Ok(false) + } +} diff --git a/nodedb/src/data/executor/handlers/transaction/redo_apply/install_refusal_tests.rs b/nodedb/src/data/executor/handlers/transaction/redo_apply/install_refusal_tests.rs index 5b93912fa..cacb3654d 100644 --- a/nodedb/src/data/executor/handlers/transaction/redo_apply/install_refusal_tests.rs +++ b/nodedb/src/data/executor/handlers/transaction/redo_apply/install_refusal_tests.rs @@ -332,7 +332,7 @@ fn fts_documents(refuse: bool) -> u32 { let op = nodedb_physical::physical_plan::TextOp::FtsIndexDoc { collection: QualifiedCollection::new(DatabaseId::DEFAULT, "docs"), surrogate: Surrogate::new(91), - text: "hello world".to_string(), + fields: vec![("body".to_string(), "hello world".to_string())], provenance: None, }; let (record_type, payload) = crate::control::server::wal_dispatch::encode_text_op_record(&op) diff --git a/nodedb/src/data/executor/handlers/transaction/resolve/text.rs b/nodedb/src/data/executor/handlers/transaction/resolve/text.rs index 719e63376..2d2db3278 100644 --- a/nodedb/src/data/executor/handlers/transaction/resolve/text.rs +++ b/nodedb/src/data/executor/handlers/transaction/resolve/text.rs @@ -4,7 +4,7 @@ //! //! **Plan-driven.** An `FtsIndexDoc` / `FtsDeleteDoc` buffered inside a //! transaction carries the complete posting input (collection, surrogate, -//! text), so it resolves to the exact `FtsIndex` / `FtsDelete` record its +//! fields), so it resolves to the exact `FtsIndex` / `FtsDelete` record its //! autocommit form journals (`wal_dispatch::encode_text_op_record`). A row the //! same transaction also writes as a document re-derives the same postings at //! install; an index upsert for one document is idempotent, so the two agree. @@ -37,7 +37,7 @@ mod tests { let op = TextOp::FtsIndexDoc { collection: QualifiedCollection::new(DatabaseId::DEFAULT, "docs"), surrogate: Surrogate::new(7), - text: "hello world".to_string(), + fields: vec![("body".to_string(), "hello world".to_string())], provenance: None, }; let mut ops = Vec::new(); diff --git a/nodedb/src/data/executor/handlers/transaction/undo/columnar_insert.rs b/nodedb/src/data/executor/handlers/transaction/undo/columnar_insert.rs index 96c2f99ce..cd94135e5 100644 --- a/nodedb/src/data/executor/handlers/transaction/undo/columnar_insert.rs +++ b/nodedb/src/data/executor/handlers/transaction/undo/columnar_insert.rs @@ -43,9 +43,10 @@ impl CoreLoop { } } ColumnarInsertIntent::Insert | ColumnarInsertIntent::Put => { - if let Some(prior_loc) = engine.pk_index().get(&pk_bytes).copied() - && prior_loc.segment_id == engine.memtable_segment_id() - { + // The insert tombstones the prior row wherever it lives, + // in the memtable or a flushed segment, so the undo puts + // back either one. + if let Some(prior_loc) = engine.pk_index().get(&pk_bytes).copied() { displaced.push((pk_bytes.clone(), prior_loc)); } inserted_pks.push(pk_bytes); diff --git a/nodedb/src/data/executor/handlers/truncate.rs b/nodedb/src/data/executor/handlers/truncate.rs index 1a36264f4..46d0ed35e 100644 --- a/nodedb/src/data/executor/handlers/truncate.rs +++ b/nodedb/src/data/executor/handlers/truncate.rs @@ -3,10 +3,11 @@ //! TRUNCATE and ESTIMATE_COUNT handlers. use nodedb_physical::physical_plan::{ResolvedSumTarget, StorageMode}; -use tracing::{debug, warn}; +use tracing::debug; use crate::bridge::envelope::{ErrorCode, Response, WriteSetEntry}; use crate::data::executor::core_loop::CoreLoop; +use crate::data::executor::core_loop::fail_stop::FailStopCause; use crate::data::executor::enforcement::materialized_sum::divergence::SumTargetCheck; use crate::data::executor::enforcement::write_hook; use crate::data::executor::handlers::partial_refusal::refusal_after_rows; @@ -61,12 +62,7 @@ impl CoreLoop { ) { Ok(ids) => ids, Err(e) => { - return self.response_error( - task, - ErrorCode::Internal { - detail: format!("scan for truncate: {e}"), - }, - ); + return self.response_error(task, ErrorCode::from(e)); } }; @@ -139,9 +135,10 @@ impl CoreLoop { // `DocumentOp::Truncate`, so these entries are the only record of the // removals WAL replay and a point-in-time restore apply. let mut write_set: Vec = Vec::new(); + // Surrogates removed so far. A refusal part-way removes their text + // from the inverted index. A full TRUNCATE empties it in one purge. + let mut removed: Vec = Vec::new(); for storage_key in &all_ids { - let doc_id = storage_key.to_string(); - // One transaction per removed row, shared with the materialized-sum // delta that row owes — identical to `execute_bulk_delete`, so a // TRUNCATE and a `DELETE` with no predicate leave the same totals. @@ -151,14 +148,37 @@ impl CoreLoop { Ok(txn) => txn, Err(e) => { let code = refusal_after_rows(truncated, e); - return self.refusal_with_landed_rows(task, code, write_set); + return self.truncate_refusal(task, tid, collection, code, &removed, write_set); } }; - let deleted_bytes = self - .sparse - .delete_in_txn(&row_txn, database_id, tid, collection, storage_key) - .ok() - .flatten(); + // A delete error refuses the TRUNCATE: read as an absent row, it + // would leave the row stored under a success. + let deleted_bytes = + match self + .sparse + .delete_in_txn(&row_txn, database_id, tid, collection, storage_key) + { + Ok(bytes) => bytes, + Err(e) => { + let code = refusal_after_rows(truncated, e); + return self + .truncate_refusal(task, tid, collection, code, &removed, write_set); + } + }; + // The row's secondary-index entries leave in the row's own + // transaction, so the row and its entries go together or not at all. + if deleted_bytes.is_some() + && let Err(e) = self.sparse.delete_indexes_for_document_in_txn( + &row_txn, + database_id, + tid, + collection, + storage_key, + ) + { + let code = refusal_after_rows(truncated, e); + return self.truncate_refusal(task, tid, collection, code, &removed, write_set); + } let mut target_writes = Vec::new(); if let Some(bytes) = deleted_bytes.as_deref() { match write_hook::run( @@ -182,18 +202,20 @@ impl CoreLoop { Ok(outcome) => target_writes = outcome.target_writes, Err(e) => { let code = refusal_after_rows(truncated, e); - return self.refusal_with_landed_rows(task, code, write_set); + return self + .truncate_refusal(task, tid, collection, code, &removed, write_set); } } } if let Err(e) = row_txn.commit() { let code = refusal_after_rows( truncated, - ErrorCode::Internal { + crate::Error::Storage { + engine: "sparse".into(), detail: format!("truncate commit: {e}"), }, ); - return self.refusal_with_landed_rows(task, code, write_set); + return self.truncate_refusal(task, tid, collection, code, &removed, write_set); } if let Some(deleted_bytes) = deleted_bytes.as_deref() { let surrogate = storage_key.surrogate(); @@ -207,22 +229,9 @@ impl CoreLoop { declared_primary_key, *storage_key, ); - if let Err(e) = self.inverted.remove_document( - database_id, - crate::types::TenantId::new(tid), - collection, - surrogate, - ) { - warn!(core = self.core_id, %collection, %doc_id, error = %e, "truncate: inverted removal failed"); - } - if let Err(e) = self.sparse.delete_indexes_for_document( - database_id, - tid, - collection, - storage_key, - ) { - warn!(core = self.core_id, %collection, %doc_id, error = %e, "truncate: index cascade failed"); - } + // The row's text leaves the inverted index with every other + // row's, in one purge once the loop ends. + removed.push(surrogate); // Cascade: secondary HNSW vector index. The put path indexed // this row's vectors under its surrogate; truncate must // soft-delete those nodes and drop the reverse-map entry, or @@ -255,23 +264,29 @@ impl CoreLoop { // events were lost, and per-row events are what ROW-level // AFTER-DELETE triggers match on (see // `event::trigger::dispatcher::single`). - let old_converted = self.resolve_event_payload( - task.request.database_id.as_u64(), - tid, - collection, - deleted_bytes, - ); self.emit_document_delete_event( task, + tid, collection, row_identity, - Some(old_converted.as_deref().unwrap_or(deleted_bytes)), + Some(deleted_bytes), ); truncated += 1; } write_set.extend(write_hook::target_write_set(&target_writes)); } + // Every row is removed: empty the collection's inverted index in one + // purge. Its analyzer, language, and fuzzy configuration stay. + if let Err(e) = self.inverted.clear_collection( + database_id, + crate::types::TenantId::new(tid), + collection, + ) { + let code = refusal_after_rows(truncated, ErrorCode::from(e)); + return self.refusal_with_landed_rows(task, code, write_set); + } + // Clear aggregate cache for this collection. self.invalidate_aggregate_cache_for_collection( task.request.database_id.as_u64(), @@ -285,17 +300,47 @@ impl CoreLoop { // removals' entries as well. let mut response = match response_codec::encode_json_as_msgpack(&result) { Ok(payload) => self.response_with_payload(task, payload), - Err(e) => self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ), + Err(e) => self.response_error(task, ErrorCode::from(e)), }; response.write_set = write_set; response } + /// The refusal of a TRUNCATE that stopped part-way. The rows removed so + /// far stay removed, so their text leaves the inverted index in one batch + /// before the refusal answers. The refusal keeps `code`, so the client + /// sees the SQLSTATE of what stopped the TRUNCATE. + /// + /// A failure of the text removal leaves the inverted index holding text + /// of removed rows. That cannot be undone here, so the core fail-stops. + /// Restart replay applies the journalled removals and rebuilds the index. + fn truncate_refusal( + &mut self, + task: &ExecutionTask, + tid: u64, + collection: &str, + code: ErrorCode, + removed: &[nodedb_types::Surrogate], + write_set: Vec, + ) -> Response { + if let Err(e) = self.inverted.remove_documents( + task.request.database_id.as_u64(), + crate::types::TenantId::new(tid), + collection, + removed, + ) { + self.fail_stop_core( + FailStopCause::PostInstallFailed, + &format!( + "TRUNCATE of '{collection}' stopped after {} rows with {code:?}, then \ + removing those rows' text from the inverted index failed: {e}", + removed.len() + ), + ); + } + self.refusal_with_landed_rows(task, code, write_set) + } + /// ESTIMATE_COUNT: return approximate row count from HLL cardinality stats. pub(in crate::data::executor) fn execute_estimate_count( &mut self, @@ -318,12 +363,7 @@ impl CoreLoop { }); match response_codec::encode_json_as_msgpack(&result) { Ok(payload) => self.response_with_payload(task, payload), - Err(e) => self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ), + Err(e) => self.response_error(task, ErrorCode::from(e)), } } Ok(None) => { @@ -336,20 +376,10 @@ impl CoreLoop { }); match response_codec::encode_json_as_msgpack(&result) { Ok(payload) => self.response_with_payload(task, payload), - Err(e) => self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ), + Err(e) => self.response_error(task, ErrorCode::from(e)), } } - Err(e) => self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ), + Err(e) => self.response_error(task, ErrorCode::from(e)), } } } diff --git a/nodedb/src/data/executor/wal_replay_fts.rs b/nodedb/src/data/executor/wal_replay_fts.rs index 7770b8df3..6dc434f35 100644 --- a/nodedb/src/data/executor/wal_replay_fts.rs +++ b/nodedb/src/data/executor/wal_replay_fts.rs @@ -212,7 +212,7 @@ impl CoreLoop { payload.collection.clone(), ), surrogate, - text: payload.text.clone(), + fields: payload.fields.clone(), provenance: Some(prov.clone()), }), ); @@ -222,7 +222,7 @@ impl CoreLoop { tenant_id, &payload.collection, surrogate, - &payload.text, + &payload.fields, Some(&prov), ); @@ -437,11 +437,15 @@ mod tests { } fn index_record(lsn: u64, text: &str) -> nodedb_wal::WalRecord { + fields_record(lsn, vec![("body".to_string(), text.to_string())]) + } + + fn fields_record(lsn: u64, fields: Vec<(String, String)>) -> nodedb_wal::WalRecord { let payload = FtsIndexPayload::new( local_provenance(), COLLECTION, format!("{SURROGATE:08x}"), - text, + fields, ) .to_bytes() .expect("encode FtsIndexPayload"); @@ -569,4 +573,42 @@ mod tests { "re-applying a durable FtsDelete record must be a no-op" ); } + + /// An `FtsIndex` record with several fields replays into each field's + /// index and into the whole-document index. + #[test] + fn an_index_record_replays_into_its_field_scopes() { + let mut h = make_core(); + replay( + &mut h.core, + &fields_record( + 10, + vec![ + ("title".to_string(), "alpha".to_string()), + ("body".to_string(), "bravo".to_string()), + ], + ), + ); + let tid = TenantId::new(TENANT); + let title = nodedb_fts::IndexScope::field(COLLECTION, "title").expect("field scope"); + let body = nodedb_fts::IndexScope::field(COLLECTION, "body").expect("field scope"); + let df = |index: nodedb_fts::IndexScope<'_>, word: &str| { + let term = h + .core + .inverted + .analyze_for_collection(DB, tid, COLLECTION, word) + .expect("analyze") + .remove(0); + h.core + .inverted + .term_df(DB, tid, index, &term) + .expect("term df") + }; + assert_eq!(df(title, "alpha"), 1); + assert_eq!(df(title, "bravo"), 0); + assert_eq!(df(body, "bravo"), 1); + assert_eq!(df(body, "alpha"), 0); + assert_eq!(df(COLLECTION.into(), "alpha"), 1); + assert_eq!(df(COLLECTION.into(), "bravo"), 1); + } } diff --git a/nodedb/src/engine/sparse/fts_redb/backend/core.rs b/nodedb/src/engine/sparse/fts_redb/backend/core.rs index 81fd7c4fc..e74ca2195 100644 --- a/nodedb/src/engine/sparse/fts_redb/backend/core.rs +++ b/nodedb/src/engine/sparse/fts_redb/backend/core.rs @@ -8,6 +8,7 @@ use std::sync::Arc; +use nodedb_fts::IndexScope; use nodedb_fts::backend::FtsBackend; use nodedb_fts::posting::Posting; use nodedb_types::Surrogate; @@ -16,15 +17,16 @@ use super::segments::CompactCommit; use super::shared::redb_err; use crate::engine::durability_gate::GatedDatabase; use crate::engine::sparse::fts_redb::tables::{ - DOC_LENGTHS, DOC_TERMS, INDEX_META, POSTINGS, SEGMENTS, STATS, + DOC_FIELDS, DOC_LENGTHS, DOC_TERMS, INDEX_META, POSTINGS, SEGMENTS, STATS, }; use crate::storage::quarantine::QuarantineRegistry; /// Redb-backed FTS backend. /// /// All persistent tables are keyed by the structural tuple -/// `(database_id, tenant_id, collection, …)` — database and tenant isolation -/// are enforced by the table schema, never by lexical-prefix ordering. +/// `(database_id, tenant_id, collection, field, …)` — database, tenant, and +/// index isolation are enforced by the table schema, never by lexical-prefix +/// ordering. pub struct RedbFtsBackend { /// The sparse engine's gated database, shared. pub(super) db: Arc, @@ -47,6 +49,9 @@ impl RedbFtsBackend { write_txn .open_table(DOC_TERMS) .map_err(|e| redb_err("create doc_terms table", e))?; + write_txn + .open_table(DOC_FIELDS) + .map_err(|e| redb_err("create doc_fields table", e))?; write_txn .open_table(INDEX_META) .map_err(|e| redb_err("create index_meta table", e))?; @@ -85,11 +90,32 @@ impl RedbFtsBackend { super::segments::compact_commit(self, params) } - /// Enumerate all `(database_id, tid, collection)` triples that have at least - /// one FTS segment. Used by maintenance to discover compaction candidates - /// without a separate registry. - pub fn list_all_fts_collections(&self) -> crate::Result> { - super::segments::list_all_collections(self) + /// Drop every indexed row of one collection in one write transaction, + /// keeping its analyzer, language, and fuzzy configuration. + pub fn clear_collection_data( + &self, + database_id: u64, + tid: u64, + collection: &str, + ) -> crate::Result { + super::purge::collection_data(self, database_id, tid, collection) + } + + /// Every document one index holds, read in one transaction. + pub fn index_members( + &self, + database_id: u64, + tid: u64, + index: IndexScope<'_>, + ) -> crate::Result> { + super::doc_lengths::members(self, database_id, tid, index) + } + + /// Enumerate every `(database_id, tid, collection, field)` index that has + /// at least one FTS segment. Used by maintenance to discover compaction + /// candidates without a separate registry. + pub fn list_all_fts_indexes(&self) -> crate::Result> { + super::segments::list_all_indexes(self) } } @@ -100,161 +126,171 @@ impl FtsBackend for RedbFtsBackend { &self, database_id: u64, tid: u64, - collection: &str, + index: IndexScope<'_>, term: &str, ) -> crate::Result> { - super::postings::read(self, database_id, tid, collection, term) + super::postings::read(self, database_id, tid, index, term) } fn write_postings( &self, database_id: u64, tid: u64, - collection: &str, + index: IndexScope<'_>, term: &str, postings: &[Posting], ) -> crate::Result<()> { - super::postings::write(self, database_id, tid, collection, term, postings) + super::postings::write(self, database_id, tid, index, term, postings) } fn remove_postings( &self, database_id: u64, tid: u64, - collection: &str, + index: IndexScope<'_>, term: &str, ) -> crate::Result<()> { - super::postings::remove(self, database_id, tid, collection, term) + super::postings::remove(self, database_id, tid, index, term) } fn read_doc_length( &self, database_id: u64, tid: u64, - collection: &str, + index: IndexScope<'_>, doc_id: Surrogate, ) -> crate::Result> { - super::doc_lengths::read(self, database_id, tid, collection, doc_id) + super::doc_lengths::read(self, database_id, tid, index, doc_id) + } + + fn read_doc_lengths( + &self, + database_id: u64, + tid: u64, + index: IndexScope<'_>, + doc_ids: &[Surrogate], + ) -> crate::Result>> { + super::doc_lengths::read_many(self, database_id, tid, index, doc_ids) } fn write_doc_length( &self, database_id: u64, tid: u64, - collection: &str, + index: IndexScope<'_>, doc_id: Surrogate, length: u32, ) -> crate::Result<()> { - super::doc_lengths::write(self, database_id, tid, collection, doc_id, length) + super::doc_lengths::write(self, database_id, tid, index, doc_id, length) } fn remove_doc_length( &self, database_id: u64, tid: u64, - collection: &str, + index: IndexScope<'_>, doc_id: Surrogate, ) -> crate::Result<()> { - super::doc_lengths::remove(self, database_id, tid, collection, doc_id) + super::doc_lengths::remove(self, database_id, tid, index, doc_id) } fn collection_terms( &self, database_id: u64, tid: u64, - collection: &str, + index: IndexScope<'_>, ) -> crate::Result> { - super::postings::collection_terms(self, database_id, tid, collection) + super::postings::collection_terms(self, database_id, tid, index) } fn collection_stats( &self, database_id: u64, tid: u64, - collection: &str, + index: IndexScope<'_>, ) -> crate::Result<(u32, u64)> { - super::stats::read(self, database_id, tid, collection) + super::stats::read(self, database_id, tid, index) } fn increment_stats( &self, database_id: u64, tid: u64, - collection: &str, + index: IndexScope<'_>, doc_len: u32, ) -> crate::Result<()> { - super::stats::increment(self, database_id, tid, collection, doc_len) + super::stats::increment(self, database_id, tid, index, doc_len) } fn decrement_stats( &self, database_id: u64, tid: u64, - collection: &str, + index: IndexScope<'_>, doc_len: u32, ) -> crate::Result<()> { - super::stats::decrement(self, database_id, tid, collection, doc_len) + super::stats::decrement(self, database_id, tid, index, doc_len) } fn read_meta( &self, database_id: u64, tid: u64, - collection: &str, + index: IndexScope<'_>, subkey: &str, ) -> crate::Result>> { - super::meta::read(self, database_id, tid, collection, subkey) + super::meta::read(self, database_id, tid, index, subkey) } fn write_meta( &self, database_id: u64, tid: u64, - collection: &str, + index: IndexScope<'_>, subkey: &str, value: &[u8], ) -> crate::Result<()> { - super::meta::write(self, database_id, tid, collection, subkey, value) + super::meta::write(self, database_id, tid, index, subkey, value) } fn write_segment( &self, database_id: u64, tid: u64, - collection: &str, + index: IndexScope<'_>, segment_id: &str, data: &[u8], ) -> crate::Result<()> { - super::segments::write(self, database_id, tid, collection, segment_id, data) + super::segments::write(self, database_id, tid, index, segment_id, data) } fn read_segment( &self, database_id: u64, tid: u64, - collection: &str, + index: IndexScope<'_>, segment_id: &str, ) -> crate::Result>> { - super::segments::read(self, database_id, tid, collection, segment_id) + super::segments::read(self, database_id, tid, index, segment_id) } fn list_segments( &self, database_id: u64, tid: u64, - collection: &str, + index: IndexScope<'_>, ) -> crate::Result> { - super::segments::list(self, database_id, tid, collection) + super::segments::list(self, database_id, tid, index) } fn remove_segment( &self, database_id: u64, tid: u64, - collection: &str, + index: IndexScope<'_>, segment_id: &str, ) -> crate::Result<()> { - super::segments::remove(self, database_id, tid, collection, segment_id) + super::segments::remove(self, database_id, tid, index, segment_id) } fn purge_collection( diff --git a/nodedb/src/engine/sparse/fts_redb/backend/doc_lengths.rs b/nodedb/src/engine/sparse/fts_redb/backend/doc_lengths.rs index 6277b25b3..4e4cacb67 100644 --- a/nodedb/src/engine/sparse/fts_redb/backend/doc_lengths.rs +++ b/nodedb/src/engine/sparse/fts_redb/backend/doc_lengths.rs @@ -1,20 +1,24 @@ // SPDX-License-Identifier: BUSL-1.1 //! Per-document token-length operations against `DOC_LENGTHS` -//! keyed by `(database_id, tenant_id, collection, surrogate_u32)`. +//! keyed by `(database_id, tenant_id, collection, field, surrogate_u32)`. +use std::collections::HashSet; + +use nodedb_fts::IndexScope; use nodedb_types::Surrogate; use redb::{ReadableDatabase, ReadableTable}; use super::core::RedbFtsBackend; use super::shared::redb_err; +use crate::engine::sparse::fts_redb::keys::doc_key; use crate::engine::sparse::fts_redb::tables::DOC_LENGTHS; pub(super) fn read( backend: &RedbFtsBackend, database_id: u64, tid: u64, - collection: &str, + index: IndexScope<'_>, doc_id: Surrogate, ) -> crate::Result> { let read_txn = backend @@ -24,7 +28,7 @@ pub(super) fn read( let table = read_txn .open_table(DOC_LENGTHS) .map_err(|e| redb_err("open doc_lengths", e))?; - match table.get((database_id, tid, collection, doc_id.as_u32())) { + match table.get(doc_key(database_id, tid, index, doc_id)) { Ok(Some(val)) => { let len: u32 = zerompk::from_msgpack(val.value()) .map_err(|e| redb_err("deserialize doc_length", e))?; @@ -35,11 +39,76 @@ pub(super) fn read( } } +/// The lengths of `doc_ids` in one index, parallel to `doc_ids`, read in one +/// transaction. +pub(super) fn read_many( + backend: &RedbFtsBackend, + database_id: u64, + tid: u64, + index: IndexScope<'_>, + doc_ids: &[Surrogate], +) -> crate::Result>> { + if doc_ids.is_empty() { + return Ok(Vec::new()); + } + let read_txn = backend + .db + .begin_read() + .map_err(|e| redb_err("read txn", e))?; + let table = read_txn + .open_table(DOC_LENGTHS) + .map_err(|e| redb_err("open doc_lengths", e))?; + let mut lengths = Vec::with_capacity(doc_ids.len()); + for doc_id in doc_ids { + let length = match table + .get(doc_key(database_id, tid, index, *doc_id)) + .map_err(|e| redb_err("get doc_length", e))? + { + Some(val) => Some( + zerompk::from_msgpack::(val.value()) + .map_err(|e| redb_err("deserialize doc_length", e))?, + ), + None => None, + }; + lengths.push(length); + } + Ok(lengths) +} + +/// Every surrogate one index records a length for, read in one transaction. +/// +/// An index holds a document exactly when it records the document's length. +pub(super) fn members( + backend: &RedbFtsBackend, + database_id: u64, + tid: u64, + index: IndexScope<'_>, +) -> crate::Result> { + let read_txn = backend + .db + .begin_read() + .map_err(|e| redb_err("read txn", e))?; + let table = read_txn + .open_table(DOC_LENGTHS) + .map_err(|e| redb_err("open doc_lengths", e))?; + let first = doc_key(database_id, tid, index, Surrogate::new(0)); + let last = doc_key(database_id, tid, index, Surrogate::new(u32::MAX)); + let mut members = HashSet::new(); + for entry in table + .range(first..=last) + .map_err(|e| redb_err("range doc_lengths", e))? + { + let (key, _) = entry.map_err(|e| redb_err("read doc_length", e))?; + members.insert(Surrogate::new(key.value().4)); + } + Ok(members) +} + pub(super) fn write( backend: &RedbFtsBackend, database_id: u64, tid: u64, - collection: &str, + index: IndexScope<'_>, doc_id: Surrogate, length: u32, ) -> crate::Result<()> { @@ -54,10 +123,7 @@ pub(super) fn write( let bytes = zerompk::to_msgpack_vec(&length).map_err(|e| redb_err("serialize doc_len", e))?; table - .insert( - (database_id, tid, collection, doc_id.as_u32()), - bytes.as_slice(), - ) + .insert(doc_key(database_id, tid, index, doc_id), bytes.as_slice()) .map_err(|e| redb_err("insert doc_len", e))?; } write_txn.commit().map_err(|e| redb_err("commit", e))?; @@ -68,7 +134,7 @@ pub(super) fn remove( backend: &RedbFtsBackend, database_id: u64, tid: u64, - collection: &str, + index: IndexScope<'_>, doc_id: Surrogate, ) -> crate::Result<()> { let write_txn = backend @@ -79,7 +145,9 @@ pub(super) fn remove( let mut table = write_txn .open_table(DOC_LENGTHS) .map_err(|e| redb_err("open doc_lengths", e))?; - let _ = table.remove((database_id, tid, collection, doc_id.as_u32())); + table + .remove(doc_key(database_id, tid, index, doc_id)) + .map_err(|e| redb_err("remove doc_len", e))?; } write_txn.commit().map_err(|e| redb_err("commit", e))?; Ok(()) diff --git a/nodedb/src/engine/sparse/fts_redb/backend/meta.rs b/nodedb/src/engine/sparse/fts_redb/backend/meta.rs index 389592293..b45ae5da8 100644 --- a/nodedb/src/engine/sparse/fts_redb/backend/meta.rs +++ b/nodedb/src/engine/sparse/fts_redb/backend/meta.rs @@ -1,18 +1,20 @@ // SPDX-License-Identifier: BUSL-1.1 -//! Opaque metadata blobs (docmap, fieldnorms, analyzer, language) -//! against `INDEX_META` keyed by `(database_id, tenant_id, collection, subkey)`. +//! Opaque metadata blobs (fieldnorms, analyzer, language, fuzzy) against +//! `INDEX_META` keyed by `(database_id, tenant_id, collection, field, subkey)`. use super::core::RedbFtsBackend; use super::shared::redb_err; +use crate::engine::sparse::fts_redb::keys::meta_key; use crate::engine::sparse::fts_redb::tables::INDEX_META; +use nodedb_fts::IndexScope; use redb::{ReadableDatabase, ReadableTable}; pub(super) fn read( backend: &RedbFtsBackend, database_id: u64, tid: u64, - collection: &str, + index: IndexScope<'_>, subkey: &str, ) -> crate::Result>> { let read_txn = backend @@ -22,7 +24,7 @@ pub(super) fn read( let table = read_txn .open_table(INDEX_META) .map_err(|e| redb_err("open index_meta", e))?; - match table.get((database_id, tid, collection, subkey)) { + match table.get(meta_key(database_id, tid, index, subkey)) { Ok(Some(val)) => Ok(Some(val.value().to_vec())), Ok(None) => Ok(None), Err(e) => Err(redb_err("get meta", e)), @@ -33,7 +35,7 @@ pub(super) fn write( backend: &RedbFtsBackend, database_id: u64, tid: u64, - collection: &str, + index: IndexScope<'_>, subkey: &str, value: &[u8], ) -> crate::Result<()> { @@ -46,7 +48,7 @@ pub(super) fn write( .open_table(INDEX_META) .map_err(|e| redb_err("open index_meta", e))?; table - .insert((database_id, tid, collection, subkey), value) + .insert(meta_key(database_id, tid, index, subkey), value) .map_err(|e| redb_err("insert meta", e))?; } write_txn.commit().map_err(|e| redb_err("commit", e))?; diff --git a/nodedb/src/engine/sparse/fts_redb/backend/postings.rs b/nodedb/src/engine/sparse/fts_redb/backend/postings.rs index e5d67b789..251453f70 100644 --- a/nodedb/src/engine/sparse/fts_redb/backend/postings.rs +++ b/nodedb/src/engine/sparse/fts_redb/backend/postings.rs @@ -1,20 +1,23 @@ // SPDX-License-Identifier: BUSL-1.1 //! Posting-list operations against the `POSTINGS` table -//! keyed by `(database_id, tenant_id, collection, term)`. +//! keyed by `(database_id, tenant_id, collection, field, term)`. +use nodedb_fts::IndexScope; use nodedb_fts::posting::Posting; use redb::{ReadableDatabase, ReadableTable}; use super::core::RedbFtsBackend; -use super::shared::{MAX_SUBKEY, redb_err}; +use super::shared::redb_err; +use crate::engine::sparse::fts_redb::keys::{KeyOwner, posting_key}; +use crate::engine::sparse::fts_redb::scan; use crate::engine::sparse::fts_redb::tables::POSTINGS; pub(super) fn read( backend: &RedbFtsBackend, database_id: u64, tid: u64, - collection: &str, + index: IndexScope<'_>, term: &str, ) -> crate::Result> { let read_txn = backend @@ -24,7 +27,7 @@ pub(super) fn read( let table = read_txn .open_table(POSTINGS) .map_err(|e| redb_err("open postings", e))?; - match table.get((database_id, tid, collection, term)) { + match table.get(posting_key(database_id, tid, index, term)) { Ok(Some(val)) => { let list: Vec = zerompk::from_msgpack(val.value()) .map_err(|e| redb_err("deserialize postings", e))?; @@ -39,7 +42,7 @@ pub(super) fn write( backend: &RedbFtsBackend, database_id: u64, tid: u64, - collection: &str, + index: IndexScope<'_>, term: &str, postings: &[Posting], ) -> crate::Result<()> { @@ -51,13 +54,16 @@ pub(super) fn write( let mut table = write_txn .open_table(POSTINGS) .map_err(|e| redb_err("open postings", e))?; + let key = posting_key(database_id, tid, index, term); if postings.is_empty() { - let _ = table.remove((database_id, tid, collection, term)); + table + .remove(key) + .map_err(|e| redb_err("remove posting", e))?; } else { let bytes = zerompk::to_msgpack_vec(&postings.to_vec()) .map_err(|e| redb_err("serialize postings", e))?; table - .insert((database_id, tid, collection, term), bytes.as_slice()) + .insert(key, bytes.as_slice()) .map_err(|e| redb_err("insert posting", e))?; } } @@ -69,7 +75,7 @@ pub(super) fn remove( backend: &RedbFtsBackend, database_id: u64, tid: u64, - collection: &str, + index: IndexScope<'_>, term: &str, ) -> crate::Result<()> { let write_txn = backend @@ -80,7 +86,9 @@ pub(super) fn remove( let mut table = write_txn .open_table(POSTINGS) .map_err(|e| redb_err("open postings", e))?; - let _ = table.remove((database_id, tid, collection, term)); + table + .remove(posting_key(database_id, tid, index, term)) + .map_err(|e| redb_err("remove posting", e))?; } write_txn.commit().map_err(|e| redb_err("commit", e))?; Ok(()) @@ -90,7 +98,7 @@ pub(super) fn collection_terms( backend: &RedbFtsBackend, database_id: u64, tid: u64, - collection: &str, + index: IndexScope<'_>, ) -> crate::Result> { let read_txn = backend .db @@ -99,11 +107,7 @@ pub(super) fn collection_terms( let table = read_txn .open_table(POSTINGS) .map_err(|e| redb_err("open postings", e))?; - - let terms: Vec = table - .range((database_id, tid, collection, "")..=(database_id, tid, collection, MAX_SUBKEY)) - .map_err(|e| redb_err("range", e))? - .filter_map(|r| r.ok().map(|(k, _)| k.value().3.to_string())) - .collect(); - Ok(terms) + let keys = scan::str_keys(&table, KeyOwner::index(database_id, tid, index)) + .map_err(|e| redb_err("postings range", e))?; + Ok(keys.into_iter().map(|(_, _, term)| term).collect()) } diff --git a/nodedb/src/engine/sparse/fts_redb/backend/purge.rs b/nodedb/src/engine/sparse/fts_redb/backend/purge.rs index 06fb9ff4b..1a4e6834c 100644 --- a/nodedb/src/engine/sparse/fts_redb/backend/purge.rs +++ b/nodedb/src/engine/sparse/fts_redb/backend/purge.rs @@ -2,27 +2,93 @@ //! Structural collection and tenant drops. //! -//! Each table is scanned by tuple range and every matching entry is -//! removed; no lexical-prefix scans appear here. `purge_tenant` performs -//! the drop in a single write transaction so the teardown is atomic +//! Each table is scanned from the owner's lower bound and every owned entry +//! is removed: every index of the collection (whole-document and field), +//! plus the per-document field sets. No lexical-prefix scans appear here. +//! Each drop runs in a single write transaction so the teardown is atomic //! across tables. //! //! A tenant purge is scoped to a single `(database_id, tenant_id)` pair — //! the same tenant in a different database is unaffected. -use redb::ReadableTable; +use nodedb_fts::IndexScope; +use nodedb_types::Surrogate; +use redb::{TableDefinition, WriteTransaction}; use super::core::RedbFtsBackend; -use super::shared::{MAX_COLLECTION, MAX_SUBKEY, redb_err}; +use super::shared::redb_err; +use crate::engine::sparse::fts_redb::keys::{ + KeyOwner, doc_fields_key, doc_key, posting_key, stats_key, +}; +use crate::engine::sparse::fts_redb::scan::{self, DocKeyDef, StrKeyDef}; use crate::engine::sparse::fts_redb::tables::{ - DOC_LENGTHS, DOC_TERMS, INDEX_META, POSTINGS, SEGMENTS, STATS, + DOC_FIELDS, DOC_LENGTHS, DOC_TERMS, INDEX_META, POSTINGS, SEGMENTS, STATS, }; +/// Index metadata sub-keys that configure a collection rather than hold its +/// indexed data: they survive a data-only purge. +const CONFIG_META_KEYS: [&str; 3] = ["analyzer", "language", "fuzzy"]; + +/// Which metadata rows a purge drops. +#[derive(Clone, Copy, PartialEq, Eq)] +enum MetaRows { + /// Every metadata row: the collection is gone. + All, + /// Data rows only. The collection's analyzer, language, and fuzzy + /// configuration stay: the collection is emptied, not dropped. + KeepConfig, +} + pub(super) fn collection( backend: &RedbFtsBackend, database_id: u64, tid: u64, coll: &str, +) -> crate::Result { + purge( + backend, + database_id, + tid, + KeyOwner::collection(database_id, tid, coll), + MetaRows::All, + ) +} + +/// Drop every indexed row of one collection, keeping its configuration. +pub(super) fn collection_data( + backend: &RedbFtsBackend, + database_id: u64, + tid: u64, + coll: &str, +) -> crate::Result { + purge( + backend, + database_id, + tid, + KeyOwner::collection(database_id, tid, coll), + MetaRows::KeepConfig, + ) +} + +pub(super) fn tenant(backend: &RedbFtsBackend, database_id: u64, tid: u64) -> crate::Result { + purge( + backend, + database_id, + tid, + KeyOwner::tenant(database_id, tid), + MetaRows::All, + ) +} + +/// Drop every row `owner` covers from every FTS table, with metadata rows +/// dropped per `meta`. Returns the number of posting lists, document-length +/// rows, and segments removed. +fn purge( + backend: &RedbFtsBackend, + database_id: u64, + tid: u64, + owner: KeyOwner<'_>, + meta: MetaRows, ) -> crate::Result { let write_txn = backend .db @@ -30,180 +96,118 @@ pub(super) fn collection( .map_err(|e| redb_err("purge write txn", e))?; let mut removed = 0; - { - let mut table = write_txn - .open_table(POSTINGS) - .map_err(|e| redb_err("open postings", e))?; - let keys: Vec = table - .range((database_id, tid, coll, "")..=(database_id, tid, coll, MAX_SUBKEY)) - .map_err(|e| redb_err("postings range", e))? - .filter_map(|r| r.ok().map(|(k, _)| k.value().3.to_string())) - .collect(); - removed += keys.len(); - for term in &keys { - let _ = table.remove((database_id, tid, coll, term.as_str())); - } - } - - // DOC_LENGTHS and DOC_TERMS are keyed by (u64, u64, &str, u32) — use - // numeric range bounds. Both are per-document rows for this collection and + removed += drop_str_rows(&write_txn, POSTINGS, database_id, tid, owner, "postings")?; + // DOC_LENGTHS and DOC_TERMS are per-document rows of the same indexes and // must go together: a surviving term set would name posting lists that no // longer exist. - removed += drop_surrogate_quad_collection(&write_txn, DOC_LENGTHS, database_id, tid, coll)?; - removed += drop_surrogate_quad_collection(&write_txn, DOC_TERMS, database_id, tid, coll)?; - - { - let mut table = write_txn - .open_table(INDEX_META) - .map_err(|e| redb_err("open index_meta", e))?; - let keys: Vec = table - .range((database_id, tid, coll, "")..=(database_id, tid, coll, MAX_SUBKEY)) - .map_err(|e| redb_err("meta range", e))? - .filter_map(|r| r.ok().map(|(k, _)| k.value().3.to_string())) - .collect(); - for sub in &keys { - let _ = table.remove((database_id, tid, coll, sub.as_str())); - } - } + removed += drop_doc_rows( + &write_txn, + DOC_LENGTHS, + database_id, + tid, + owner, + "doc_lengths", + )?; + drop_doc_rows(&write_txn, DOC_TERMS, database_id, tid, owner, "doc_terms")?; + drop_meta_rows(&write_txn, database_id, tid, owner, meta)?; + removed += drop_str_rows(&write_txn, SEGMENTS, database_id, tid, owner, "segments")?; { - let mut table = write_txn + let mut stats = write_txn .open_table(STATS) .map_err(|e| redb_err("open stats", e))?; - let _ = table.remove((database_id, tid, coll)); - } - - { - let mut table = write_txn - .open_table(SEGMENTS) - .map_err(|e| redb_err("open segments", e))?; - let ids: Vec = table - .range((database_id, tid, coll, "")..=(database_id, tid, coll, MAX_SUBKEY)) - .map_err(|e| redb_err("segments range", e))? - .filter_map(|r| r.ok().map(|(k, _)| k.value().3.to_string())) - .collect(); - removed += ids.len(); - for id in &ids { - let _ = table.remove((database_id, tid, coll, id.as_str())); + let keys = scan::stats_keys(&stats, owner).map_err(|e| redb_err("stats range", e))?; + for (c, f) in &keys { + stats + .remove(stats_key(database_id, tid, IndexScope::from_key(c, f))) + .map_err(|e| redb_err("remove stats", e))?; } } - - write_txn - .commit() - .map_err(|e| redb_err("commit purge", e))?; - Ok(removed) -} - -pub(super) fn tenant(backend: &RedbFtsBackend, database_id: u64, tid: u64) -> crate::Result { - let write_txn = backend - .db - .begin_write() - .map_err(|e| redb_err("purge_tenant write txn", e))?; - let mut removed = 0; - - removed += drop_str_quad_range(&write_txn, POSTINGS, database_id, tid)?; - removed += drop_surrogate_quad_tenant(&write_txn, DOC_LENGTHS, database_id, tid)?; - removed += drop_surrogate_quad_tenant(&write_txn, DOC_TERMS, database_id, tid)?; - let _ = drop_str_quad_range(&write_txn, INDEX_META, database_id, tid)?; - { - let mut stats = write_txn - .open_table(STATS) - .map_err(|e| redb_err("open stats", e))?; - let colls: Vec = stats - .range((database_id, tid, "")..=(database_id, tid, MAX_COLLECTION)) - .map_err(|e| redb_err("stats range", e))? - .filter_map(|r| r.ok().map(|(k, _)| k.value().2.to_string())) - .collect(); - for c in &colls { - let _ = stats.remove((database_id, tid, c.as_str())); + let mut fields = write_txn + .open_table(DOC_FIELDS) + .map_err(|e| redb_err("open doc_fields", e))?; + let keys = + scan::doc_fields_keys(&fields, owner).map_err(|e| redb_err("doc_fields range", e))?; + for (c, s) in &keys { + fields + .remove(doc_fields_key(database_id, tid, c, Surrogate::new(*s))) + .map_err(|e| redb_err("remove doc_fields", e))?; } } - removed += drop_str_quad_range(&write_txn, SEGMENTS, database_id, tid)?; - write_txn .commit() - .map_err(|e| redb_err("commit purge_tenant", e))?; + .map_err(|e| redb_err("commit purge", e))?; Ok(removed) } -/// Delete every `(database_id, tid, *, *)` row from a -/// `TableDefinition<(u64, u64, &str, &str), &[u8]>`. -fn drop_str_quad_range( - txn: &redb::WriteTransaction, - def: redb::TableDefinition<(u64, u64, &str, &str), &[u8]>, +/// Delete every row `owner` covers from a POSTINGS-shaped table. +fn drop_str_rows( + txn: &WriteTransaction, + def: TableDefinition<'static, StrKeyDef, &'static [u8]>, database_id: u64, tid: u64, + owner: KeyOwner<'_>, + name: &str, ) -> crate::Result { let mut table = txn .open_table(def) - .map_err(|e| redb_err("open quad table", e))?; - let keys: Vec<(String, String)> = table - .range((database_id, tid, "", "")..=(database_id, tid, MAX_COLLECTION, MAX_SUBKEY)) - .map_err(|e| redb_err("quad range", e))? - .filter_map(|r| { - r.ok().map(|(k, _)| { - let (_, _, c, s) = k.value(); - (c.to_string(), s.to_string()) - }) - }) - .collect(); - let n = keys.len(); - for (c, s) in &keys { - let _ = table.remove((database_id, tid, c.as_str(), s.as_str())); + .map_err(|e| redb_err(&format!("open {name}"), e))?; + let keys = scan::str_keys(&table, owner).map_err(|e| redb_err(&format!("{name} range"), e))?; + for (c, f, s) in &keys { + table + .remove(posting_key(database_id, tid, IndexScope::from_key(c, f), s)) + .map_err(|e| redb_err(&format!("remove {name}"), e))?; } - Ok(n) + Ok(keys.len()) } -/// Delete every `(database_id, tid, coll, *)` row from a per-document table -/// keyed by `(u64, u64, &str, u32)` (DOC_LENGTHS, DOC_TERMS). -fn drop_surrogate_quad_collection( - txn: &redb::WriteTransaction, - def: redb::TableDefinition<(u64, u64, &str, u32), &[u8]>, +/// Delete the INDEX_META rows `owner` covers that `meta` drops. +fn drop_meta_rows( + txn: &WriteTransaction, database_id: u64, tid: u64, - coll: &str, -) -> crate::Result { + owner: KeyOwner<'_>, + meta: MetaRows, +) -> crate::Result<()> { let mut table = txn - .open_table(def) - .map_err(|e| redb_err("open surrogate-keyed table", e))?; - let surrogates: Vec = table - .range((database_id, tid, coll, 0u32)..=(database_id, tid, coll, u32::MAX)) - .map_err(|e| redb_err("surrogate-keyed collection range", e))? - .filter_map(|r| r.ok().map(|(k, _)| k.value().3)) - .collect(); - let n = surrogates.len(); - for s in surrogates { - let _ = table.remove((database_id, tid, coll, s)); + .open_table(INDEX_META) + .map_err(|e| redb_err("open index_meta", e))?; + let keys = scan::str_keys(&table, owner).map_err(|e| redb_err("index_meta range", e))?; + for (c, f, s) in &keys { + if meta == MetaRows::KeepConfig && CONFIG_META_KEYS.contains(&s.as_str()) { + continue; + } + table + .remove(posting_key(database_id, tid, IndexScope::from_key(c, f), s)) + .map_err(|e| redb_err("remove index_meta", e))?; } - Ok(n) + Ok(()) } -/// Delete every `(database_id, tid, *, *)` row from a per-document table -/// keyed by `(u64, u64, &str, u32)` (DOC_LENGTHS, DOC_TERMS). -fn drop_surrogate_quad_tenant( - txn: &redb::WriteTransaction, - def: redb::TableDefinition<(u64, u64, &str, u32), &[u8]>, +/// Delete every row `owner` covers from a DOC_LENGTHS-shaped table. +fn drop_doc_rows( + txn: &WriteTransaction, + def: TableDefinition<'static, DocKeyDef, &'static [u8]>, database_id: u64, tid: u64, + owner: KeyOwner<'_>, + name: &str, ) -> crate::Result { let mut table = txn .open_table(def) - .map_err(|e| redb_err("open surrogate-keyed table", e))?; - let keys: Vec<(String, u32)> = table - .range((database_id, tid, "", 0u32)..=(database_id, tid, MAX_COLLECTION, u32::MAX)) - .map_err(|e| redb_err("surrogate-keyed tenant range", e))? - .filter_map(|r| { - r.ok().map(|(k, _)| { - let (_, _, c, s) = k.value(); - (c.to_string(), s) - }) - }) - .collect(); - let n = keys.len(); - for (c, s) in &keys { - let _ = table.remove((database_id, tid, c.as_str(), *s)); + .map_err(|e| redb_err(&format!("open {name}"), e))?; + let keys = scan::doc_keys(&table, owner).map_err(|e| redb_err(&format!("{name} range"), e))?; + for (c, f, s) in &keys { + table + .remove(doc_key( + database_id, + tid, + IndexScope::from_key(c, f), + Surrogate::new(*s), + )) + .map_err(|e| redb_err(&format!("remove {name}"), e))?; } - Ok(n) + Ok(keys.len()) } diff --git a/nodedb/src/engine/sparse/fts_redb/backend/segments.rs b/nodedb/src/engine/sparse/fts_redb/backend/segments.rs index e92187a83..91f857167 100644 --- a/nodedb/src/engine/sparse/fts_redb/backend/segments.rs +++ b/nodedb/src/engine/sparse/fts_redb/backend/segments.rs @@ -1,13 +1,16 @@ // SPDX-License-Identifier: BUSL-1.1 //! LSM segment blobs against `SEGMENTS` -//! keyed by `(database_id, tenant_id, collection, segment_id)`. +//! keyed by `(database_id, tenant_id, collection, field, segment_id)`. +use nodedb_fts::IndexScope; use redb::ReadableDatabase; use redb::ReadableTable as _; use super::core::RedbFtsBackend; -use super::shared::{MAX_SUBKEY, redb_err}; +use super::shared::redb_err; +use crate::engine::sparse::fts_redb::keys::{KeyOwner, segment_key}; +use crate::engine::sparse::fts_redb::scan; use crate::engine::sparse::fts_redb::tables::SEGMENTS; use crate::storage::quarantine::engines::validate_fts_segment_bytes; @@ -15,7 +18,7 @@ pub(super) fn write( backend: &RedbFtsBackend, database_id: u64, tid: u64, - collection: &str, + index: IndexScope<'_>, segment_id: &str, data: &[u8], ) -> crate::Result<()> { @@ -28,7 +31,7 @@ pub(super) fn write( .open_table(SEGMENTS) .map_err(|e| redb_err("open segments", e))?; table - .insert((database_id, tid, collection, segment_id), data) + .insert(segment_key(database_id, tid, index, segment_id), data) .map_err(|e| redb_err("insert segment", e))?; } write_txn.commit().map_err(|e| redb_err("commit", e))?; @@ -39,7 +42,7 @@ pub(super) fn read( backend: &RedbFtsBackend, database_id: u64, tid: u64, - collection: &str, + index: IndexScope<'_>, segment_id: &str, ) -> crate::Result>> { let read_txn = backend @@ -49,18 +52,16 @@ pub(super) fn read( let table = read_txn .open_table(SEGMENTS) .map_err(|e| redb_err("open segments", e))?; - let bytes = match table.get((database_id, tid, collection, segment_id)) { + let bytes = match table.get(segment_key(database_id, tid, index, segment_id)) { Ok(Some(val)) => val.value().to_vec(), Ok(None) => return Ok(None), Err(e) => return Err(redb_err("get segment", e)), }; if let Some(reg) = &backend.quarantine_registry { - let validated = - validate_fts_segment_bytes(reg, bytes, collection, segment_id).map_err(|e| { - crate::Error::SegmentCorrupted { - detail: e.to_string(), - } + let validated = validate_fts_segment_bytes(reg, bytes, index.collection(), segment_id) + .map_err(|e| crate::Error::SegmentCorrupted { + detail: e.to_string(), })?; Ok(Some(validated)) } else { @@ -72,7 +73,7 @@ pub(super) fn list( backend: &RedbFtsBackend, database_id: u64, tid: u64, - collection: &str, + index: IndexScope<'_>, ) -> crate::Result> { let read_txn = backend .db @@ -81,19 +82,16 @@ pub(super) fn list( let table = read_txn .open_table(SEGMENTS) .map_err(|e| redb_err("open segments", e))?; - let ids: Vec = table - .range((database_id, tid, collection, "")..=(database_id, tid, collection, MAX_SUBKEY)) - .map_err(|e| redb_err("range", e))? - .filter_map(|r| r.ok().map(|(k, _)| k.value().3.to_string())) - .collect(); - Ok(ids) + let keys = scan::str_keys(&table, KeyOwner::index(database_id, tid, index)) + .map_err(|e| redb_err("segments range", e))?; + Ok(keys.into_iter().map(|(_, _, id)| id).collect()) } pub(super) fn remove( backend: &RedbFtsBackend, database_id: u64, tid: u64, - collection: &str, + index: IndexScope<'_>, segment_id: &str, ) -> crate::Result<()> { let write_txn = backend @@ -108,7 +106,7 @@ pub(super) fn remove( // the error so a corrupt index or aborted txn surfaces; the absent-key // case (`Ok(None)`) is intentionally a no-op, not an error. table - .remove((database_id, tid, collection, segment_id)) + .remove(segment_key(database_id, tid, index, segment_id)) .map_err(|e| redb_err("remove segment", e))?; } write_txn.commit().map_err(|e| redb_err("commit", e))?; @@ -116,14 +114,14 @@ pub(super) fn remove( } /// Inputs to [`compact_commit`]: the merged-segment write plus the source -/// segment ids to remove, scoped to one `(database, tenant, collection)`. +/// segment ids to remove, scoped to one `(database, tenant, index)`. pub struct CompactCommit<'a> { /// Owning database id. pub database_id: u64, /// Owning tenant id (raw storage form). pub tid: u64, - /// Collection whose segments are being compacted. - pub collection: &'a str, + /// Index whose segments are being compacted. + pub index: IndexScope<'a>, /// Id of the new merged segment to write. pub new_segment_id: &'a str, /// Bytes of the new merged segment. @@ -147,7 +145,7 @@ pub(super) fn compact_commit( let CompactCommit { database_id, tid, - collection, + index, new_segment_id, new_segment_data, merged_ids, @@ -162,7 +160,7 @@ pub(super) fn compact_commit( .map_err(|e| redb_err("open segments for compact", e))?; table .insert( - (database_id, tid, collection, new_segment_id), + segment_key(database_id, tid, index, new_segment_id), new_segment_data, ) .map_err(|e| redb_err("insert merged segment", e))?; @@ -172,7 +170,7 @@ pub(super) fn compact_commit( // which both old and new segments are visible to readers, causing // double-counted postings during BM25 scoring. table - .remove((database_id, tid, collection, id.as_str())) + .remove(segment_key(database_id, tid, index, id.as_str())) .map_err(|e| redb_err("remove merged source segment", e))?; } } @@ -182,15 +180,15 @@ pub(super) fn compact_commit( Ok(()) } -/// Enumerate all `(database_id, tid, collection)` triples that have at least one -/// segment stored in the SEGMENTS table. +/// Enumerate every `(database_id, tid, collection, field)` index that has at +/// least one segment stored in the SEGMENTS table. /// -/// Used by maintenance to discover which collections need FTS LSM compaction +/// Used by maintenance to discover which indexes need FTS LSM compaction /// without requiring the executor to maintain a separate registry of /// FTS-indexed collections. -pub(super) fn list_all_collections( +pub(super) fn list_all_indexes( backend: &RedbFtsBackend, -) -> crate::Result> { +) -> crate::Result> { let read_txn = backend .db .begin_read() @@ -199,20 +197,16 @@ pub(super) fn list_all_collections( .open_table(SEGMENTS) .map_err(|e| redb_err("open segments", e))?; - let mut collections: Vec<(u64, u64, String)> = Vec::new(); - let mut last: Option<(u64, u64, String)> = None; - + let mut indexes: Vec<(u64, u64, String, String)> = Vec::new(); for entry in table.iter().map_err(|e| redb_err("iter segments", e))? { let (key, _) = entry.map_err(|e| redb_err("next segment key", e))?; - let (database_id, tid, collection, _) = key.value(); - let coll = collection.to_string(); - match &last { - Some((d, t, c)) if *d == database_id && *t == tid && c == &coll => {} - _ => { - collections.push((database_id, tid, coll.clone())); - last = Some((database_id, tid, coll)); - } + let (database_id, tid, collection, field, _) = key.value(); + let same_as_last = indexes.last().is_some_and(|(d, t, c, f)| { + *d == database_id && *t == tid && c == collection && f == field + }); + if !same_as_last { + indexes.push((database_id, tid, collection.to_string(), field.to_string())); } } - Ok(collections) + Ok(indexes) } diff --git a/nodedb/src/engine/sparse/fts_redb/backend/shared.rs b/nodedb/src/engine/sparse/fts_redb/backend/shared.rs index 850ce39b9..1ff04d949 100644 --- a/nodedb/src/engine/sparse/fts_redb/backend/shared.rs +++ b/nodedb/src/engine/sparse/fts_redb/backend/shared.rs @@ -1,18 +1,6 @@ // SPDX-License-Identifier: BUSL-1.1 -//! Shared helpers for the FTS backend: error construction and the -//! inclusive upper bound for `(database_id, tid, collection, subkey)` range scans. - -/// Upper-bound sentinel for the `subkey` tuple component of a -/// `(database_id, tid, collection, subkey)` range scan. `"\u{10ffff}"` is the -/// highest scalar value representable in UTF-8, so every valid subkey sorts -/// below it. -pub(super) const MAX_SUBKEY: &str = "\u{10ffff}"; - -/// Upper-bound sentinel for the `collection` tuple component of a -/// `(database_id, tid, collection, …)` range scan. Same reasoning as -/// [`MAX_SUBKEY`]. -pub(super) const MAX_COLLECTION: &str = "\u{10ffff}"; +//! Shared helpers for the FTS backend: error construction. pub(super) fn redb_err(ctx: &str, e: impl std::fmt::Display) -> crate::Error { crate::Error::Storage { diff --git a/nodedb/src/engine/sparse/fts_redb/backend/stats.rs b/nodedb/src/engine/sparse/fts_redb/backend/stats.rs index 8186a998e..d98ecb232 100644 --- a/nodedb/src/engine/sparse/fts_redb/backend/stats.rs +++ b/nodedb/src/engine/sparse/fts_redb/backend/stats.rs @@ -1,19 +1,21 @@ // SPDX-License-Identifier: BUSL-1.1 //! Corpus statistics against the dedicated `STATS` table -//! keyed by `(database_id, tenant_id, collection)`. +//! keyed by `(database_id, tenant_id, collection, field)`. +use nodedb_fts::IndexScope; use redb::{ReadableDatabase, ReadableTable}; use super::core::RedbFtsBackend; use super::shared::redb_err; +use crate::engine::sparse::fts_redb::keys::stats_key; use crate::engine::sparse::fts_redb::tables::STATS; pub(super) fn read( backend: &RedbFtsBackend, database_id: u64, tid: u64, - collection: &str, + index: IndexScope<'_>, ) -> crate::Result<(u32, u64)> { let read_txn = backend .db @@ -22,7 +24,7 @@ pub(super) fn read( let table = read_txn .open_table(STATS) .map_err(|e| redb_err("open stats", e))?; - match table.get((database_id, tid, collection)) { + match table.get(stats_key(database_id, tid, index)) { Ok(Some(val)) => { let stats: (u32, u64) = zerompk::from_msgpack(val.value()).map_err(|e| redb_err("deserialize stats", e))?; @@ -37,28 +39,31 @@ pub(super) fn increment( backend: &RedbFtsBackend, database_id: u64, tid: u64, - collection: &str, + index: IndexScope<'_>, doc_len: u32, ) -> crate::Result<()> { - bump(backend, database_id, tid, collection, doc_len as i64) + bump(backend, database_id, tid, index, 1, i64::from(doc_len)) } pub(super) fn decrement( backend: &RedbFtsBackend, database_id: u64, tid: u64, - collection: &str, + index: IndexScope<'_>, doc_len: u32, ) -> crate::Result<()> { - bump(backend, database_id, tid, collection, -(doc_len as i64)) + bump(backend, database_id, tid, index, -1, -i64::from(doc_len)) } +/// Move one index's `(doc_count, total_token_sum)` by explicit deltas, +/// clamped at zero. fn bump( backend: &RedbFtsBackend, database_id: u64, tid: u64, - collection: &str, - delta: i64, + index: IndexScope<'_>, + count_delta: i64, + total_delta: i64, ) -> crate::Result<()> { let write_txn = backend .db @@ -68,23 +73,18 @@ fn bump( let mut table = write_txn .open_table(STATS) .map_err(|e| redb_err("open stats", e))?; - let (mut count, mut total) = table - .get((database_id, tid, collection)) - .ok() - .flatten() - .and_then(|v| zerompk::from_msgpack::<(u32, u64)>(v.value()).ok()) - .unwrap_or((0, 0)); - if delta > 0 { - count += 1; - total += delta as u64; - } else { - count = count.saturating_sub(1); - total = total.saturating_sub((-delta) as u64); - } + let key = stats_key(database_id, tid, index); + let (count, total) = match table.get(key).map_err(|e| redb_err("get stats", e))? { + Some(v) => zerompk::from_msgpack::<(u32, u64)>(v.value()) + .map_err(|e| redb_err("deserialize stats", e))?, + None => (0, 0), + }; + let count = (i64::from(count) + count_delta).clamp(0, i64::from(u32::MAX)) as u32; + let total = (total as i64).saturating_add(total_delta).max(0) as u64; let bytes = zerompk::to_msgpack_vec(&(count, total)).map_err(|e| redb_err("serialize stats", e))?; table - .insert((database_id, tid, collection), bytes.as_slice()) + .insert(key, bytes.as_slice()) .map_err(|e| redb_err("insert stats", e))?; } write_txn.commit().map_err(|e| redb_err("commit", e))?; diff --git a/nodedb/src/engine/sparse/fts_redb/keys.rs b/nodedb/src/engine/sparse/fts_redb/keys.rs new file mode 100644 index 000000000..c36f73930 --- /dev/null +++ b/nodedb/src/engine/sparse/fts_redb/keys.rs @@ -0,0 +1,239 @@ +// SPDX-License-Identifier: BUSL-1.1 + +//! The single place FTS table keys are formed. +//! +//! Per-index keys carry `(database_id, tenant_id, collection, field)`, where +//! `field` is the index's field key: empty for the whole-document index. +//! Scans over many indexes start at a [`KeyOwner`]'s lower bound and stop at +//! the first key it does not own, so no name needs an upper-bound sentinel. + +use nodedb_fts::IndexScope; +use nodedb_types::Surrogate; + +/// `(db, tid, collection, field, term | subkey | segment_id)`: POSTINGS, +/// INDEX_META, SEGMENTS. +pub(crate) type FieldStrKey<'a> = (u64, u64, &'a str, &'a str, &'a str); +/// `(db, tid, collection, field, surrogate)`: DOC_LENGTHS, DOC_TERMS. +pub(crate) type FieldDocKey<'a> = (u64, u64, &'a str, &'a str, u32); +/// `(db, tid, collection, field)`: STATS. +pub(crate) type StatsKey<'a> = (u64, u64, &'a str, &'a str); +/// `(db, tid, collection, surrogate)`: DOC_FIELDS. +pub(crate) type DocFieldsKey<'a> = (u64, u64, &'a str, u32); + +/// POSTINGS key of `term` in one index. +pub(crate) fn posting_key<'a>( + database_id: u64, + tid: u64, + index: IndexScope<'a>, + term: &'a str, +) -> FieldStrKey<'a> { + ( + database_id, + tid, + index.collection(), + index.field_key(), + term, + ) +} + +/// INDEX_META key of `subkey` in one index. +pub(crate) fn meta_key<'a>( + database_id: u64, + tid: u64, + index: IndexScope<'a>, + subkey: &'a str, +) -> FieldStrKey<'a> { + posting_key(database_id, tid, index, subkey) +} + +/// SEGMENTS key of `segment_id` in one index. +pub(crate) fn segment_key<'a>( + database_id: u64, + tid: u64, + index: IndexScope<'a>, + segment_id: &'a str, +) -> FieldStrKey<'a> { + posting_key(database_id, tid, index, segment_id) +} + +/// DOC_LENGTHS / DOC_TERMS key of one document in one index. +pub(crate) fn doc_key<'a>( + database_id: u64, + tid: u64, + index: IndexScope<'a>, + surrogate: Surrogate, +) -> FieldDocKey<'a> { + ( + database_id, + tid, + index.collection(), + index.field_key(), + surrogate.as_u32(), + ) +} + +/// STATS key of one index. +pub(crate) fn stats_key<'a>(database_id: u64, tid: u64, index: IndexScope<'a>) -> StatsKey<'a> { + (database_id, tid, index.collection(), index.field_key()) +} + +/// DOC_FIELDS key of one document. +pub(crate) fn doc_fields_key( + database_id: u64, + tid: u64, + collection: &str, + surrogate: Surrogate, +) -> DocFieldsKey<'_> { + (database_id, tid, collection, surrogate.as_u32()) +} + +/// The set of keys a scan covers: one index, one collection, or one tenant. +#[derive(Debug, Clone, Copy)] +pub(crate) struct KeyOwner<'a> { + database_id: u64, + tid: u64, + collection: Option<&'a str>, + field: Option<&'a str>, +} + +impl<'a> KeyOwner<'a> { + /// Every key of one index. + pub(crate) fn index(database_id: u64, tid: u64, index: IndexScope<'a>) -> Self { + Self { + database_id, + tid, + collection: Some(index.collection()), + field: Some(index.field_key()), + } + } + + /// Every key of every index of one collection. + pub(crate) fn collection(database_id: u64, tid: u64, collection: &'a str) -> Self { + Self { + database_id, + tid, + collection: Some(collection), + field: None, + } + } + + /// Every key of one `(database, tenant)`. + pub(crate) fn tenant(database_id: u64, tid: u64) -> Self { + Self { + database_id, + tid, + collection: None, + field: None, + } + } + + /// Whether a key with these leading components belongs to this owner. + /// `field` is `None` for a table without a field component. + pub(crate) fn owns( + &self, + database_id: u64, + tid: u64, + collection: &str, + field: Option<&str>, + ) -> bool { + database_id == self.database_id + && tid == self.tid + && self.collection.is_none_or(|c| c == collection) + && match (self.field, field) { + (Some(own), Some(key)) => own == key, + _ => true, + } + } + + fn collection_bound(&self) -> &'a str { + self.collection.unwrap_or("") + } + + fn field_bound(&self) -> &'a str { + self.field.unwrap_or("") + } + + /// Lowest POSTINGS / INDEX_META / SEGMENTS key this owner covers. + pub(crate) fn str_start(&self) -> FieldStrKey<'a> { + ( + self.database_id, + self.tid, + self.collection_bound(), + self.field_bound(), + "", + ) + } + + /// Lowest DOC_LENGTHS / DOC_TERMS key this owner covers. + pub(crate) fn doc_start(&self) -> FieldDocKey<'a> { + ( + self.database_id, + self.tid, + self.collection_bound(), + self.field_bound(), + 0, + ) + } + + /// Lowest STATS key this owner covers. + pub(crate) fn stats_start(&self) -> StatsKey<'a> { + ( + self.database_id, + self.tid, + self.collection_bound(), + self.field_bound(), + ) + } + + /// Lowest DOC_FIELDS key this owner covers. + pub(crate) fn doc_fields_start(&self) -> DocFieldsKey<'a> { + (self.database_id, self.tid, self.collection_bound(), 0) + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn document_and_field_keys_differ_only_in_the_field_component() { + let doc = IndexScope::document("c"); + let title = IndexScope::field("c", "title").unwrap(); + assert_eq!(posting_key(1, 2, doc, "t"), (1, 2, "c", "", "t")); + assert_eq!(posting_key(1, 2, title, "t"), (1, 2, "c", "title", "t")); + assert_eq!(stats_key(1, 2, title), (1, 2, "c", "title")); + assert_eq!(doc_key(1, 2, doc, Surrogate::new(7)), (1, 2, "c", "", 7)); + } + + #[test] + fn owners_cover_exactly_their_keys() { + let title = IndexScope::field("c", "title").unwrap(); + let index = KeyOwner::index(1, 2, title); + assert!(index.owns(1, 2, "c", Some("title"))); + assert!(!index.owns(1, 2, "c", Some(""))); + assert!(!index.owns(1, 2, "c", Some("title2"))); + assert!(!index.owns(1, 2, "cc", Some("title"))); + + let coll = KeyOwner::collection(1, 2, "c"); + assert!(coll.owns(1, 2, "c", Some(""))); + assert!(coll.owns(1, 2, "c", Some("\u{10ffff}field"))); + assert!(coll.owns(1, 2, "c", None)); + assert!(!coll.owns(1, 2, "c\0", Some(""))); + assert!(!coll.owns(1, 3, "c", Some(""))); + + let tenant = KeyOwner::tenant(1, 2); + assert!(tenant.owns(1, 2, "anything", Some("x"))); + assert!(!tenant.owns(2, 2, "anything", Some("x"))); + } + + #[test] + fn owner_lower_bounds_sort_before_every_owned_key() { + let coll = KeyOwner::collection(1, 2, "c"); + assert!(coll.str_start() <= (1, 2, "c", "", "")); + assert!(coll.doc_start() <= (1, 2, "c", "", 0)); + assert!(coll.stats_start() <= (1, 2, "c", "")); + assert!(coll.doc_fields_start() <= (1, 2, "c", 0)); + let tenant = KeyOwner::tenant(1, 2); + assert!(tenant.str_start() <= (1, 2, "", "", "")); + } +} diff --git a/nodedb/src/engine/sparse/fts_redb/mod.rs b/nodedb/src/engine/sparse/fts_redb/mod.rs index 01c0129f0..978232bf7 100644 --- a/nodedb/src/engine/sparse/fts_redb/mod.rs +++ b/nodedb/src/engine/sparse/fts_redb/mod.rs @@ -1,6 +1,8 @@ // SPDX-License-Identifier: BUSL-1.1 pub mod backend; +pub mod keys; +pub mod scan; pub mod tables; pub use backend::{CompactCommit, RedbFtsBackend}; diff --git a/nodedb/src/engine/sparse/fts_redb/scan.rs b/nodedb/src/engine/sparse/fts_redb/scan.rs new file mode 100644 index 000000000..ec228a480 --- /dev/null +++ b/nodedb/src/engine/sparse/fts_redb/scan.rs @@ -0,0 +1,205 @@ +// SPDX-License-Identifier: BUSL-1.1 + +//! Owned-key scans over the FTS tables. +//! +//! Each scan starts at a [`KeyOwner`]'s lower bound and stops at the first +//! key the owner does not own. A failed read surfaces as an error: no entry +//! is skipped. + +use redb::{ReadableTable, StorageError}; + +use super::keys::KeyOwner; + +/// Key definition of POSTINGS / INDEX_META / SEGMENTS. +pub(crate) type StrKeyDef = (u64, u64, &'static str, &'static str, &'static str); +/// Key definition of DOC_LENGTHS / DOC_TERMS. +pub(crate) type DocKeyDef = (u64, u64, &'static str, &'static str, u32); +/// Key definition of STATS. +pub(crate) type StatsKeyDef = (u64, u64, &'static str, &'static str); +/// Key definition of DOC_FIELDS. +pub(crate) type DocFieldsKeyDef = (u64, u64, &'static str, u32); + +/// One row of a string-suffixed table within one collection: +/// `(field, suffix, value)`. +pub(crate) struct StrRow { + pub(crate) field: String, + pub(crate) suffix: String, + pub(crate) value: Vec, +} + +/// One row of a document table within one collection: +/// `(field, surrogate, value)`. +pub(crate) struct DocRow { + pub(crate) field: String, + pub(crate) surrogate: u32, + pub(crate) value: Vec, +} + +/// Every POSTINGS / INDEX_META / SEGMENTS row of one collection, with values. +pub(crate) fn str_rows( + table: &T, + database_id: u64, + tid: u64, + collection: &str, +) -> Result, StorageError> +where + T: ReadableTable, +{ + let owner = KeyOwner::collection(database_id, tid, collection); + let mut out = Vec::new(); + for entry in table.range(owner.str_start()..)? { + let (key, value) = entry?; + let (db, tid, collection, field, suffix) = key.value(); + if !owner.owns(db, tid, collection, Some(field)) { + break; + } + out.push(StrRow { + field: field.to_string(), + suffix: suffix.to_string(), + value: value.value().to_vec(), + }); + } + Ok(out) +} + +/// Every `(collection, field, suffix)` key `owner` owns in a string-suffixed table. +pub(crate) fn str_keys( + table: &T, + owner: KeyOwner<'_>, +) -> Result, StorageError> +where + T: ReadableTable, +{ + let mut out = Vec::new(); + for entry in table.range(owner.str_start()..)? { + let (key, _) = entry?; + let (db, tid, collection, field, suffix) = key.value(); + if !owner.owns(db, tid, collection, Some(field)) { + break; + } + out.push(( + collection.to_string(), + field.to_string(), + suffix.to_string(), + )); + } + Ok(out) +} + +/// Whether `owner` owns any row of a string-suffixed table. +pub(crate) fn has_str_rows(table: &T, owner: KeyOwner<'_>) -> Result +where + T: ReadableTable, +{ + match table.range(owner.str_start()..)?.next() { + Some(entry) => { + let (key, _) = entry?; + let (db, tid, collection, field, _) = key.value(); + Ok(owner.owns(db, tid, collection, Some(field))) + } + None => Ok(false), + } +} + +/// Whether `owner` owns any row of a document table. +pub(crate) fn has_doc_rows(table: &T, owner: KeyOwner<'_>) -> Result +where + T: ReadableTable, +{ + match table.range(owner.doc_start()..)?.next() { + Some(entry) => { + let (key, _) = entry?; + let (db, tid, collection, field, _) = key.value(); + Ok(owner.owns(db, tid, collection, Some(field))) + } + None => Ok(false), + } +} + +/// Every DOC_LENGTHS / DOC_TERMS row of one collection, with values. +pub(crate) fn doc_rows( + table: &T, + database_id: u64, + tid: u64, + collection: &str, +) -> Result, StorageError> +where + T: ReadableTable, +{ + let owner = KeyOwner::collection(database_id, tid, collection); + let mut out = Vec::new(); + for entry in table.range(owner.doc_start()..)? { + let (key, value) = entry?; + let (db, tid, collection, field, surrogate) = key.value(); + if !owner.owns(db, tid, collection, Some(field)) { + break; + } + out.push(DocRow { + field: field.to_string(), + surrogate, + value: value.value().to_vec(), + }); + } + Ok(out) +} + +/// Every `(collection, field, surrogate)` key `owner` owns in a document table. +pub(crate) fn doc_keys( + table: &T, + owner: KeyOwner<'_>, +) -> Result, StorageError> +where + T: ReadableTable, +{ + let mut out = Vec::new(); + for entry in table.range(owner.doc_start()..)? { + let (key, _) = entry?; + let (db, tid, collection, field, surrogate) = key.value(); + if !owner.owns(db, tid, collection, Some(field)) { + break; + } + out.push((collection.to_string(), field.to_string(), surrogate)); + } + Ok(out) +} + +/// Every `(collection, field)` STATS key `owner` owns. +pub(crate) fn stats_keys( + table: &T, + owner: KeyOwner<'_>, +) -> Result, StorageError> +where + T: ReadableTable, +{ + let mut out = Vec::new(); + for entry in table.range(owner.stats_start()..)? { + let (key, _) = entry?; + let (db, tid, collection, field) = key.value(); + if !owner.owns(db, tid, collection, Some(field)) { + break; + } + out.push((collection.to_string(), field.to_string())); + } + Ok(out) +} + +/// Every `(collection, surrogate)` DOC_FIELDS key `owner` owns. A field +/// owner covers its whole collection here: DOC_FIELDS has no field component. +pub(crate) fn doc_fields_keys( + table: &T, + owner: KeyOwner<'_>, +) -> Result, StorageError> +where + T: ReadableTable, +{ + let mut out = Vec::new(); + for entry in table.range(owner.doc_fields_start()..)? { + let (key, _) = entry?; + let (db, tid, collection, surrogate) = key.value(); + if !owner.owns(db, tid, collection, None) { + break; + } + out.push((collection.to_string(), surrogate)); + } + Ok(out) +} diff --git a/nodedb/src/engine/sparse/fts_redb/tables.rs b/nodedb/src/engine/sparse/fts_redb/tables.rs index 284ed2b03..4beefebc8 100644 --- a/nodedb/src/engine/sparse/fts_redb/tables.rs +++ b/nodedb/src/engine/sparse/fts_redb/tables.rs @@ -2,48 +2,59 @@ //! Redb table definitions for the full-text search backend. //! -//! Every table is keyed by a structural `(database_id, tenant_id, collection, …)` -//! tuple — matching the EdgeStore pattern. Per-`(database, tenant)` drops become -//! range scans `(db, tid, ..)..(db, tid+1, ..)` instead of fragile lexical-prefix -//! scans. +//! Every table is keyed by a structural +//! `(database_id, tenant_id, collection, field, …)` tuple — matching the +//! EdgeStore pattern. `field` is the index's field key: empty for the +//! collection's whole-document index, the field name for a field index. +//! Per-`(database, tenant)` and per-collection drops become range scans over +//! the tuple instead of fragile lexical-prefix scans. Keys are formed only in +//! the sibling `keys` module. use redb::TableDefinition; -/// Inverted index: key = `(database_id, tenant_id, collection, term)`, +/// Inverted index: key = `(database_id, tenant_id, collection, field, term)`, /// value = MessagePack-encoded `Vec`. -pub const POSTINGS: TableDefinition<(u64, u64, &str, &str), &[u8]> = +pub const POSTINGS: TableDefinition<(u64, u64, &str, &str, &str), &[u8]> = TableDefinition::new("text.postings"); -/// Document lengths: key = `(database_id, tenant_id, collection, surrogate)`, +/// Document lengths: key = `(database_id, tenant_id, collection, field, surrogate)`, /// value = MessagePack-encoded `u32` token count. /// The surrogate is stored as its raw `u32` value (redb native key type). -pub const DOC_LENGTHS: TableDefinition<(u64, u64, &str, u32), &[u8]> = +pub const DOC_LENGTHS: TableDefinition<(u64, u64, &str, &str, u32), &[u8]> = TableDefinition::new("text.doc_lengths"); /// Per-document indexed term set: key = -/// `(database_id, tenant_id, collection, surrogate)`, value = +/// `(database_id, tenant_id, collection, field, surrogate)`, value = /// MessagePack-encoded `Vec` holding the distinct analyzed terms the -/// document currently contributes a posting to. +/// document currently contributes a posting to in that index. /// /// This is what makes a re-index able to SUBTRACT. A re-index only sees the /// new text's terms; without a record of the previous version's terms it /// cannot tell which posting lists the document must be dropped from, so /// words removed by an update keep matching the document forever. Keyed /// exactly like `DOC_LENGTHS` so both are read in the same lookup pattern. -pub const DOC_TERMS: TableDefinition<(u64, u64, &str, u32), &[u8]> = +pub const DOC_TERMS: TableDefinition<(u64, u64, &str, &str, u32), &[u8]> = TableDefinition::new("text.doc_terms"); -/// Index metadata blobs: key = `(database_id, tenant_id, collection, sub_key)`, -/// value = opaque blob. Sub-keys: `"docmap"`, `"fieldnorms"`, `"analyzer"`, -/// `"language"`. -pub const INDEX_META: TableDefinition<(u64, u64, &str, &str), &[u8]> = +/// Field indexes a document occupies: key = +/// `(database_id, tenant_id, collection, surrogate)`, value = MessagePack +/// `Vec` of non-empty field names. A re-index retracts the document +/// from every field index listed here that its new text no longer fills. +pub const DOC_FIELDS: TableDefinition<(u64, u64, &str, u32), &[u8]> = + TableDefinition::new("text.doc_fields"); + +/// Index metadata blobs: key = `(database_id, tenant_id, collection, field, sub_key)`, +/// value = opaque blob. Sub-keys: `"fieldnorms"` per index; `"analyzer"`, +/// `"language"`, `"fuzzy"` under the whole-document index. +pub const INDEX_META: TableDefinition<(u64, u64, &str, &str, &str), &[u8]> = TableDefinition::new("text.meta"); -/// Corpus stats: key = `(database_id, tenant_id, collection)`, +/// Corpus stats: key = `(database_id, tenant_id, collection, field)`, /// value = MessagePack-encoded `(doc_count, total_token_sum)`. -pub const STATS: TableDefinition<(u64, u64, &str), &[u8]> = TableDefinition::new("text.stats"); +pub const STATS: TableDefinition<(u64, u64, &str, &str), &[u8]> = + TableDefinition::new("text.stats"); -/// Segment blobs: key = `(database_id, tenant_id, collection, segment_id)`, +/// Segment blobs: key = `(database_id, tenant_id, collection, field, segment_id)`, /// value = compressed segment bytes. `segment_id` format `"L{level}:{id:016x}"`. -pub const SEGMENTS: TableDefinition<(u64, u64, &str, &str), &[u8]> = +pub const SEGMENTS: TableDefinition<(u64, u64, &str, &str, &str), &[u8]> = TableDefinition::new("text.segments"); diff --git a/nodedb/src/engine/sparse/inverted/batch_removal.rs b/nodedb/src/engine/sparse/inverted/batch_removal.rs new file mode 100644 index 000000000..1ee4af432 --- /dev/null +++ b/nodedb/src/engine/sparse/inverted/batch_removal.rs @@ -0,0 +1,206 @@ +// SPDX-License-Identifier: BUSL-1.1 + +//! Removal of many documents from the inverted index in one write +//! transaction. +//! +//! A bulk delete removes its rows' text in batches. Each posting list a batch +//! occupies is rewritten once for the whole batch, and each index's corpus +//! counters are updated once, so a batch costs the lists it touches rather +//! than one rewrite per document per term. + +use std::collections::{BTreeMap, HashSet}; + +use nodedb_fts::IndexScope; +use nodedb_types::{Surrogate, TenantId}; + +use super::core::InvertedIndex; +use super::doc_fields; +use super::doc_terms; +use super::errors::inverted_err; +use super::indexing::{IndexDocScope, prior_doc_length}; +use crate::engine::sparse::fts_redb::keys::doc_key; +use crate::engine::sparse::fts_redb::tables::DOC_LENGTHS; + +/// The edits one index takes from a batch. +#[derive(Default)] +struct IndexEdits { + /// Term → documents leaving its posting list. + postings: BTreeMap>, + /// Documents leaving the index. + count: i64, + /// Token length leaving the index. + total: i64, +} + +impl InvertedIndex { + /// Remove `surrogates` from every index of `collection` in one write + /// transaction. A document not in an index is left alone there. Either + /// every removal commits or none does. + pub fn remove_documents( + &self, + database_id: u64, + tid: TenantId, + collection: &str, + surrogates: &[Surrogate], + ) -> crate::Result<()> { + if surrogates.is_empty() { + return Ok(()); + } + let db = self.inner.backend().db(); + let txn = db.begin_write().map_err(|e| inverted_err("write txn", e))?; + // Keyed by field key: empty for the whole-document index. + let mut edits: BTreeMap = BTreeMap::new(); + for surrogate in surrogates { + let doc = IndexDocScope { + database_id, + tid, + collection, + surrogate: *surrogate, + }; + self.note_doc_write(doc); + let mut field_keys = doc_fields::read(&txn, doc)?; + field_keys.push(String::new()); + for field_key in field_keys { + let index = scope_of(collection, &field_key)?; + let Some(old_len) = prior_doc_length(&txn, doc, index)? else { + continue; + }; + let terms = doc_terms::occupied_terms(&txn, doc, index, true)?; + doc_terms::clear(&txn, doc, index)?; + { + let mut lengths = txn + .open_table(DOC_LENGTHS) + .map_err(|e| inverted_err("open doc_lengths", e))?; + lengths + .remove(doc_key(database_id, tid.as_u64(), index, *surrogate)) + .map_err(|e| inverted_err("remove doc length", e))?; + } + let entry = edits.entry(field_key).or_default(); + for term in terms { + entry.postings.entry(term).or_default().insert(*surrogate); + } + entry.count -= 1; + entry.total -= i64::from(old_len); + } + doc_fields::clear(&txn, doc)?; + } + for (field_key, index_edits) in &edits { + let index = scope_of(collection, field_key)?; + doc_terms::strip_postings_many( + &txn, + database_id, + tid.as_u64(), + index, + &index_edits.postings, + )?; + Self::update_stats_in_txn( + &txn, + database_id, + tid, + index, + index_edits.count, + index_edits.total, + )?; + } + txn.commit() + .map_err(|e| inverted_err("commit batch remove", e))?; + Ok(()) + } +} + +/// The index a stored field key names: the whole-document index for the +/// empty key. +fn scope_of<'a>(collection: &'a str, field_key: &'a str) -> crate::Result> { + if field_key.is_empty() { + return Ok(IndexScope::document(collection)); + } + IndexScope::field(collection, field_key) + .ok_or_else(|| inverted_err("doc_fields", "stored field name is empty")) +} + +#[cfg(test)] +mod tests { + use std::sync::Arc; + + use nodedb_fts::FtsSearchParams; + use nodedb_fts::posting::QueryMode; + + use super::*; + use crate::engine::durability_gate::GatedDatabase; + use crate::engine::sparse::inverted::test_support::{body, fields}; + + const DB: u64 = 0; + const T: TenantId = TenantId::new(1); + + fn open_temp() -> (InvertedIndex, tempfile::TempDir) { + let dir = tempfile::tempdir().unwrap(); + let path = dir.path().join("test-inverted.redb"); + let db = Arc::new(GatedDatabase::new(redb::Database::create(&path).unwrap())); + let idx = + InvertedIndex::open(db, crate::data::executor::core_loop::test_governor()).unwrap(); + (idx, dir) + } + + fn hits(idx: &InvertedIndex, index: IndexScope<'_>, query: &str) -> Vec { + idx.search( + DB, + T, + index, + FtsSearchParams { + query, + top_k: 100, + fuzzy_enabled: false, + mode: QueryMode::And, + prefilter: None, + }, + ) + .unwrap() + .into_iter() + .map(|h| h.doc_id) + .collect() + } + + /// A batch removal leaves the index exactly as removing each document + /// one by one does: postings, counters, and field indexes. + #[test] + fn batch_removal_matches_one_by_one_removal() { + let (idx, _dir) = open_temp(); + for i in 1..=5u32 { + idx.index_document( + DB, + T, + "docs", + Surrogate::new(i), + &fields(&[("title", "rust"), ("body", "shared words here")]), + ) + .unwrap(); + } + idx.index_document(DB, T, "docs", Surrogate::new(6), &body("rust only")) + .unwrap(); + + let removed: Vec = (1..=5u32).map(Surrogate::new).collect(); + idx.remove_documents(DB, T, "docs", &removed).unwrap(); + + assert_eq!( + hits(&idx, IndexScope::document("docs"), "rust"), + vec![Surrogate::new(6)] + ); + assert!(hits(&idx, IndexScope::document("docs"), "shared").is_empty()); + let title = IndexScope::field("docs", "title").unwrap(); + assert!(hits(&idx, title, "rust").is_empty()); + assert_eq!(idx.corpus_stats(DB, T, "docs").unwrap().0, 1); + assert_eq!(idx.corpus_stats(DB, T, title).unwrap().0, 0); + } + + /// A surrogate the index does not hold is a no-op, not a second + /// decrement. + #[test] + fn removing_an_unindexed_document_changes_nothing() { + let (idx, _dir) = open_temp(); + idx.index_document(DB, T, "docs", Surrogate::new(1), &body("alpha")) + .unwrap(); + idx.remove_documents(DB, T, "docs", &[Surrogate::new(9)]) + .unwrap(); + assert_eq!(idx.corpus_stats(DB, T, "docs").unwrap().0, 1); + } +} diff --git a/nodedb/src/engine/sparse/inverted/compaction.rs b/nodedb/src/engine/sparse/inverted/compaction.rs index 52c183751..9100d562d 100644 --- a/nodedb/src/engine/sparse/inverted/compaction.rs +++ b/nodedb/src/engine/sparse/inverted/compaction.rs @@ -23,21 +23,19 @@ impl InvertedIndex { self.inner.backend().compact_commit(params) } - /// Enumerate every `(database_id, TenantId, collection)` triple that has at - /// least one FTS segment in the backing store. + /// Enumerate every `(database_id, TenantId, collection, field)` index + /// that has at least one FTS segment in the backing store. `field` is + /// empty for a whole-document index. /// /// Used by the maintenance cycle to discover compaction candidates /// without requiring a separate in-memory registry of FTS-indexed /// collections. - pub fn list_all_fts_collections(&self) -> crate::Result> { - self.inner - .backend() - .list_all_fts_collections() - .map(|triples| { - triples - .into_iter() - .map(|(d, t, c)| (d, TenantId::new(t), c)) - .collect() - }) + pub fn list_all_fts_indexes(&self) -> crate::Result> { + self.inner.backend().list_all_fts_indexes().map(|indexes| { + indexes + .into_iter() + .map(|(d, t, c, f)| (d, TenantId::new(t), c, f)) + .collect() + }) } } diff --git a/nodedb/src/engine/sparse/inverted/core.rs b/nodedb/src/engine/sparse/inverted/core.rs index b78b13325..8353ac450 100644 --- a/nodedb/src/engine/sparse/inverted/core.rs +++ b/nodedb/src/engine/sparse/inverted/core.rs @@ -77,6 +77,24 @@ impl InvertedIndex { .purge_collection(database_id, tid.as_u64(), collection) .map_err(into_result_err) } + + /// Empty every index of a `(database, tenant, collection)` in one write + /// transaction, keeping the collection's analyzer, language, and fuzzy + /// configuration. TRUNCATE empties a collection this way. + pub fn clear_collection( + &self, + database_id: u64, + tid: TenantId, + collection: &str, + ) -> crate::Result { + self.note_purge(database_id, tid.as_u64(), Some(collection)); + self.inner + .memtable() + .drain_collection(database_id, tid.as_u64(), collection); + self.inner + .backend() + .clear_collection_data(database_id, tid.as_u64(), collection) + } } #[cfg(test)] @@ -86,6 +104,7 @@ mod tests { use nodedb_types::Surrogate; use super::*; + use crate::engine::sparse::inverted::test_support::body; const DB: u64 = 0; @@ -98,14 +117,61 @@ mod tests { (idx, dir) } + /// Clearing a collection empties its indexes and keeps its fuzzy + /// configuration. + #[test] + fn clear_collection_empties_indexes_and_keeps_config() { + let (idx, _dir) = open_temp(); + let t = TenantId::new(1); + idx.set_collection_fuzzy(DB, t, "docs", true).unwrap(); + idx.index_document(DB, t, "docs", Surrogate::new(1), &body("alpha bravo")) + .unwrap(); + idx.index_document(DB, t, "other", Surrogate::new(1), &body("alpha")) + .unwrap(); + + idx.clear_collection(DB, t, "docs").unwrap(); + + let search = |collection: &str| { + idx.search( + DB, + t, + collection, + FtsSearchParams { + query: "alpha", + top_k: 10, + fuzzy_enabled: false, + mode: QueryMode::And, + prefilter: None, + }, + ) + .unwrap() + }; + assert!( + search("docs").is_empty(), + "the cleared collection holds no text" + ); + assert_eq!( + search("other").len(), + 1, + "another collection keeps its text" + ); + assert_eq!(idx.corpus_stats(DB, t, "docs").unwrap().0, 0); + assert!( + idx.inner + .get_collection_fuzzy(DB, t.as_u64(), "docs") + .unwrap(), + "the fuzzy configuration survives" + ); + } + #[test] fn purge_tenant_structurally_drops_data() { let (idx, _dir) = open_temp(); let t1 = TenantId::new(1); let t2 = TenantId::new(2); - idx.index_document(DB, t1, "docs", Surrogate::new(1), "alpha bravo") + idx.index_document(DB, t1, "docs", Surrogate::new(1), &body("alpha bravo")) .unwrap(); - idx.index_document(DB, t2, "docs", Surrogate::new(1), "alpha bravo") + idx.index_document(DB, t2, "docs", Surrogate::new(1), &body("alpha bravo")) .unwrap(); idx.purge_tenant(DB, t1).unwrap(); diff --git a/nodedb/src/engine/sparse/inverted/corpus_stats.rs b/nodedb/src/engine/sparse/inverted/corpus_stats.rs index d5be98dac..686ff4f39 100644 --- a/nodedb/src/engine/sparse/inverted/corpus_stats.rs +++ b/nodedb/src/engine/sparse/inverted/corpus_stats.rs @@ -7,40 +7,67 @@ //! base search used — a staged doc must not shift the corpus itself, only //! be scored against it. +use nodedb_fts::IndexScope; use nodedb_fts::backend::FtsBackend; use nodedb_types::TenantId; use super::core::InvertedIndex; impl InvertedIndex { - /// Total document count and average document length for a collection, - /// as read by the base BM25 search (`FtsIndex::index_stats`). - pub fn corpus_stats( + /// Total document count and average document length of one index, as + /// read by the base BM25 search (`FtsIndex::index_stats`). + pub fn corpus_stats<'a>( &self, database_id: u64, tid: TenantId, - collection: &str, + index: impl Into>, ) -> crate::Result<(u32, f32)> { - self.inner - .index_stats(database_id, tid.as_u64(), collection) + self.inner.index_stats(database_id, tid.as_u64(), index) } /// Document frequency (number of documents containing `term`) for a - /// single already-analyzed term, read from the same POSTINGS table the - /// base search scores against. - pub fn term_df( + /// single already-analyzed term in one index, read from the same POSTINGS + /// table the base search scores against. + pub fn term_df<'a>( &self, database_id: u64, tid: TenantId, - collection: &str, + index: impl Into>, term: &str, ) -> crate::Result { let postings = self.inner .backend() - .read_postings(database_id, tid.as_u64(), collection, term)?; + .read_postings(database_id, tid.as_u64(), index.into(), term)?; Ok(postings.len() as u32) } + + /// Every document one index holds, read in one transaction. An index + /// holds a document exactly when it records the document's length. + pub fn index_members<'a>( + &self, + database_id: u64, + tid: TenantId, + index: impl Into>, + ) -> crate::Result> { + self.inner + .backend() + .index_members(database_id, tid.as_u64(), index.into()) + } + + /// Whether any document currently holds text in one index. + pub fn has_text<'a>( + &self, + database_id: u64, + tid: TenantId, + index: impl Into>, + ) -> crate::Result { + let (count, _) = + self.inner + .backend() + .collection_stats(database_id, tid.as_u64(), index.into())?; + Ok(count > 0) + } } #[cfg(test)] @@ -52,6 +79,7 @@ mod tests { use nodedb_types::Surrogate; use super::*; + use crate::engine::sparse::inverted::test_support::body; const DB: u64 = 0; const T: TenantId = TenantId::new(1); @@ -67,6 +95,26 @@ mod tests { (idx, dir) } + #[test] + fn index_members_are_the_documents_with_a_recorded_length() { + let (idx, _dir) = open_temp(); + idx.index_document(DB, T, "docs", Surrogate::new(1), &body("alpha bravo")) + .unwrap(); + idx.index_document(DB, T, "docs", Surrogate::new(7), &body("charlie")) + .unwrap(); + idx.index_document(DB, T, "other", Surrogate::new(3), &body("delta")) + .unwrap(); + + let members = idx.index_members(DB, T, "docs").unwrap(); + assert_eq!( + members, + [Surrogate::new(1), Surrogate::new(7)] + .into_iter() + .collect::>() + ); + assert!(idx.index_members(DB, T, "absent").unwrap().is_empty()); + } + /// STATS (doc count, avg doc length) must not double-count when the SAME /// surrogate is indexed more than once — this is exactly what happens on /// WAL replay, which re-invokes `index_document` for already-durable @@ -79,8 +127,14 @@ mod tests { let (idx, _dir) = open_temp(); // "alpha bravo charlie" tokenizes to 3 terms. - idx.index_document(DB, T, "docs", Surrogate::new(1), "alpha bravo charlie") - .unwrap(); + idx.index_document( + DB, + T, + "docs", + Surrogate::new(1), + &body("alpha bravo charlie"), + ) + .unwrap(); let (count, avg_len) = idx.corpus_stats(DB, T, "docs").unwrap(); assert_eq!(count, 1, "first index must count the doc once"); assert_eq!(avg_len, 3.0, "avg doc len == the single doc's length"); @@ -88,8 +142,14 @@ mod tests { // Re-index the SAME surrogate with IDENTICAL content, simulating a WAL // replay of an already-durable FtsIndex record. Doc count and total // token sum must be unchanged (net zero), not doubled. - idx.index_document(DB, T, "docs", Surrogate::new(1), "alpha bravo charlie") - .unwrap(); + idx.index_document( + DB, + T, + "docs", + Surrogate::new(1), + &body("alpha bravo charlie"), + ) + .unwrap(); let (count, avg_len) = idx.corpus_stats(DB, T, "docs").unwrap(); assert_eq!( count, 1, @@ -109,8 +169,14 @@ mod tests { fn reindex_same_surrogate_different_length_adjusts_total_by_delta() { let (idx, _dir) = open_temp(); - idx.index_document(DB, T, "docs", Surrogate::new(1), "alpha bravo charlie") - .unwrap(); + idx.index_document( + DB, + T, + "docs", + Surrogate::new(1), + &body("alpha bravo charlie"), + ) + .unwrap(); let (count, avg_len) = idx.corpus_stats(DB, T, "docs").unwrap(); assert_eq!(count, 1); assert_eq!(avg_len, 3.0); @@ -122,7 +188,7 @@ mod tests { T, "docs", Surrogate::new(1), - "alpha bravo charlie delta echo", + &body("alpha bravo charlie delta echo"), ) .unwrap(); let (count, avg_len) = idx.corpus_stats(DB, T, "docs").unwrap(); diff --git a/nodedb/src/engine/sparse/inverted/doc_fields.rs b/nodedb/src/engine/sparse/inverted/doc_fields.rs new file mode 100644 index 000000000..28d600a2a --- /dev/null +++ b/nodedb/src/engine/sparse/inverted/doc_fields.rs @@ -0,0 +1,84 @@ +// SPDX-License-Identifier: BUSL-1.1 + +//! Per-document field sets: the field indexes a document holds postings in. +//! +//! A re-index sees only the new text's fields. A field the previous version +//! filled and the new one does not would keep its postings without this +//! record. The whole-document index is not listed: every indexed document +//! is tracked there by its `DOC_LENGTHS` row. + +use std::collections::BTreeSet; + +use redb::{ReadableTable as _, WriteTransaction}; + +use super::errors::inverted_err; +use super::indexing::IndexDocScope; +use crate::engine::sparse::fts_redb::keys::doc_fields_key; +use crate::engine::sparse::fts_redb::tables::DOC_FIELDS; + +/// The non-empty field names the document holds postings in. Empty when no +/// row is stored. +pub(super) fn read(txn: &WriteTransaction, doc: IndexDocScope<'_>) -> crate::Result> { + let table = txn + .open_table(DOC_FIELDS) + .map_err(|e| inverted_err("open doc_fields", e))?; + let key = doc_fields_key( + doc.database_id, + doc.tid.as_u64(), + doc.collection, + doc.surrogate, + ); + match table + .get(key) + .map_err(|e| inverted_err("read doc_fields", e))? + { + Some(v) => zerompk::from_msgpack::>(v.value()) + .map_err(|e| inverted_err("decode doc_fields", e)), + None => Ok(Vec::new()), + } +} + +/// Record `fields` as the document's field set. An empty set removes the row. +pub(super) fn put( + txn: &WriteTransaction, + doc: IndexDocScope<'_>, + fields: &BTreeSet<&str>, +) -> crate::Result<()> { + if fields.is_empty() { + return clear(txn, doc); + } + let encoded: Vec<&str> = fields.iter().copied().collect(); + let bytes = + zerompk::to_msgpack_vec(&encoded).map_err(|e| inverted_err("serialize doc_fields", e))?; + let mut table = txn + .open_table(DOC_FIELDS) + .map_err(|e| inverted_err("open doc_fields", e))?; + table + .insert( + doc_fields_key( + doc.database_id, + doc.tid.as_u64(), + doc.collection, + doc.surrogate, + ), + bytes.as_slice(), + ) + .map_err(|e| inverted_err("insert doc_fields", e))?; + Ok(()) +} + +/// Drop the document's field-set row. +pub(super) fn clear(txn: &WriteTransaction, doc: IndexDocScope<'_>) -> crate::Result<()> { + let mut table = txn + .open_table(DOC_FIELDS) + .map_err(|e| inverted_err("open doc_fields", e))?; + table + .remove(doc_fields_key( + doc.database_id, + doc.tid.as_u64(), + doc.collection, + doc.surrogate, + )) + .map_err(|e| inverted_err("remove doc_fields", e))?; + Ok(()) +} diff --git a/nodedb/src/engine/sparse/inverted/doc_image.rs b/nodedb/src/engine/sparse/inverted/doc_image.rs index 98e22951d..233230630 100644 --- a/nodedb/src/engine/sparse/inverted/doc_image.rs +++ b/nodedb/src/engine/sparse/inverted/doc_image.rs @@ -3,34 +3,34 @@ //! A document's complete index footprint, captured so a rollback can put it //! back exactly. //! -//! The index does not keep a document's text. It keeps, per term, the -//! positions the term occupies in the document's analyzed token stream, and -//! the stream's length. Those positions are dense (`0..len`), so the token -//! stream rebuilds from them exactly. Re-indexing that stream reproduces the -//! document's postings, its length row, its term set, and the collection's -//! corpus counters. +//! The index does not keep a document's text. It keeps, per index and per +//! term, the positions the term occupies in the document's analyzed token +//! stream, and the stream's length. Those positions are dense (`0..len`), so +//! each index's token stream rebuilds from them exactly. Re-indexing those +//! streams reproduces the document's postings, its length rows, its term +//! sets, its field set, and every index's corpus counters. +use std::collections::BTreeSet; + +use nodedb_fts::IndexScope; use nodedb_fts::posting::Posting; use nodedb_types::{Surrogate, TenantId}; use redb::ReadableTable as _; use super::core::InvertedIndex; +use super::doc_fields; use super::doc_terms; use super::errors::inverted_err; use super::indexing::{IndexDocScope, prior_doc_length}; +use crate::engine::sparse::fts_redb::keys::posting_key; use crate::engine::sparse::fts_redb::tables::POSTINGS; -/// The analyzed token stream one document is indexed with. +/// The analyzed token streams one document is indexed with, one per index. #[derive(Debug, Clone, PartialEq, Eq)] pub struct FtsDocImage { - tokens: Vec, -} - -impl FtsDocImage { - /// The analyzed token stream, in document order. - pub(super) fn tokens(&self) -> &[String] { - &self.tokens - } + /// `(field_key, tokens)` per index the document holds postings in. The + /// empty field key is the whole-document index. + scopes: Vec<(String, Vec)>, } impl InvertedIndex { @@ -44,7 +44,7 @@ impl InvertedIndex { collection: &str, surrogate: Surrogate, ) -> crate::Result> { - let scope = IndexDocScope { + let doc = IndexDocScope { database_id, tid, collection, @@ -52,7 +52,7 @@ impl InvertedIndex { }; let db = self.inner.backend().db(); let txn = db.begin_write().map_err(|e| inverted_err("write txn", e))?; - let image = Self::read_document_image(&txn, scope)?; + let image = Self::read_document_image(&txn, doc)?; txn.abort() .map_err(|e| inverted_err("abort image read", e))?; Ok(image) @@ -68,10 +68,10 @@ impl InvertedIndex { surrogate: Surrogate, image: Option<&FtsDocImage>, ) -> crate::Result<()> { - let Some(image) = image.filter(|image| !image.tokens.is_empty()) else { + let Some(image) = image.filter(|image| !image.scopes.is_empty()) else { return self.remove_document(database_id, tid, collection, surrogate); }; - let scope = IndexDocScope { + let doc = IndexDocScope { database_id, tid, collection, @@ -79,59 +79,99 @@ impl InvertedIndex { }; let db = self.inner.backend().db(); let txn = db.begin_write().map_err(|e| inverted_err("write txn", e))?; - self.write_index_data(&txn, scope, &image.tokens)?; + self.write_document_image(&txn, doc, image)?; txn.commit() .map_err(|e| inverted_err("commit image restore", e))?; Ok(()) } + /// Replace the document's footprint in every index with `image`. + pub(super) fn write_document_image( + &self, + txn: &redb::WriteTransaction, + doc: IndexDocScope<'_>, + image: &FtsDocImage, + ) -> crate::Result<()> { + self.remove_document_in_txn(txn, doc)?; + let mut fields: BTreeSet<&str> = BTreeSet::new(); + for (field, tokens) in &image.scopes { + let index = IndexScope::from_key(doc.collection, field); + self.write_index_data(txn, doc, index, tokens)?; + if let Some(name) = index.field_name() { + fields.insert(name); + } + } + doc_fields::put(txn, doc, &fields) + } + + /// The footprint of `doc` across the whole-document index and every field + /// index its field set names. `None` when it holds no postings anywhere. pub(super) fn read_document_image( txn: &redb::WriteTransaction, - scope: IndexDocScope<'_>, + doc: IndexDocScope<'_>, ) -> crate::Result> { - let Some(doc_len) = prior_doc_length(txn, scope)? else { - return Ok(None); - }; - let terms = doc_terms::occupied_terms(txn, scope, true)?; - let postings = txn - .open_table(POSTINGS) - .map_err(|e| inverted_err("open postings", e))?; - let mut tokens = vec![String::new(); doc_len as usize]; - for term in terms { - let list: Vec = postings - .get(( - scope.database_id, - scope.tid.as_u64(), - scope.collection, - term.as_str(), - )) - .map_err(|e| inverted_err("read postings", e))? - .map(|bytes| zerompk::from_msgpack(bytes.value())) - .transpose() - .map_err(|e| inverted_err("decode postings", e))? - .unwrap_or_default(); - let Some(posting) = list.iter().find(|p| p.doc_id == scope.surrogate) else { - continue; - }; - for &position in &posting.positions { - let slot = tokens.get_mut(position as usize).ok_or_else(|| { - inverted_err( - "document image", - format!( - "term '{term}' sits at position {position} past the document \ - length {doc_len}" - ), - ) - })?; - slot.clone_from(&term); + let mut keys = vec![String::new()]; + keys.extend(doc_fields::read(txn, doc)?); + let mut scopes = Vec::with_capacity(keys.len()); + for key in keys { + let index = IndexScope::from_key(doc.collection, &key); + if let Some(tokens) = read_scope_tokens(txn, doc, index)? { + scopes.push((key, tokens)); } } - if let Some(position) = tokens.iter().position(String::is_empty) { - return Err(inverted_err( - "document image", - format!("no posting covers position {position} of {doc_len}"), - )); + Ok((!scopes.is_empty()).then_some(FtsDocImage { scopes })) + } +} + +/// The analyzed token stream of `doc` in one index, or `None` when it is not +/// in that index. +fn read_scope_tokens( + txn: &redb::WriteTransaction, + doc: IndexDocScope<'_>, + index: IndexScope<'_>, +) -> crate::Result>> { + let Some(doc_len) = prior_doc_length(txn, doc, index)? else { + return Ok(None); + }; + let terms = doc_terms::occupied_terms(txn, doc, index, true)?; + let postings = txn + .open_table(POSTINGS) + .map_err(|e| inverted_err("open postings", e))?; + let mut tokens = vec![String::new(); doc_len as usize]; + for term in terms { + let list: Vec = postings + .get(posting_key( + doc.database_id, + doc.tid.as_u64(), + index, + term.as_str(), + )) + .map_err(|e| inverted_err("read postings", e))? + .map(|bytes| zerompk::from_msgpack(bytes.value())) + .transpose() + .map_err(|e| inverted_err("decode postings", e))? + .unwrap_or_default(); + let Some(posting) = list.iter().find(|p| p.doc_id == doc.surrogate) else { + continue; + }; + for &position in &posting.positions { + let slot = tokens.get_mut(position as usize).ok_or_else(|| { + inverted_err( + "document image", + format!( + "term '{term}' sits at position {position} past the document \ + length {doc_len}" + ), + ) + })?; + slot.clone_from(&term); } - Ok(Some(FtsDocImage { tokens })) } + if let Some(position) = tokens.iter().position(String::is_empty) { + return Err(inverted_err( + "document image", + format!("no posting covers position {position} of {doc_len}"), + )); + } + Ok(Some(tokens)) } diff --git a/nodedb/src/engine/sparse/inverted/doc_terms.rs b/nodedb/src/engine/sparse/inverted/doc_terms.rs index b15c90f17..8c6e03ef7 100644 --- a/nodedb/src/engine/sparse/inverted/doc_terms.rs +++ b/nodedb/src/engine/sparse/inverted/doc_terms.rs @@ -1,7 +1,7 @@ // SPDX-License-Identifier: BUSL-1.1 -//! Per-document term sets: the record of which posting lists a document is -//! currently a member of. +//! Per-document term sets: the record of which posting lists of one index a +//! document is currently a member of. //! //! ## Why this exists //! @@ -15,9 +15,9 @@ //! //! Fixing that needs the old term set. Two ways to obtain it: //! -//! * Range-scan every posting list in the collection and strip the surrogate +//! * Range-scan every posting list in the index and strip the surrogate //! from any list that still names it. Needs no extra storage, but costs -//! O(terms in the collection) on every single re-index — an unbounded, +//! O(terms in the index) on every single re-index — an unbounded, //! corpus-sized cost on the write path. //! * Persist the document's own term set beside its `DOC_LENGTHS` row, so a //! re-index removes exactly the terms that dropped out: O(terms in the old @@ -25,61 +25,59 @@ //! //! The second wins: the cost scales with the document being written rather //! than with the corpus it lands in. Storage cost is one row per indexed -//! document holding its DISTINCT analyzed terms — bounded by the document's -//! own text (post-stemming, deduplicated, so typically well under the raw -//! field bytes) and independent of corpus size. +//! document per index holding its DISTINCT analyzed terms — bounded by the +//! document's own text (post-stemming, deduplicated, so typically well under +//! the raw field bytes) and independent of corpus size. //! -//! The scan survives only as the bounded fallback for documents indexed -//! before term sets were recorded: it runs at most once per such document -//! (the re-index that triggers it writes the term set, so the next one takes -//! the targeted path) and never for a document that is not already in the -//! index — `stored` is consulted only when `DOC_LENGTHS` proves a prior -//! index exists. +//! The scan survives only as the bounded fallback for a document whose term +//! set row is missing: it runs at most once per such document (the re-index +//! that triggers it writes the term set, so the next one takes the targeted +//! path) and never for a document that is not already in the index — +//! `stored` is consulted only when `DOC_LENGTHS` proves a prior index exists. use std::collections::BTreeSet; use redb::{ReadableTable as _, WriteTransaction}; +use nodedb_fts::IndexScope; use nodedb_fts::posting::Posting; use super::errors::inverted_err; use super::indexing::IndexDocScope; +use crate::engine::sparse::fts_redb::keys::{KeyOwner, doc_key, posting_key}; +use crate::engine::sparse::fts_redb::scan; use crate::engine::sparse::fts_redb::tables::{DOC_TERMS, POSTINGS}; -/// Upper-bound sentinel for the `term` component of a -/// `(database_id, tid, collection, term)` range scan: the highest scalar -/// value representable in UTF-8, so every real term sorts below it. -const MAX_TERM: &str = "\u{10ffff}"; - -/// The term set a previous index of this surrogate recorded, or `None` when -/// none was ever stored (a document first indexed by a build that predates -/// the term set). +/// The term set a previous index of this surrogate recorded in `index`, or +/// `None` when none is stored. pub(super) fn stored( txn: &WriteTransaction, - scope: IndexDocScope<'_>, + doc: IndexDocScope<'_>, + index: IndexScope<'_>, ) -> crate::Result>> { let table = txn .open_table(DOC_TERMS) .map_err(|e| inverted_err("open doc_terms", e))?; - let key = ( - scope.database_id, - scope.tid.as_u64(), - scope.collection, - scope.surrogate.as_u32(), - ); - let stored = table + let key = doc_key(doc.database_id, doc.tid.as_u64(), index, doc.surrogate); + match table .get(key) .map_err(|e| inverted_err("read doc_terms", e))? - .and_then(|v| zerompk::from_msgpack::>(v.value()).ok()); - Ok(stored) + { + Some(v) => zerompk::from_msgpack::>(v.value()) + .map(Some) + .map_err(|e| inverted_err("decode doc_terms", e)), + None => Ok(None), + } } -/// Record `terms` as the document's current term set, replacing any prior -/// one. Must be called in the same transaction as the posting writes it -/// describes, so a crash can never leave the set disagreeing with POSTINGS. +/// Record `terms` as the document's current term set in `index`, replacing +/// any prior one. Must be called in the same transaction as the posting +/// writes it describes, so a crash can never leave the set disagreeing with +/// POSTINGS. pub(super) fn put( txn: &WriteTransaction, - scope: IndexDocScope<'_>, + doc: IndexDocScope<'_>, + index: IndexScope<'_>, terms: &BTreeSet<&str>, ) -> crate::Result<()> { let encoded: Vec<&str> = terms.iter().copied().collect(); @@ -90,73 +88,70 @@ pub(super) fn put( .map_err(|e| inverted_err("open doc_terms", e))?; table .insert( - ( - scope.database_id, - scope.tid.as_u64(), - scope.collection, - scope.surrogate.as_u32(), - ), + doc_key(doc.database_id, doc.tid.as_u64(), index, doc.surrogate), bytes.as_slice(), ) .map_err(|e| inverted_err("insert doc_terms", e))?; Ok(()) } -/// Drop the document's term-set row. Paired with the removal of its -/// `DOC_LENGTHS` row so the two never disagree about whether the document is -/// in the index. -pub(super) fn clear(txn: &WriteTransaction, scope: IndexDocScope<'_>) -> crate::Result<()> { +/// Drop the document's term-set row in `index`. Paired with the removal of +/// its `DOC_LENGTHS` row so the two never disagree about whether the +/// document is in the index. +pub(super) fn clear( + txn: &WriteTransaction, + doc: IndexDocScope<'_>, + index: IndexScope<'_>, +) -> crate::Result<()> { let mut table = txn .open_table(DOC_TERMS) .map_err(|e| inverted_err("open doc_terms", e))?; table - .remove(( - scope.database_id, - scope.tid.as_u64(), - scope.collection, - scope.surrogate.as_u32(), + .remove(doc_key( + doc.database_id, + doc.tid.as_u64(), + index, + doc.surrogate, )) .map_err(|e| inverted_err("remove doc_terms", e))?; Ok(()) } -/// Every term the collection currently has a posting list for. +/// Every term `index` currently has a posting list for. /// /// The fallback source of "which lists might name this surrogate" for a /// document that has no stored term set. Callers must gate it on the document /// actually being indexed — see the module docs for why this is not the /// general path. -pub(super) fn collection_terms( +pub(super) fn index_terms( txn: &WriteTransaction, - scope: IndexDocScope<'_>, + doc: IndexDocScope<'_>, + index: IndexScope<'_>, ) -> crate::Result> { let table = txn .open_table(POSTINGS) .map_err(|e| inverted_err("open postings", e))?; - let t = scope.tid.as_u64(); - let terms = table - .range( - (scope.database_id, t, scope.collection, "") - ..=(scope.database_id, t, scope.collection, MAX_TERM), - ) - .map_err(|e| inverted_err("postings range", e))? - .filter_map(|r| r.ok().map(|(k, _)| k.value().3.to_string())) - .collect(); - Ok(terms) + let keys = scan::str_keys( + &table, + KeyOwner::index(doc.database_id, doc.tid.as_u64(), index), + ) + .map_err(|e| inverted_err("postings range", e))?; + Ok(keys.into_iter().map(|(_, _, term)| term).collect()) } -/// Remove the surrogate's posting from each of `terms`, deleting a list that -/// this empties so the term's `df` returns to its true value instead of -/// counting a document that no longer contains it. +/// Remove the surrogate's posting from each of `terms` in `index`, deleting +/// a list that this empties so the term's `df` returns to its true value +/// instead of counting a document that no longer contains it. pub(super) fn strip_postings( txn: &WriteTransaction, - scope: IndexDocScope<'_>, + doc: IndexDocScope<'_>, + index: IndexScope<'_>, terms: &[String], ) -> crate::Result<()> { if terms.is_empty() { return Ok(()); } - let t = scope.tid.as_u64(); + let t = doc.tid.as_u64(); let mut table = txn .open_table(POSTINGS) .map_err(|e| inverted_err("open postings", e))?; @@ -166,17 +161,18 @@ pub(super) fn strip_postings( // interleaved with the reads. let mut updates: Vec<(&str, Option>)> = Vec::new(); for term in terms { - let key = (scope.database_id, t, scope.collection, term.as_str()); - let Some(existing) = table + let key = posting_key(doc.database_id, t, index, term.as_str()); + let existing = match table .get(key) .map_err(|e| inverted_err("read postings", e))? - .and_then(|v| zerompk::from_msgpack::>(v.value()).ok()) - else { - continue; + { + Some(v) => zerompk::from_msgpack::>(v.value()) + .map_err(|e| inverted_err("decode postings", e))?, + None => continue, }; let mut remaining = existing; let before = remaining.len(); - remaining.retain(|p| p.doc_id != scope.surrogate); + remaining.retain(|p| p.doc_id != doc.surrogate); if remaining.len() == before { continue; } @@ -190,7 +186,7 @@ pub(super) fn strip_postings( } for (term, new_value) in updates { - let key = (scope.database_id, t, scope.collection, term); + let key = posting_key(doc.database_id, t, index, term); match new_value { None => { table @@ -207,22 +203,67 @@ pub(super) fn strip_postings( Ok(()) } -/// The terms a document currently occupies, for a caller that is about to -/// remove it from some or all of them. +/// Remove every surrogate `removals` names from its term's posting list in +/// `index`, rewriting each list once. A list this empties is deleted. +pub(super) fn strip_postings_many( + txn: &WriteTransaction, + database_id: u64, + tid: u64, + index: IndexScope<'_>, + removals: &std::collections::BTreeMap< + String, + std::collections::HashSet, + >, +) -> crate::Result<()> { + let mut table = txn + .open_table(POSTINGS) + .map_err(|e| inverted_err("open postings", e))?; + for (term, docs) in removals { + let key = posting_key(database_id, tid, index, term.as_str()); + let remaining = match table + .get(key) + .map_err(|e| inverted_err("read postings", e))? + { + Some(v) => { + let mut postings = zerompk::from_msgpack::>(v.value()) + .map_err(|e| inverted_err("decode postings", e))?; + postings.retain(|p| !docs.contains(&p.doc_id)); + postings + } + None => continue, + }; + if remaining.is_empty() { + table + .remove(key) + .map_err(|e| inverted_err("remove posting", e))?; + } else { + let bytes = zerompk::to_msgpack_vec(&remaining) + .map_err(|e| inverted_err("serialize postings", e))?; + table + .insert(key, bytes.as_slice()) + .map_err(|e| inverted_err("update posting", e))?; + } + } + Ok(()) +} + +/// The terms a document currently occupies in `index`, for a caller that is +/// about to remove it from some or all of them. /// /// Returns the stored term set when there is one, and otherwise falls back to -/// the collection's full term list — bounded to documents that a prior index +/// the index's full term list — bounded to documents that a prior index /// actually recorded, which the caller proves by passing `previously_indexed`. pub(super) fn occupied_terms( txn: &WriteTransaction, - scope: IndexDocScope<'_>, + doc: IndexDocScope<'_>, + index: IndexScope<'_>, previously_indexed: bool, ) -> crate::Result> { if !previously_indexed { return Ok(Vec::new()); } - match stored(txn, scope)? { + match stored(txn, doc, index)? { Some(terms) => Ok(terms), - None => collection_terms(txn, scope), + None => index_terms(txn, doc, index), } } diff --git a/nodedb/src/engine/sparse/inverted/document.rs b/nodedb/src/engine/sparse/inverted/document.rs new file mode 100644 index 000000000..311f2191d --- /dev/null +++ b/nodedb/src/engine/sparse/inverted/document.rs @@ -0,0 +1,275 @@ +// SPDX-License-Identifier: BUSL-1.1 + +//! Per-document orchestration across a collection's indexes: the +//! whole-document index and one index per top-level string field. + +use std::collections::BTreeSet; + +use redb::WriteTransaction; + +use nodedb_fts::{DocumentText, IndexScope}; + +use super::core::InvertedIndex; +use super::doc_fields; +use super::errors::inverted_err; +use super::indexing::IndexDocScope; + +impl InvertedIndex { + /// Replace a document's footprint in every index: the whole-document index + /// and one index per top-level string field. A field the previous version + /// held and this one does not is retracted from that field's index. + pub fn index_document_in_txn( + &self, + txn: &WriteTransaction, + doc: IndexDocScope<'_>, + text: &DocumentText, + ) -> crate::Result<()> { + self.note_doc_write(doc); + let prior = doc_fields::read(txn, doc)?; + self.index_scope_in_txn( + txn, + doc, + IndexScope::document(doc.collection), + &text.whole(), + )?; + let mut held: BTreeSet<&str> = BTreeSet::new(); + for (index, value) in text.field_scopes(doc.collection) { + if self.index_scope_in_txn(txn, doc, index, value)? + && let Some(name) = index.field_name() + { + held.insert(name); + } + } + for gone in prior.iter().filter(|f| !held.contains(f.as_str())) { + let index = IndexScope::field(doc.collection, gone) + .ok_or_else(|| inverted_err("doc_fields", "stored field name is empty"))?; + self.remove_scope_in_txn(txn, doc, index)?; + } + doc_fields::put(txn, doc, &held) + } + + /// Index `text` under one scope. Text with no terms removes the document + /// from that scope. Returns whether the document holds postings there. + fn index_scope_in_txn( + &self, + txn: &WriteTransaction, + doc: IndexDocScope<'_>, + index: IndexScope<'_>, + text: &str, + ) -> crate::Result { + let tokens = self.analyze_for_collection(doc.database_id, doc.tid, doc.collection, text)?; + if tokens.is_empty() { + self.remove_scope_in_txn(txn, doc, index)?; + return Ok(false); + } + self.write_index_data(txn, doc, index, &tokens)?; + Ok(true) + } +} + +#[cfg(test)] +mod tests { + use std::sync::Arc; + + use redb::{Database, ReadableDatabase as _, ReadableTable as _}; + + use nodedb_fts::FtsSearchParams; + use nodedb_fts::posting::QueryMode; + use nodedb_types::{Surrogate, TenantId}; + + use super::*; + use crate::engine::sparse::fts_redb::keys::doc_fields_key; + use crate::engine::sparse::fts_redb::tables::DOC_FIELDS; + use crate::engine::sparse::inverted::test_support::fields; + + const DB: u64 = 0; + const T: TenantId = TenantId::new(1); + const DOCS: IndexScope<'static> = IndexScope::document("docs"); + + fn open_temp() -> (InvertedIndex, tempfile::TempDir) { + let dir = tempfile::tempdir().unwrap(); + let path = dir.path().join("test-inverted.redb"); + let db = Arc::new(crate::engine::durability_gate::GatedDatabase::new( + Database::create(&path).unwrap(), + )); + let idx = + InvertedIndex::open(db, crate::data::executor::core_loop::test_governor()).unwrap(); + (idx, dir) + } + + fn scope(field: &str) -> IndexScope<'_> { + IndexScope::field("docs", field).unwrap() + } + + fn hits(idx: &InvertedIndex, index: IndexScope<'_>, query: &str) -> Vec { + let mut ids: Vec = idx + .search( + DB, + T, + index, + FtsSearchParams { + query, + top_k: 100, + fuzzy_enabled: false, + mode: QueryMode::Or, + prefilter: None, + }, + ) + .unwrap() + .into_iter() + .map(|r| r.doc_id.as_u32()) + .collect(); + ids.sort_unstable(); + ids + } + + fn doc_fields_row(idx: &InvertedIndex, surrogate: u32) -> Option> { + let txn = idx.backend().db().begin_read().unwrap(); + let table = txn.open_table(DOC_FIELDS).unwrap(); + table + .get(doc_fields_key( + DB, + T.as_u64(), + "docs", + Surrogate::new(surrogate), + )) + .unwrap() + .map(|v| zerompk::from_msgpack::>(v.value()).unwrap()) + } + + fn crossed(idx: &InvertedIndex) { + idx.index_document( + DB, + T, + "docs", + Surrogate::new(1), + &fields(&[("title", "rust"), ("body", "python guide")]), + ) + .unwrap(); + idx.index_document( + DB, + T, + "docs", + Surrogate::new(2), + &fields(&[("body", "rust handbook")]), + ) + .unwrap(); + } + + #[test] + fn a_field_search_returns_only_documents_holding_the_term_in_that_field() { + let (idx, _dir) = open_temp(); + crossed(&idx); + assert_eq!(hits(&idx, scope("title"), "rust"), vec![1]); + assert_eq!(hits(&idx, scope("body"), "rust"), vec![2]); + assert_eq!(hits(&idx, DOCS, "rust"), vec![1, 2]); + } + + #[test] + fn a_field_index_counts_only_documents_holding_the_field() { + let (idx, _dir) = open_temp(); + crossed(&idx); + assert_eq!(idx.corpus_stats(DB, T, scope("title")).unwrap(), (1, 1.0)); + assert_eq!(idx.corpus_stats(DB, T, scope("body")).unwrap(), (2, 2.0)); + assert_eq!(idx.corpus_stats(DB, T, DOCS).unwrap(), (2, 2.5)); + } + + #[test] + fn moving_a_word_between_fields_retracts_it_and_keeps_the_whole_document() { + let (idx, _dir) = open_temp(); + idx.index_document( + DB, + T, + "docs", + Surrogate::new(1), + &fields(&[("body", "guide"), ("title", "rust")]), + ) + .unwrap(); + let whole_before = ( + idx.corpus_stats(DB, T, DOCS).unwrap(), + hits(&idx, DOCS, "rust"), + hits(&idx, DOCS, "guide"), + ); + + idx.index_document( + DB, + T, + "docs", + Surrogate::new(1), + &fields(&[("body", "guide rust")]), + ) + .unwrap(); + + assert!(hits(&idx, scope("title"), "rust").is_empty()); + assert_eq!(idx.corpus_stats(DB, T, scope("title")).unwrap().0, 0); + assert_eq!(hits(&idx, scope("body"), "rust"), vec![1]); + assert_eq!(idx.term_df(DB, T, scope("title"), "rust").unwrap(), 0); + let whole_after = ( + idx.corpus_stats(DB, T, DOCS).unwrap(), + hits(&idx, DOCS, "rust"), + hits(&idx, DOCS, "guide"), + ); + assert_eq!(whole_before, whole_after); + assert_eq!(doc_fields_row(&idx, 1), Some(vec!["body".to_string()])); + } + + #[test] + fn delete_removes_every_scope_and_the_field_set() { + let (idx, _dir) = open_temp(); + crossed(&idx); + assert_eq!( + doc_fields_row(&idx, 1), + Some(vec!["body".to_string(), "title".to_string()]) + ); + + idx.remove_document(DB, T, "docs", Surrogate::new(1)) + .unwrap(); + + assert!(hits(&idx, scope("title"), "rust").is_empty()); + assert!(hits(&idx, scope("body"), "python").is_empty()); + assert_eq!(hits(&idx, DOCS, "rust"), vec![2]); + assert_eq!(idx.corpus_stats(DB, T, scope("title")).unwrap().0, 0); + assert_eq!(idx.corpus_stats(DB, T, scope("body")).unwrap().0, 1); + assert_eq!(idx.corpus_stats(DB, T, DOCS).unwrap().0, 1); + assert_eq!(doc_fields_row(&idx, 1), None); + } + + #[test] + fn purge_collection_removes_field_scopes() { + let (idx, _dir) = open_temp(); + crossed(&idx); + idx.index_document( + DB, + T, + "other", + Surrogate::new(3), + &fields(&[("title", "rust")]), + ) + .unwrap(); + + idx.purge_collection(DB, T, "docs").unwrap(); + + for index in [scope("title"), scope("body"), DOCS] { + assert!(!idx.has_text(DB, T, index).unwrap()); + assert!(hits(&idx, index, "rust").is_empty()); + } + assert_eq!(doc_fields_row(&idx, 1), None); + assert_eq!(doc_fields_row(&idx, 2), None); + let other_title = IndexScope::field("other", "title").unwrap(); + assert!(idx.has_text(DB, T, other_title).unwrap()); + } + + #[test] + fn has_text_is_false_after_the_last_holder_is_deleted() { + let (idx, _dir) = open_temp(); + crossed(&idx); + assert!(idx.has_text(DB, T, scope("title")).unwrap()); + assert!(!idx.has_text(DB, T, scope("missing")).unwrap()); + + idx.remove_document(DB, T, "docs", Surrogate::new(1)) + .unwrap(); + + assert!(!idx.has_text(DB, T, scope("title")).unwrap()); + assert!(idx.has_text(DB, T, scope("body")).unwrap()); + } +} diff --git a/nodedb/src/engine/sparse/inverted/errors.rs b/nodedb/src/engine/sparse/inverted/errors.rs index 77698501f..e0cb37807 100644 --- a/nodedb/src/engine/sparse/inverted/errors.rs +++ b/nodedb/src/engine/sparse/inverted/errors.rs @@ -24,12 +24,20 @@ pub(super) fn fts_index_err(e: nodedb_fts::FtsIndexError) -> crate FtsIndexError::BudgetExhausted(_) => crate::Error::MemoryExhausted { engine: "fts".into(), }, - other @ (FtsIndexError::SurrogateOutOfRange { .. } | FtsIndexError::Segment(_)) => { - crate::Error::Storage { - engine: "inverted".into(), - detail: other.to_string(), - } - } + // On-disk index state that is wrong: a segment that fails + // validation, a listed segment that is gone, or a state blob that + // does not decode. + other @ (FtsIndexError::CorruptSegment { .. } + | FtsIndexError::MissingSegment { .. } + | FtsIndexError::CorruptState { .. }) => crate::Error::SegmentCorrupted { + detail: other.to_string(), + }, + other @ (FtsIndexError::SurrogateOutOfRange { .. } + | FtsIndexError::Segment(_) + | FtsIndexError::StateEncode { .. }) => crate::Error::Storage { + engine: "inverted".into(), + detail: other.to_string(), + }, // `FtsIndexError` is `#[non_exhaustive]` and lives in another crate, // so the compiler requires this arm. A variant this build cannot name // is a storage fault. @@ -76,4 +84,27 @@ mod tests { } )); } + + /// A corrupt or missing segment is on-disk corruption, not a generic + /// storage fault. + #[test] + fn a_corrupt_segment_is_segment_corruption() { + let corrupt: nodedb_fts::FtsIndexError = + nodedb_fts::FtsIndexError::CorruptSegment { + segment_id: "L0:0000000000000001".into(), + source: nodedb_fts::lsm::segment::error::SegmentError::Truncated, + }; + assert!(matches!( + fts_index_err(corrupt), + crate::Error::SegmentCorrupted { ref detail } if detail.contains("L0:0000000000000001") + )); + let missing: nodedb_fts::FtsIndexError = + nodedb_fts::FtsIndexError::MissingSegment { + segment_id: "L0:0000000000000002".into(), + }; + assert!(matches!( + fts_index_err(missing), + crate::Error::SegmentCorrupted { .. } + )); + } } diff --git a/nodedb/src/engine/sparse/inverted/indexing.rs b/nodedb/src/engine/sparse/inverted/indexing.rs index 3be9380f2..4613f9149 100644 --- a/nodedb/src/engine/sparse/inverted/indexing.rs +++ b/nodedb/src/engine/sparse/inverted/indexing.rs @@ -1,15 +1,17 @@ // SPDX-License-Identifier: BUSL-1.1 -//! Document indexing for the inverted index. +//! Per-index writes for the inverted index. //! //! All writes bypass the LSM memtable and go directly to the persistent //! POSTINGS / DOC_LENGTHS / DOC_TERMS / STATS tables so they can participate //! in the caller's redb write transaction (Origin transactional indexing). //! -//! A re-index is a full replacement of the document's index footprint, not an -//! overlay: the terms the new text no longer contains are retracted first -//! (see the `doc_terms` sibling module), then the new text's postings are -//! written. Removal lives in the `removal` sibling module. +//! A re-index of one index is a full replacement of the document's footprint +//! there, not an overlay: the terms the new text no longer contains are +//! retracted first (see the `doc_terms` sibling module), then the new text's +//! postings are written. The per-document orchestration across indexes lives +//! in the `document` sibling module. Removal lives in the `removal` sibling +//! module. use std::collections::{BTreeSet, HashMap}; @@ -17,11 +19,13 @@ use redb::{ReadableTable as _, WriteTransaction}; use tracing::debug; use nodedb_fts::posting::Posting; +use nodedb_fts::{DocumentText, IndexScope}; use nodedb_types::{Surrogate, TenantId}; use super::core::InvertedIndex; use super::doc_terms; use super::errors::inverted_err; +use crate::engine::sparse::fts_redb::keys::{doc_key, posting_key, stats_key}; use crate::engine::sparse::fts_redb::tables::{DOC_LENGTHS, POSTINGS, STATS}; /// `(database_id, tenant, collection, surrogate)` scope shared by the @@ -47,7 +51,8 @@ impl InvertedIndex { /// `index_document_in_txn`) and query-term canonicalization /// (`phrase_search`) all call through here so a document is always /// tokenized the same way it is later matched against, whether the write - /// is durable or still staged in an open transaction. + /// is durable or still staged in an open transaction. Every index of a + /// collection shares its analyzer. pub fn analyze_for_collection( &self, database_id: u64, @@ -87,73 +92,53 @@ impl InvertedIndex { .set_collection_fuzzy(database_id, tid.as_u64(), collection, fuzzy) } - /// Index a document's text content. + /// Index a document's text into its whole-document index and one index + /// per top-level string field, in its own write transaction. /// /// Text that analyzes to no terms is a removal, not a no-op: an update - /// that strips a document of every indexable word must take it out of the - /// index, or it keeps matching the words it used to contain. + /// that strips a document (or one field) of every indexable word must take + /// it out of that index, or it keeps matching the words it used to + /// contain. pub fn index_document( &self, database_id: u64, tid: TenantId, collection: &str, surrogate: Surrogate, - text: &str, + text: &DocumentText, ) -> crate::Result<()> { - let tokens = self.analyze_for_collection(database_id, tid, collection, text)?; - if tokens.is_empty() { - return self.remove_document(database_id, tid, collection, surrogate); - } - let scope = IndexDocScope { + let doc = IndexDocScope { database_id, tid, collection, surrogate, }; - let db = self.inner.backend().db(); let write_txn = db.begin_write().map_err(|e| inverted_err("write txn", e))?; - self.write_index_data(&write_txn, scope, &tokens)?; + self.index_document_in_txn(&write_txn, doc, text)?; write_txn .commit() .map_err(|e| inverted_err("commit index", e))?; Ok(()) } - /// Index a document within an externally-owned write transaction. - /// - /// Same empty-text semantics as [`Self::index_document`]. - pub fn index_document_in_txn( - &self, - txn: &WriteTransaction, - scope: IndexDocScope<'_>, - text: &str, - ) -> crate::Result<()> { - let tokens = - self.analyze_for_collection(scope.database_id, scope.tid, scope.collection, text)?; - if tokens.is_empty() { - return self.remove_document_in_txn(txn, scope); - } - self.write_index_data(txn, scope, &tokens) - } - - /// Core indexing logic: writes postings, doc length, and stats within - /// a transaction. Bypasses the LSM memtable so Origin transactions can - /// stay atomic with the document write. + /// Core per-index write: postings, doc length, term set, and stats of + /// `doc` in `index`, within a transaction. Bypasses the LSM memtable so + /// Origin transactions can stay atomic with the document write. pub(super) fn write_index_data( &self, txn: &WriteTransaction, - scope: IndexDocScope<'_>, + doc: IndexDocScope<'_>, + index: IndexScope<'_>, tokens: &[String], ) -> crate::Result<()> { let IndexDocScope { database_id, tid, - collection, surrogate, - } = scope; + .. + } = doc; let t = tid.as_u64(); - self.note_doc_write(scope); let mut term_postings: HashMap<&str, (u32, Vec)> = HashMap::new(); for (pos, token) in tokens.iter().enumerate() { @@ -170,41 +155,45 @@ impl InvertedIndex { // the same write transaction as every mutation below, so the // check-and-increment is atomic (no TOCTOU). Its presence is the // idempotency key AND the "is this an update?" signal: it means this - // surrogate was already counted into STATS by a prior index (live - // write or an earlier WAL replay pass), so a repeat index of the SAME - // surrogate must not increment `count` again — and it means a previous - // version of the document holds postings that may need retracting. - let prior_len = prior_doc_length(txn, scope)?; + // surrogate was already counted into this index's STATS by a prior + // index (live write or an earlier WAL replay pass), so a repeat index + // of the SAME surrogate must not increment `count` again — and it + // means a previous version of the document holds postings that may + // need retracting. + let prior_len = prior_doc_length(txn, doc, index)?; // Retract the document from every term the previous version put it in - // that the new text does not: `write_index_data` below only touches - // terms present in `tokens`, so without this a word deleted by an - // update keeps its posting and its inflated `df` forever. + // that the new text does not: the loop below only touches terms + // present in `tokens`, so without this a word deleted by an update + // keeps its posting and its inflated `df` forever. let new_terms: BTreeSet<&str> = term_postings.keys().copied().collect(); - let dropped: Vec = doc_terms::occupied_terms(txn, scope, prior_len.is_some())? + let dropped: Vec = doc_terms::occupied_terms(txn, doc, index, prior_len.is_some())? .into_iter() .filter(|term| !new_terms.contains(term.as_str())) .collect(); - doc_terms::strip_postings(txn, scope, &dropped)?; - doc_terms::put(txn, scope, &new_terms)?; + doc_terms::strip_postings(txn, doc, index, &dropped)?; + doc_terms::put(txn, doc, index, &new_terms)?; let mut postings_table = txn .open_table(POSTINGS) .map_err(|e| inverted_err("open postings", e))?; for (term, (freq, positions)) in &term_postings { + let key = posting_key(database_id, t, index, term); let posting = Posting { doc_id: surrogate, term_freq: *freq, positions: positions.clone(), }; - let mut existing: Vec = postings_table - .get((database_id, t, collection, *term)) - .ok() - .flatten() - .and_then(|v| zerompk::from_msgpack(v.value()).ok()) - .unwrap_or_default(); + let mut existing: Vec = match postings_table + .get(key) + .map_err(|e| inverted_err("read postings", e))? + { + Some(v) => zerompk::from_msgpack(v.value()) + .map_err(|e| inverted_err("decode postings", e))?, + None => Vec::new(), + }; existing.retain(|p| p.doc_id != surrogate); existing.push(posting); @@ -212,7 +201,7 @@ impl InvertedIndex { let bytes = zerompk::to_msgpack_vec(&existing) .map_err(|e| inverted_err("serialize postings", e))?; postings_table - .insert((database_id, t, collection, *term), bytes.as_slice()) + .insert(key, bytes.as_slice()) .map_err(|e| inverted_err("insert posting", e))?; } drop(postings_table); @@ -225,7 +214,7 @@ impl InvertedIndex { zerompk::to_msgpack_vec(&doc_len).map_err(|e| inverted_err("serialize doc_len", e))?; lengths .insert( - (database_id, t, collection, surrogate.as_u32()), + doc_key(database_id, t, index, surrogate), len_bytes.as_slice(), ) .map_err(|e| inverted_err("insert doc_len", e))?; @@ -242,76 +231,83 @@ impl InvertedIndex { Some(prior) => (0i64, doc_len as i64 - prior as i64), }; - Self::update_stats_in_txn(txn, database_id, tid, collection, count_delta, total_delta)?; + Self::update_stats_in_txn(txn, database_id, tid, index, count_delta, total_delta)?; - debug!(database_id, tid = t, %collection, surrogate = surrogate.as_u32(), tokens = tokens.len(), terms = term_postings.len(), "indexed document"); + debug!( + database_id, + tid = t, + collection = index.collection(), + field = index.field_key(), + surrogate = surrogate.as_u32(), + tokens = tokens.len(), + terms = term_postings.len(), + "indexed document" + ); Ok(()) } - /// Atomically update `(doc_count, total_token_sum)` in STATS by the given - /// explicit deltas. + /// Atomically update one index's `(doc_count, total_token_sum)` in STATS + /// by the given explicit deltas. /// /// Callers compute `count_delta` / `total_delta` themselves rather than /// this function inferring "new doc vs. removal" from the sign of a /// single combined delta: a re-index of an already-counted surrogate /// (e.g. WAL replay) needs `count_delta == 0` with a `total_delta` that - /// may be positive, negative, or zero — a case the old sign-based - /// inference could not express, which is what caused STATS to be - /// double-counted on replay. + /// may be positive, negative, or zero. pub(super) fn update_stats_in_txn( txn: &WriteTransaction, database_id: u64, tid: TenantId, - collection: &str, + index: IndexScope<'_>, count_delta: i64, total_delta: i64, ) -> crate::Result<()> { - let t = tid.as_u64(); + let key = stats_key(database_id, tid.as_u64(), index); let mut stats = txn .open_table(STATS) .map_err(|e| inverted_err("open stats", e))?; - let (count, total) = stats - .get((database_id, t, collection)) - .ok() - .flatten() - .and_then(|v| zerompk::from_msgpack::<(u32, u64)>(v.value()).ok()) - .unwrap_or((0, 0)); + let (count, total) = match stats.get(key).map_err(|e| inverted_err("read stats", e))? { + Some(v) => zerompk::from_msgpack::<(u32, u64)>(v.value()) + .map_err(|e| inverted_err("decode stats", e))?, + None => (0, 0), + }; - let new_count = (i64::from(count) + count_delta).max(0) as u32; - let new_total = (total as i64 + total_delta).max(0) as u64; + let new_count = (i64::from(count) + count_delta).clamp(0, i64::from(u32::MAX)) as u32; + let new_total = (total as i64).saturating_add(total_delta).max(0) as u64; let bytes = zerompk::to_msgpack_vec(&(new_count, new_total)) .map_err(|e| inverted_err("serialize stats", e))?; stats - .insert((database_id, t, collection), bytes.as_slice()) + .insert(key, bytes.as_slice()) .map_err(|e| inverted_err("insert stats", e))?; Ok(()) } } -/// The token length a previous index recorded for this surrogate, or `None` -/// when the document is not in the index. +/// The token length a previous index recorded for this surrogate in `index`, +/// or `None` when the document is not in that index. /// /// Shared by the index and removal paths: both need it as the authoritative /// "is this document already counted into STATS?" answer, and both must read /// it inside the same write transaction as the mutation it gates. pub(super) fn prior_doc_length( txn: &WriteTransaction, - scope: IndexDocScope<'_>, + doc: IndexDocScope<'_>, + index: IndexScope<'_>, ) -> crate::Result> { let table = txn .open_table(DOC_LENGTHS) .map_err(|e| inverted_err("open doc_lengths", e))?; - let prior = table - .get(( - scope.database_id, - scope.tid.as_u64(), - scope.collection, - scope.surrogate.as_u32(), - )) + let key = doc_key(doc.database_id, doc.tid.as_u64(), index, doc.surrogate); + match table + .get(key) .map_err(|e| inverted_err("read doc_length", e))? - .and_then(|v| zerompk::from_msgpack::(v.value()).ok()); - Ok(prior) + { + Some(v) => zerompk::from_msgpack::(v.value()) + .map(Some) + .map_err(|e| inverted_err("decode doc_length", e)), + None => Ok(None), + } } #[cfg(test)] @@ -324,6 +320,7 @@ mod tests { use nodedb_fts::posting::QueryMode; use super::*; + use crate::engine::sparse::inverted::test_support::body; const DB: u64 = 0; const T: TenantId = TenantId::new(1); @@ -346,14 +343,14 @@ mod tests { #[test] fn update_dropping_a_term_removes_its_posting_and_restores_df() { let (idx, _dir) = open_temp(); - idx.index_document(DB, T, "docs", Surrogate::new(1), "alpha bravo") + idx.index_document(DB, T, "docs", Surrogate::new(1), &body("alpha bravo")) .unwrap(); - idx.index_document(DB, T, "docs", Surrogate::new(2), "bravo charlie") + idx.index_document(DB, T, "docs", Surrogate::new(2), &body("bravo charlie")) .unwrap(); assert_eq!(idx.term_df(DB, T, "docs", "bravo").unwrap(), 2); // Document 1 loses "bravo" and gains "delta". - idx.index_document(DB, T, "docs", Surrogate::new(1), "alpha delta") + idx.index_document(DB, T, "docs", Surrogate::new(1), &body("alpha delta")) .unwrap(); assert_eq!( @@ -389,11 +386,11 @@ mod tests { #[test] fn update_adding_a_term_indexes_it() { let (idx, _dir) = open_temp(); - idx.index_document(DB, T, "docs", Surrogate::new(1), "alpha") + idx.index_document(DB, T, "docs", Surrogate::new(1), &body("alpha")) .unwrap(); assert_eq!(idx.term_df(DB, T, "docs", "delta").unwrap(), 0); - idx.index_document(DB, T, "docs", Surrogate::new(1), "alpha delta") + idx.index_document(DB, T, "docs", Surrogate::new(1), &body("alpha delta")) .unwrap(); assert_eq!(idx.term_df(DB, T, "docs", "alpha").unwrap(), 1); @@ -409,9 +406,15 @@ mod tests { #[test] fn reindex_with_unchanged_tokens_is_a_no_op_for_counts() { let (idx, _dir) = open_temp(); - idx.index_document(DB, T, "docs", Surrogate::new(1), "alpha bravo charlie") - .unwrap(); - idx.index_document(DB, T, "docs", Surrogate::new(2), "bravo") + idx.index_document( + DB, + T, + "docs", + Surrogate::new(1), + &body("alpha bravo charlie"), + ) + .unwrap(); + idx.index_document(DB, T, "docs", Surrogate::new(2), &body("bravo")) .unwrap(); let before = ( @@ -421,8 +424,14 @@ mod tests { idx.term_df(DB, T, "docs", "charlie").unwrap(), ); - idx.index_document(DB, T, "docs", Surrogate::new(1), "alpha bravo charlie") - .unwrap(); + idx.index_document( + DB, + T, + "docs", + Surrogate::new(1), + &body("alpha bravo charlie"), + ) + .unwrap(); let after = ( idx.corpus_stats(DB, T, "docs").unwrap(), @@ -433,29 +442,35 @@ mod tests { assert_eq!(before, after, "an identical re-index must change nothing"); } - /// A document indexed before term sets were recorded has no stored set. Its - /// first re-index must still retract dropped terms, via the fallback scan. + /// A document whose stored term set is missing must still retract dropped + /// terms on its next re-index, via the fallback scan. #[test] fn update_of_a_document_without_a_stored_term_set_still_retracts() { use crate::engine::sparse::fts_redb::tables::DOC_TERMS; let (idx, _dir) = open_temp(); - idx.index_document(DB, T, "docs", Surrogate::new(1), "alpha bravo") + idx.index_document(DB, T, "docs", Surrogate::new(1), &body("alpha bravo")) .unwrap(); - // Drop the term-set row to reproduce a document written by a build that - // did not maintain one. + // Drop the term-set row of the whole-document index. { let db = idx.backend().db(); let txn = db.begin_write().unwrap(); { let mut table = txn.open_table(DOC_TERMS).unwrap(); - table.remove((DB, T.as_u64(), "docs", 1u32)).unwrap(); + table + .remove(doc_key( + DB, + T.as_u64(), + IndexScope::document("docs"), + Surrogate::new(1), + )) + .unwrap(); } txn.commit().unwrap(); } - idx.index_document(DB, T, "docs", Surrogate::new(1), "alpha") + idx.index_document(DB, T, "docs", Surrogate::new(1), &body("alpha")) .unwrap(); assert_eq!( @@ -472,12 +487,12 @@ mod tests { #[test] fn update_to_empty_text_removes_the_document() { let (idx, _dir) = open_temp(); - idx.index_document(DB, T, "docs", Surrogate::new(1), "alpha bravo") + idx.index_document(DB, T, "docs", Surrogate::new(1), &body("alpha bravo")) .unwrap(); - idx.index_document(DB, T, "docs", Surrogate::new(2), "bravo") + idx.index_document(DB, T, "docs", Surrogate::new(2), &body("bravo")) .unwrap(); - idx.index_document(DB, T, "docs", Surrogate::new(1), "") + idx.index_document(DB, T, "docs", Surrogate::new(1), &body("")) .unwrap(); assert_eq!(idx.term_df(DB, T, "docs", "alpha").unwrap(), 0); @@ -489,5 +504,8 @@ mod tests { let (count, avg_len) = idx.corpus_stats(DB, T, "docs").unwrap(); assert_eq!(count, 1, "the emptied document is no longer in the corpus"); assert_eq!(avg_len, 1.0); + let body_index = IndexScope::field("docs", "body").unwrap(); + assert_eq!(idx.term_df(DB, T, body_index, "alpha").unwrap(), 0); + assert_eq!(idx.corpus_stats(DB, T, body_index).unwrap().0, 1); } } diff --git a/nodedb/src/engine/sparse/inverted/mod.rs b/nodedb/src/engine/sparse/inverted/mod.rs index c1d719fc4..9513ad37d 100644 --- a/nodedb/src/engine/sparse/inverted/mod.rs +++ b/nodedb/src/engine/sparse/inverted/mod.rs @@ -10,14 +10,18 @@ //! //! The public API takes `TenantId` as a first-class parameter. Every //! persistent redb table is keyed by the structural tuple -//! `(tenant_id, collection, …)` — per-tenant drops are a tuple range scan, -//! not a lexical-prefix scan. +//! `(database_id, tenant_id, collection, field, …)` — per-tenant drops are a +//! tuple range scan, not a lexical-prefix scan. Each collection has a +//! whole-document index and one index per top-level string field. +mod batch_removal; mod compaction; mod core; mod corpus_stats; +mod doc_fields; mod doc_image; mod doc_terms; +mod document; mod errors; mod indexing; mod rebuild_install; @@ -25,14 +29,18 @@ mod rebuild_journal; mod rebuild_snapshot; mod removal; mod search; +mod staged_search; mod synonyms; +#[cfg(test)] +pub(crate) mod test_support; pub use core::InvertedIndex; pub use doc_image::FtsDocImage; pub use indexing::IndexDocScope; -pub use nodedb_fts::FtsSearchParams; pub use nodedb_fts::posting::{MatchOffset, Posting, QueryMode, TextSearchResult}; +pub use nodedb_fts::{DocumentText, FtsSearchParams, IndexScope}; pub use rebuild_install::{FtsInstallOutcome, FtsRebuildRefusal}; pub use rebuild_journal::FTS_REBUILD_JOURNAL_MAX_DOCS; pub use rebuild_snapshot::{FtsRebuildTicket, FtsRebuilt, FtsSnapshot}; pub use search::PhraseSearchParams; +pub use staged_search::TextDocScorer; diff --git a/nodedb/src/engine/sparse/inverted/rebuild_install.rs b/nodedb/src/engine/sparse/inverted/rebuild_install.rs index 45368b579..acf133cdc 100644 --- a/nodedb/src/engine/sparse/inverted/rebuild_install.rs +++ b/nodedb/src/engine/sparse/inverted/rebuild_install.rs @@ -5,17 +5,19 @@ //! One redb write transaction: //! //! 1. reads the live footprint of every document the journal noted, -//! 2. replaces the collection's `POSTINGS`, `DOC_LENGTHS`, `DOC_TERMS` and -//! `STATS` rows with the rebuilt ones, +//! 2. replaces the `POSTINGS`, `DOC_LENGTHS`, `DOC_TERMS` and `STATS` rows of +//! every index of the collection, and its `DOC_FIELDS` rows, with the +//! rebuilt ones, //! 3. writes each noted document's live footprint over the rebuilt rows, //! or removes the document when it is no longer indexed. //! //! Readers see the transaction's state before or after the commit, never a -//! mix. `INDEX_META` (analyzer, fuzzy flag, synonyms) and `SEGMENTS` are -//! left alone. +//! mix. `INDEX_META` (analyzer, fuzzy flag, synonyms, fieldnorms) and +//! `SEGMENTS` are left alone. -use redb::{ReadableTable as _, WriteTransaction}; +use redb::{TableDefinition, WriteTransaction}; +use nodedb_fts::IndexScope; use nodedb_types::{Surrogate, TenantId}; use super::core::InvertedIndex; @@ -23,11 +25,14 @@ use super::doc_image::FtsDocImage; use super::errors::inverted_err; use super::indexing::IndexDocScope; use super::rebuild_journal::JournalState; -use super::rebuild_snapshot::FtsRebuilt; -use crate::engine::sparse::fts_redb::tables::{DOC_LENGTHS, DOC_TERMS, POSTINGS, STATS}; - -/// Upper bound for the `term` component of a posting range scan. -const MAX_TERM: &str = "\u{10ffff}"; +use super::rebuild_snapshot::{FtsRebuilt, RebuiltScope}; +use crate::engine::sparse::fts_redb::keys::{ + KeyOwner, doc_fields_key, doc_key, posting_key, stats_key, +}; +use crate::engine::sparse::fts_redb::scan::{self, DocKeyDef}; +use crate::engine::sparse::fts_redb::tables::{ + DOC_FIELDS, DOC_LENGTHS, DOC_TERMS, POSTINGS, STATS, +}; /// Why a cutover left the live index as it is. #[derive(Debug, Clone, Copy, PartialEq, Eq, thiserror::Error)] @@ -53,9 +58,9 @@ pub enum FtsRebuildRefusal { pub enum FtsInstallOutcome { /// The rebuilt rows are live. Installed { - /// Terms with a posting list. + /// Term posting lists across every index of the collection. terms: usize, - /// Documents in the rebuilt rows before replay. + /// Documents in the rebuilt whole-document index before replay. docs: usize, /// Documents written during the rebuild and replayed. replayed: usize, @@ -91,7 +96,7 @@ impl InvertedIndex { touched.sort_unstable(); let tid = TenantId::new(rebuilt.tid); - let scope_of = |surrogate: u32| IndexDocScope { + let doc_of = |surrogate: u32| IndexDocScope { database_id: rebuilt.database_id, tid, collection: &rebuilt.collection, @@ -107,7 +112,7 @@ impl InvertedIndex { for &surrogate in &touched { live.push(( surrogate, - Self::read_document_image(&txn, scope_of(surrogate))?, + Self::read_document_image(&txn, doc_of(surrogate))?, )); } @@ -116,83 +121,149 @@ impl InvertedIndex { for (surrogate, image) in &live { match image { - Some(image) => self.write_index_data(&txn, scope_of(*surrogate), image.tokens())?, - None => self.remove_document_in_txn(&txn, scope_of(*surrogate))?, + Some(image) => self.write_document_image(&txn, doc_of(*surrogate), image)?, + None => self.remove_document_in_txn(&txn, doc_of(*surrogate))?, } } txn.commit() .map_err(|e| inverted_err("rebuild cutover commit", e))?; Ok(FtsInstallOutcome::Installed { - terms: rebuilt.postings.len(), - docs: rebuilt.doc_lengths.len(), + terms: rebuilt.scopes.iter().map(|s| s.postings.len()).sum(), + docs: rebuilt + .scopes + .iter() + .find(|s| s.field.is_empty()) + .map_or(0, |s| s.doc_lengths.len()), replayed: live.len(), }) } } -/// Remove the collection's rows from every table the rebuild owns. +/// Remove the rows of every index of the collection from every table the +/// rebuild owns, and the collection's field sets. fn clear_collection_rows( txn: &WriteTransaction, database_id: u64, tid: u64, collection: &str, ) -> crate::Result<()> { + let owner = KeyOwner::collection(database_id, tid, collection); { let mut table = txn .open_table(POSTINGS) .map_err(|e| inverted_err("rebuild open postings", e))?; - let terms: Vec = table - .range((database_id, tid, collection, "")..=(database_id, tid, collection, MAX_TERM)) - .map_err(|e| inverted_err("rebuild postings range", e))? - .map(|entry| entry.map(|(k, _)| k.value().3.to_string())) - .collect::>() - .map_err(|e| inverted_err("rebuild postings entry", e))?; - for term in &terms { + let keys = + scan::str_keys(&table, owner).map_err(|e| inverted_err("rebuild postings range", e))?; + for (c, f, term) in &keys { table - .remove((database_id, tid, collection, term.as_str())) + .remove(posting_key( + database_id, + tid, + IndexScope::from_key(c, f), + term, + )) .map_err(|e| inverted_err("rebuild remove postings", e))?; } } - for (def, name) in [(DOC_LENGTHS, "doc_lengths"), (DOC_TERMS, "doc_terms")] { - let mut table = txn - .open_table(def) - .map_err(|e| inverted_err(&format!("rebuild open {name}"), e))?; - let docs: Vec = table - .range((database_id, tid, collection, 0u32)..=(database_id, tid, collection, u32::MAX)) - .map_err(|e| inverted_err(&format!("rebuild {name} range"), e))? - .map(|entry| entry.map(|(k, _)| k.value().3)) - .collect::>() - .map_err(|e| inverted_err(&format!("rebuild {name} entry"), e))?; - for doc in docs { - table - .remove((database_id, tid, collection, doc)) - .map_err(|e| inverted_err(&format!("rebuild remove {name}"), e))?; + clear_doc_rows(txn, DOC_LENGTHS, "doc_lengths", database_id, tid, owner)?; + clear_doc_rows(txn, DOC_TERMS, "doc_terms", database_id, tid, owner)?; + { + let mut stats = txn + .open_table(STATS) + .map_err(|e| inverted_err("rebuild open stats", e))?; + let keys = + scan::stats_keys(&stats, owner).map_err(|e| inverted_err("rebuild stats range", e))?; + for (c, f) in &keys { + stats + .remove(stats_key(database_id, tid, IndexScope::from_key(c, f))) + .map_err(|e| inverted_err("rebuild remove stats", e))?; } } - let mut stats = txn - .open_table(STATS) - .map_err(|e| inverted_err("rebuild open stats", e))?; - stats - .remove((database_id, tid, collection)) - .map_err(|e| inverted_err("rebuild remove stats", e))?; + let mut fields = txn + .open_table(DOC_FIELDS) + .map_err(|e| inverted_err("rebuild open doc_fields", e))?; + let keys = scan::doc_fields_keys(&fields, owner) + .map_err(|e| inverted_err("rebuild doc_fields range", e))?; + for (c, s) in &keys { + fields + .remove(doc_fields_key(database_id, tid, c, Surrogate::new(*s))) + .map_err(|e| inverted_err("rebuild remove doc_fields", e))?; + } Ok(()) } -/// Write the rebuilt rows of the collection. +/// Remove every row `owner` covers from a DOC_LENGTHS-shaped table. +fn clear_doc_rows( + txn: &WriteTransaction, + def: TableDefinition<'static, DocKeyDef, &'static [u8]>, + name: &str, + database_id: u64, + tid: u64, + owner: KeyOwner<'_>, +) -> crate::Result<()> { + let mut table = txn + .open_table(def) + .map_err(|e| inverted_err(&format!("rebuild open {name}"), e))?; + let keys = scan::doc_keys(&table, owner) + .map_err(|e| inverted_err(&format!("rebuild {name} range"), e))?; + for (c, f, doc) in &keys { + table + .remove(doc_key( + database_id, + tid, + IndexScope::from_key(c, f), + Surrogate::new(*doc), + )) + .map_err(|e| inverted_err(&format!("rebuild remove {name}"), e))?; + } + Ok(()) +} + +/// Write the rebuilt rows of every index of the collection. fn write_rebuilt_rows(txn: &WriteTransaction, rebuilt: &FtsRebuilt) -> crate::Result<()> { + for scope in &rebuilt.scopes { + write_rebuilt_scope(txn, rebuilt, scope)?; + } + let mut table = txn + .open_table(DOC_FIELDS) + .map_err(|e| inverted_err("rebuild open doc_fields", e))?; + for (doc, fields) in &rebuilt.doc_fields { + let bytes = zerompk::to_msgpack_vec(fields) + .map_err(|e| inverted_err("rebuild serialize doc_fields", e))?; + table + .insert( + doc_fields_key( + rebuilt.database_id, + rebuilt.tid, + &rebuilt.collection, + Surrogate::new(*doc), + ), + bytes.as_slice(), + ) + .map_err(|e| inverted_err("rebuild insert doc_fields", e))?; + } + Ok(()) +} + +/// Write the rebuilt rows of one index. +fn write_rebuilt_scope( + txn: &WriteTransaction, + rebuilt: &FtsRebuilt, + scope: &RebuiltScope, +) -> crate::Result<()> { let db = rebuilt.database_id; let t = rebuilt.tid; - let coll = rebuilt.collection.as_str(); + let index = IndexScope::from_key(&rebuilt.collection, &scope.field); { let mut table = txn .open_table(POSTINGS) .map_err(|e| inverted_err("rebuild open postings", e))?; - for (term, list) in &rebuilt.postings { + for (term, list) in &scope.postings { let bytes = zerompk::to_msgpack_vec(list) .map_err(|e| inverted_err("rebuild serialize postings", e))?; table - .insert((db, t, coll, term.as_str()), bytes.as_slice()) + .insert(posting_key(db, t, index, term), bytes.as_slice()) .map_err(|e| inverted_err("rebuild insert postings", e))?; } } @@ -200,11 +271,11 @@ fn write_rebuilt_rows(txn: &WriteTransaction, rebuilt: &FtsRebuilt) -> crate::Re let mut table = txn .open_table(DOC_LENGTHS) .map_err(|e| inverted_err("rebuild open doc_lengths", e))?; - for &(doc, len) in &rebuilt.doc_lengths { + for &(doc, len) in &scope.doc_lengths { let bytes = zerompk::to_msgpack_vec(&len) .map_err(|e| inverted_err("rebuild serialize doc_length", e))?; table - .insert((db, t, coll, doc), bytes.as_slice()) + .insert(doc_key(db, t, index, Surrogate::new(doc)), bytes.as_slice()) .map_err(|e| inverted_err("rebuild insert doc_length", e))?; } } @@ -212,21 +283,24 @@ fn write_rebuilt_rows(txn: &WriteTransaction, rebuilt: &FtsRebuilt) -> crate::Re let mut table = txn .open_table(DOC_TERMS) .map_err(|e| inverted_err("rebuild open doc_terms", e))?; - for (doc, terms) in &rebuilt.doc_terms { + for (doc, terms) in &scope.doc_terms { let bytes = zerompk::to_msgpack_vec(terms) .map_err(|e| inverted_err("rebuild serialize doc_terms", e))?; table - .insert((db, t, coll, *doc), bytes.as_slice()) + .insert( + doc_key(db, t, index, Surrogate::new(*doc)), + bytes.as_slice(), + ) .map_err(|e| inverted_err("rebuild insert doc_terms", e))?; } } let mut stats = txn .open_table(STATS) .map_err(|e| inverted_err("rebuild open stats", e))?; - let bytes = zerompk::to_msgpack_vec(&(rebuilt.doc_count, rebuilt.total_tokens)) + let bytes = zerompk::to_msgpack_vec(&(scope.doc_count, scope.total_tokens)) .map_err(|e| inverted_err("rebuild serialize stats", e))?; stats - .insert((db, t, coll), bytes.as_slice()) + .insert(stats_key(db, t, index), bytes.as_slice()) .map_err(|e| inverted_err("rebuild insert stats", e))?; Ok(()) } @@ -235,13 +309,14 @@ fn write_rebuilt_rows(txn: &WriteTransaction, rebuilt: &FtsRebuilt) -> crate::Re mod tests { use std::sync::Arc; - use redb::{Database, ReadableDatabase as _}; + use redb::{Database, ReadableDatabase as _, ReadableTable as _}; use nodedb_fts::FtsSearchParams; use nodedb_fts::posting::QueryMode; use super::*; use crate::engine::sparse::inverted::FTS_REBUILD_JOURNAL_MAX_DOCS; + use crate::engine::sparse::inverted::test_support::{body, fields}; const DB: u64 = 0; const T: TenantId = TenantId::new(1); @@ -257,12 +332,12 @@ mod tests { (idx, dir) } - fn hits(idx: &InvertedIndex, query: &str) -> Vec { + fn hits_in(idx: &InvertedIndex, index: IndexScope<'_>, query: &str) -> Vec { let mut ids: Vec = idx .search( DB, T, - "docs", + index, FtsSearchParams { query, top_k: 100, @@ -279,6 +354,10 @@ mod tests { ids } + fn hits(idx: &InvertedIndex, query: &str) -> Vec { + hits_in(idx, IndexScope::document("docs"), query) + } + fn rebuild_with(idx: &InvertedIndex, during: impl FnOnce(&InvertedIndex)) -> FtsInstallOutcome { let ticket = idx .begin_rebuild(DB, T, "docs", FTS_REBUILD_JOURNAL_MAX_DOCS) @@ -292,13 +371,13 @@ mod tests { fn writes_during_the_rebuild_survive_the_cutover() { let (idx, _dir) = open_temp(); for i in 1..=3u32 { - idx.index_document(DB, T, "docs", Surrogate::new(i), "alpha") + idx.index_document(DB, T, "docs", Surrogate::new(i), &body("alpha")) .unwrap(); } let outcome = rebuild_with(&idx, |idx| { - idx.index_document(DB, T, "docs", Surrogate::new(4), "alpha") + idx.index_document(DB, T, "docs", Surrogate::new(4), &body("alpha")) .unwrap(); - idx.index_document(DB, T, "docs", Surrogate::new(1), "beta") + idx.index_document(DB, T, "docs", Surrogate::new(1), &body("beta")) .unwrap(); idx.remove_document(DB, T, "docs", Surrogate::new(2)) .unwrap(); @@ -317,13 +396,13 @@ mod tests { #[test] fn a_term_dropped_during_the_rebuild_leaves_its_posting_list() { let (idx, _dir) = open_temp(); - idx.index_document(DB, T, "docs", Surrogate::new(1), "alpha bravo") + idx.index_document(DB, T, "docs", Surrogate::new(1), &body("alpha bravo")) .unwrap(); - idx.index_document(DB, T, "docs", Surrogate::new(2), "alpha") + idx.index_document(DB, T, "docs", Surrogate::new(2), &body("alpha")) .unwrap(); let outcome = rebuild_with(&idx, |idx| { // The snapshot holds document 1 under `alpha`; the update drops it. - idx.index_document(DB, T, "docs", Surrogate::new(1), "bravo charlie") + idx.index_document(DB, T, "docs", Surrogate::new(1), &body("bravo charlie")) .unwrap(); }); assert!(matches!( @@ -337,7 +416,12 @@ mod tests { .unwrap() .open_table(POSTINGS) .unwrap() - .get((DB, T.as_u64(), "docs", "alpha")) + .get(posting_key( + DB, + T.as_u64(), + IndexScope::document("docs"), + "alpha", + )) .unwrap() .map(|v| zerompk::from_msgpack::>(v.value()).unwrap()) .unwrap_or_default(); @@ -354,10 +438,59 @@ mod tests { assert_eq!(avg_len, 1.5, "(2 + 1) tokens over 2 documents"); } + #[test] + fn a_rebuild_preserves_field_scopes() { + let (idx, _dir) = open_temp(); + idx.index_document( + DB, + T, + "docs", + Surrogate::new(1), + &fields(&[("title", "rust"), ("body", "guide")]), + ) + .unwrap(); + idx.index_document(DB, T, "docs", Surrogate::new(2), &body("rust")) + .unwrap(); + let title = IndexScope::field("docs", "title").unwrap(); + let body_index = IndexScope::field("docs", "body").unwrap(); + let outcome = rebuild_with(&idx, |idx| { + // Moves `rust` out of document 1's title during the rebuild. + idx.index_document( + DB, + T, + "docs", + Surrogate::new(1), + &fields(&[("body", "guide rust")]), + ) + .unwrap(); + idx.index_document( + DB, + T, + "docs", + Surrogate::new(3), + &fields(&[("title", "rust")]), + ) + .unwrap(); + }); + assert!(matches!(outcome, FtsInstallOutcome::Installed { .. })); + assert_eq!(hits_in(&idx, title, "rust"), vec![3]); + assert_eq!(hits_in(&idx, body_index, "rust"), vec![1, 2]); + assert_eq!(hits(&idx, "rust"), vec![1, 2, 3]); + assert_eq!(idx.corpus_stats(DB, T, title).unwrap().0, 1); + assert_eq!(idx.corpus_stats(DB, T, body_index).unwrap().0, 2); + assert_eq!(idx.corpus_stats(DB, T, "docs").unwrap().0, 3); + + // The rebuilt field set of document 1 lets a delete clear every scope. + idx.remove_document(DB, T, "docs", Surrogate::new(1)) + .unwrap(); + assert_eq!(hits_in(&idx, body_index, "guide"), Vec::::new()); + assert_eq!(idx.corpus_stats(DB, T, body_index).unwrap().0, 1); + } + #[test] fn a_purge_during_the_rebuild_refuses_the_cutover() { let (idx, _dir) = open_temp(); - idx.index_document(DB, T, "docs", Surrogate::new(1), "alpha") + idx.index_document(DB, T, "docs", Surrogate::new(1), &body("alpha")) .unwrap(); let outcome = rebuild_with(&idx, |idx| { idx.purge_collection(DB, T, "docs").unwrap(); @@ -373,9 +506,9 @@ mod tests { fn journal_overflow_refuses_the_cutover_and_keeps_live_writes() { let (idx, _dir) = open_temp(); let ticket = idx.begin_rebuild(DB, T, "docs", 1).unwrap(); - idx.index_document(DB, T, "docs", Surrogate::new(1), "alpha") + idx.index_document(DB, T, "docs", Surrogate::new(1), &body("alpha")) .unwrap(); - idx.index_document(DB, T, "docs", Surrogate::new(2), "alpha") + idx.index_document(DB, T, "docs", Surrogate::new(2), &body("alpha")) .unwrap(); let outcome = idx .install_rebuild(ticket.read().unwrap().compact()) diff --git a/nodedb/src/engine/sparse/inverted/rebuild_journal.rs b/nodedb/src/engine/sparse/inverted/rebuild_journal.rs index d0024a104..c39240eb8 100644 --- a/nodedb/src/engine/sparse/inverted/rebuild_journal.rs +++ b/nodedb/src/engine/sparse/inverted/rebuild_journal.rs @@ -18,7 +18,7 @@ use std::collections::HashSet; -use redb::{ReadableDatabase as _, ReadableTable as _}; +use redb::ReadableDatabase as _; use nodedb_types::TenantId; @@ -26,11 +26,10 @@ use super::core::InvertedIndex; use super::errors::inverted_err; use super::indexing::IndexDocScope; use super::rebuild_snapshot::FtsRebuildTicket; +use crate::engine::sparse::fts_redb::keys::KeyOwner; +use crate::engine::sparse::fts_redb::scan; use crate::engine::sparse::fts_redb::tables::{DOC_LENGTHS, POSTINGS}; -/// Upper bound for the `term` component of a posting range scan. -const MAX_TERM: &str = "\u{10ffff}"; - /// Default bound on the distinct documents one rebuild journal records. pub const FTS_REBUILD_JOURNAL_MAX_DOCS: usize = 1 << 20; @@ -148,26 +147,20 @@ impl InvertedIndex { .db() .begin_read() .map_err(|e| inverted_err("rebuild probe txn", e))?; + let owner = KeyOwner::collection(database_id, t, collection); let postings = txn .open_table(POSTINGS) .map_err(|e| inverted_err("rebuild probe postings", e))?; - if postings - .range((database_id, t, collection, "")..=(database_id, t, collection, MAX_TERM)) + if scan::has_str_rows(&postings, owner) .map_err(|e| inverted_err("rebuild probe postings range", e))? - .next() - .is_some() { return Ok(true); } let lengths = txn .open_table(DOC_LENGTHS) .map_err(|e| inverted_err("rebuild probe doc_lengths", e))?; - let any = lengths - .range((database_id, t, collection, 0u32)..=(database_id, t, collection, u32::MAX)) - .map_err(|e| inverted_err("rebuild probe doc_lengths range", e))? - .next() - .is_some(); - Ok(any) + scan::has_doc_rows(&lengths, owner) + .map_err(|e| inverted_err("rebuild probe doc_lengths range", e)) } /// Close the journal of rebuild `token`. The index is unchanged. diff --git a/nodedb/src/engine/sparse/inverted/rebuild_snapshot.rs b/nodedb/src/engine/sparse/inverted/rebuild_snapshot.rs index 226389819..659d8308e 100644 --- a/nodedb/src/engine/sparse/inverted/rebuild_snapshot.rs +++ b/nodedb/src/engine/sparse/inverted/rebuild_snapshot.rs @@ -1,7 +1,7 @@ // SPDX-License-Identifier: BUSL-1.1 //! The off-core half of a collection rebuild: read the pinned snapshot and -//! derive the collection's canonical index rows from it. +//! derive the canonical rows of every index of the collection from it. //! //! Every type here is `Send` and touches no `!Send` engine state. The redb //! read transaction was opened on the owning core by @@ -10,16 +10,14 @@ use std::collections::{BTreeMap, BTreeSet}; -use redb::{ReadTransaction, ReadableTable as _}; +use redb::ReadTransaction; use nodedb_fts::posting::Posting; use super::errors::inverted_err; +use crate::engine::sparse::fts_redb::scan; use crate::engine::sparse::fts_redb::tables::{DOC_LENGTHS, POSTINGS}; -/// Upper bound for the `term` component of a posting range scan. -const MAX_TERM: &str = "\u{10ffff}"; - /// A pinned snapshot of one collection's index, ready to read. pub struct FtsRebuildTicket { token: u64, @@ -29,22 +27,26 @@ pub struct FtsRebuildTicket { txn: ReadTransaction, } -/// One collection's postings and document lengths as the snapshot holds them. +/// One index's postings and document lengths as the snapshot holds them. +struct SnapshotScope { + postings: Vec<(String, Vec)>, + doc_lengths: Vec<(u32, u32)>, +} + +/// Every index of one collection as the snapshot holds it, keyed by field +/// key (empty for the whole-document index). pub struct FtsSnapshot { token: u64, database_id: u64, tid: u64, collection: String, - postings: Vec<(String, Vec)>, - doc_lengths: Vec<(u32, u32)>, + scopes: BTreeMap, } -/// The canonical index rows of one collection, derived from a snapshot. -pub struct FtsRebuilt { - pub(super) token: u64, - pub(super) database_id: u64, - pub(super) tid: u64, - pub(super) collection: String, +/// The canonical rows of one index, derived from a snapshot. +pub(super) struct RebuiltScope { + /// Field key: empty for the whole-document index. + pub(super) field: String, /// One list per term, one posting per document, ordered by surrogate. pub(super) postings: Vec<(String, Vec)>, /// `(surrogate, token count)` per indexed document. @@ -57,6 +59,18 @@ pub struct FtsRebuilt { pub(super) total_tokens: u64, } +/// The canonical rows of every index of one collection. +pub struct FtsRebuilt { + pub(super) token: u64, + pub(super) database_id: u64, + pub(super) tid: u64, + pub(super) collection: String, + /// One entry per index, whole-document index first. + pub(super) scopes: Vec, + /// `(surrogate, field names)` per document holding any field index. + pub(super) doc_fields: Vec<(u32, Vec)>, +} + impl FtsRebuildTicket { pub(super) fn new( token: u64, @@ -79,41 +93,39 @@ impl FtsRebuildTicket { self.token } - /// Read the collection's postings and document lengths from the pinned - /// snapshot. Runs on any thread. The snapshot is released on return. + /// Read every index's postings and document lengths of the collection + /// from the pinned snapshot. Runs on any thread. The snapshot is released + /// on return. pub fn read(self) -> crate::Result { - let db = self.database_id; - let t = self.tid; - let coll = self.collection.as_str(); + let (db, t) = (self.database_id, self.tid); + let mut scopes: BTreeMap = BTreeMap::new(); let postings_table = self .txn .open_table(POSTINGS) .map_err(|e| inverted_err("rebuild open postings", e))?; - let mut postings = Vec::new(); - for entry in postings_table - .range((db, t, coll, "")..=(db, t, coll, MAX_TERM)) + for row in scan::str_rows(&postings_table, db, t, &self.collection) .map_err(|e| inverted_err("rebuild postings range", e))? { - let (key, value) = entry.map_err(|e| inverted_err("rebuild postings entry", e))?; - let list: Vec = zerompk::from_msgpack(value.value()) + let list: Vec = zerompk::from_msgpack(&row.value) .map_err(|e| inverted_err("rebuild decode postings", e))?; - postings.push((key.value().3.to_string(), list)); + scope_entry(&mut scopes, row.field) + .postings + .push((row.suffix, list)); } let lengths_table = self .txn .open_table(DOC_LENGTHS) .map_err(|e| inverted_err("rebuild open doc_lengths", e))?; - let mut doc_lengths = Vec::new(); - for entry in lengths_table - .range((db, t, coll, 0u32)..=(db, t, coll, u32::MAX)) + for row in scan::doc_rows(&lengths_table, db, t, &self.collection) .map_err(|e| inverted_err("rebuild doc_lengths range", e))? { - let (key, value) = entry.map_err(|e| inverted_err("rebuild doc_length entry", e))?; - let len: u32 = zerompk::from_msgpack(value.value()) + let len: u32 = zerompk::from_msgpack(&row.value) .map_err(|e| inverted_err("rebuild decode doc_length", e))?; - doc_lengths.push((key.value().3, len)); + scope_entry(&mut scopes, row.field) + .doc_lengths + .push((row.surrogate, len)); } Ok(FtsSnapshot { @@ -121,67 +133,93 @@ impl FtsRebuildTicket { database_id: self.database_id, tid: self.tid, collection: self.collection, - postings, - doc_lengths, + scopes, }) } } +fn scope_entry(scopes: &mut BTreeMap, field: String) -> &mut SnapshotScope { + scopes.entry(field).or_insert_with(|| SnapshotScope { + postings: Vec::new(), + doc_lengths: Vec::new(), + }) +} + impl FtsSnapshot { /// The token that ties this rebuild to its journal. pub fn token(&self) -> u64 { self.token } - /// Derive the canonical rows: one posting per document per term, the - /// term set of each document, and corpus stats counted from the - /// document lengths. Runs on any thread. + /// Derive the canonical rows of every index: one posting per document per + /// term, the term set of each document, corpus stats counted from the + /// document lengths, and each document's field set. Runs on any thread. pub fn compact(self) -> FtsRebuilt { - let mut postings = Vec::with_capacity(self.postings.len()); - let mut terms_by_doc: BTreeMap> = BTreeMap::new(); - for (term, mut list) in self.postings { - list.sort_unstable_by(|a, b| { - a.doc_id - .as_u32() - .cmp(&b.doc_id.as_u32()) - .then(b.term_freq.cmp(&a.term_freq)) - }); - list.dedup_by_key(|p| p.doc_id.as_u32()); - if list.is_empty() { - continue; + let mut fields_by_doc: BTreeMap> = BTreeMap::new(); + let mut scopes = Vec::with_capacity(self.scopes.len()); + // The BTreeMap yields the empty (whole-document) key first. + for (field, scope) in self.scopes { + if !field.is_empty() { + for &(doc, _) in &scope.doc_lengths { + fields_by_doc.entry(doc).or_default().push(field.clone()); + } } - for posting in &list { - terms_by_doc - .entry(posting.doc_id.as_u32()) - .or_default() - .insert(term.clone()); - } - postings.push((term, list)); + scopes.push(compact_scope(field, scope)); } - let doc_terms = terms_by_doc - .into_iter() - .map(|(doc, terms)| (doc, terms.into_iter().collect())) - .collect(); - let doc_count = u32::try_from(self.doc_lengths.len()).unwrap_or(u32::MAX); - let total_tokens = self - .doc_lengths - .iter() - .map(|&(_, len)| u64::from(len)) - .sum(); FtsRebuilt { token: self.token, database_id: self.database_id, tid: self.tid, collection: self.collection, - postings, - doc_lengths: self.doc_lengths, - doc_terms, - doc_count, - total_tokens, + scopes, + doc_fields: fields_by_doc.into_iter().collect(), } } } +/// Canonical rows of one index. +fn compact_scope(field: String, scope: SnapshotScope) -> RebuiltScope { + let mut postings = Vec::with_capacity(scope.postings.len()); + let mut terms_by_doc: BTreeMap> = BTreeMap::new(); + for (term, mut list) in scope.postings { + list.sort_unstable_by(|a, b| { + a.doc_id + .as_u32() + .cmp(&b.doc_id.as_u32()) + .then(b.term_freq.cmp(&a.term_freq)) + }); + list.dedup_by_key(|p| p.doc_id.as_u32()); + if list.is_empty() { + continue; + } + for posting in &list { + terms_by_doc + .entry(posting.doc_id.as_u32()) + .or_default() + .insert(term.clone()); + } + postings.push((term, list)); + } + let doc_terms = terms_by_doc + .into_iter() + .map(|(doc, terms)| (doc, terms.into_iter().collect())) + .collect(); + let doc_count = u32::try_from(scope.doc_lengths.len()).unwrap_or(u32::MAX); + let total_tokens = scope + .doc_lengths + .iter() + .map(|&(_, len)| u64::from(len)) + .sum(); + RebuiltScope { + field, + postings, + doc_lengths: scope.doc_lengths, + doc_terms, + doc_count, + total_tokens, + } +} + impl FtsRebuilt { /// The token that ties this rebuild to its journal. pub fn token(&self) -> u64 { diff --git a/nodedb/src/engine/sparse/inverted/removal.rs b/nodedb/src/engine/sparse/inverted/removal.rs index 3028c3386..12e374548 100644 --- a/nodedb/src/engine/sparse/inverted/removal.rs +++ b/nodedb/src/engine/sparse/inverted/removal.rs @@ -11,16 +11,19 @@ use redb::WriteTransaction; +use nodedb_fts::IndexScope; use nodedb_types::{Surrogate, TenantId}; use super::core::InvertedIndex; +use super::doc_fields; use super::doc_terms; use super::errors::inverted_err; use super::indexing::{IndexDocScope, prior_doc_length}; +use crate::engine::sparse::fts_redb::keys::doc_key; use crate::engine::sparse::fts_redb::tables::DOC_LENGTHS; impl InvertedIndex { - /// Remove a document from the inverted index. + /// Remove a document from every index of its collection. pub fn remove_document( &self, database_id: u64, @@ -45,57 +48,62 @@ impl InvertedIndex { Ok(()) } - /// Remove a document within an externally-owned write transaction. + /// Remove a document from every index of its collection within an + /// externally-owned write transaction: each field index its field set + /// names, the whole-document index, and the field-set row. + pub fn remove_document_in_txn( + &self, + txn: &WriteTransaction, + doc: IndexDocScope<'_>, + ) -> crate::Result<()> { + self.note_doc_write(doc); + for field in doc_fields::read(txn, doc)? { + let index = IndexScope::field(doc.collection, &field) + .ok_or_else(|| inverted_err("doc_fields", "stored field name is empty"))?; + self.remove_scope_in_txn(txn, doc, index)?; + } + self.remove_scope_in_txn(txn, doc, IndexScope::document(doc.collection))?; + doc_fields::clear(txn, doc) + } + + /// Remove a document from one index. /// /// A document that is not in the index (no `DOC_LENGTHS` row) is left /// entirely alone — that is what makes a repeated delete a no-op rather /// than a second decrement of the corpus counters, and it is also what /// keeps the term-set fallback scan off the path of documents that were /// never indexed. - pub fn remove_document_in_txn( + pub(super) fn remove_scope_in_txn( &self, txn: &WriteTransaction, - scope: IndexDocScope<'_>, + doc: IndexDocScope<'_>, + index: IndexScope<'_>, ) -> crate::Result<()> { - self.note_doc_write(scope); - let Some(old_len) = prior_doc_length(txn, scope)? else { + let Some(old_len) = prior_doc_length(txn, doc, index)? else { return Ok(()); }; // The stored term set names exactly the lists this document occupies; - // the fallback scan covers documents indexed before term sets were - // recorded. - let terms = doc_terms::occupied_terms(txn, scope, true)?; - doc_terms::strip_postings(txn, scope, &terms)?; - doc_terms::clear(txn, scope)?; + // the fallback scan covers a document whose term set row is missing. + let terms = doc_terms::occupied_terms(txn, doc, index, true)?; + doc_terms::strip_postings(txn, doc, index, &terms)?; + doc_terms::clear(txn, doc, index)?; { let mut lengths = txn .open_table(DOC_LENGTHS) .map_err(|e| inverted_err("open doc_lengths", e))?; lengths - .remove(( - scope.database_id, - scope.tid.as_u64(), - scope.collection, - scope.surrogate.as_u32(), + .remove(doc_key( + doc.database_id, + doc.tid.as_u64(), + index, + doc.surrogate, )) .map_err(|e| inverted_err("remove doc length", e))?; } - Self::update_stats_in_txn( - txn, - scope.database_id, - scope.tid, - scope.collection, - -1, - -(old_len as i64), - )?; - - // Note: the docmap sub-key in INDEX_META (previously maintained by the - // old DocIdMap abstraction) is no longer updated. Searches filter via - // Surrogate prefilter bitmaps instead. - Ok(()) + Self::update_stats_in_txn(txn, doc.database_id, doc.tid, index, -1, -(old_len as i64)) } } @@ -109,6 +117,7 @@ mod tests { use nodedb_fts::posting::QueryMode; use super::*; + use crate::engine::sparse::inverted::test_support::body; const DB: u64 = 0; const T: TenantId = TenantId::new(1); @@ -127,9 +136,9 @@ mod tests { #[test] fn remove_document() { let (idx, _dir) = open_temp(); - idx.index_document(DB, T, "docs", Surrogate::new(1), "hello world") + idx.index_document(DB, T, "docs", Surrogate::new(1), &body("hello world")) .unwrap(); - idx.index_document(DB, T, "docs", Surrogate::new(2), "hello rust") + idx.index_document(DB, T, "docs", Surrogate::new(2), &body("hello rust")) .unwrap(); idx.remove_document(DB, T, "docs", Surrogate::new(1)) @@ -159,9 +168,15 @@ mod tests { fn remove_document_decrements_stats() { let (idx, _dir) = open_temp(); - idx.index_document(DB, T, "docs", Surrogate::new(1), "alpha bravo charlie") - .unwrap(); - idx.index_document(DB, T, "docs", Surrogate::new(2), "delta echo") + idx.index_document( + DB, + T, + "docs", + Surrogate::new(1), + &body("alpha bravo charlie"), + ) + .unwrap(); + idx.index_document(DB, T, "docs", Surrogate::new(2), &body("delta echo")) .unwrap(); let (count, avg_len) = idx.corpus_stats(DB, T, "docs").unwrap(); assert_eq!(count, 2); diff --git a/nodedb/src/engine/sparse/inverted/search.rs b/nodedb/src/engine/sparse/inverted/search.rs index e86864064..013decaf0 100644 --- a/nodedb/src/engine/sparse/inverted/search.rs +++ b/nodedb/src/engine/sparse/inverted/search.rs @@ -6,25 +6,30 @@ use redb::{ReadableDatabase, ReadableTable}; use tracing::debug; -use nodedb_fts::FtsSearchParams; use nodedb_fts::posting::{MatchOffset, Posting, TextSearchResult}; +use nodedb_fts::{FtsSearchParams, IndexScope}; use nodedb_types::{Surrogate, TenantId}; use super::core::InvertedIndex; use super::errors::{fts_index_err, inverted_err}; +use crate::engine::sparse::fts_redb::keys::posting_key; use crate::engine::sparse::fts_redb::tables::POSTINGS; /// Query and tuning parameters for an inverted-index phrase search. /// -/// The `(database_id, tid, collection)` scope is passed separately so callers +/// The `(database_id, tid, index)` scope is passed separately so callers /// can reuse their existing scope variables. pub struct PhraseSearchParams<'a> { - /// Ordered terms that must appear as a contiguous sequence. + /// The phrase as analyzed tokens, in order. They must appear as a + /// contiguous sequence. The caller analyzes the phrase once with the + /// collection's analyzer, the same analysis its indexed text took. pub terms: &'a [String], /// Maximum number of results to return. pub top_k: usize, /// Optional surrogate bitmap restricting candidates before position match. pub prefilter: Option<&'a nodedb_types::SurrogateBitmap>, + /// Documents that never match, excluded before the top-k cut. + pub exclude: Option<&'a nodedb_types::SurrogateBitmap>, } impl InvertedIndex { @@ -36,17 +41,20 @@ impl InvertedIndex { /// /// The result is scored by position rank (earlier = higher). An optional /// `prefilter` bitmap restricts the candidate set before position matching. - pub fn phrase_search( + pub fn phrase_search<'a>( &self, database_id: u64, tid: TenantId, - collection: &str, + index: impl Into>, params: PhraseSearchParams<'_>, ) -> crate::Result> { + let index = index.into(); + let collection = index.collection(); let PhraseSearchParams { terms, top_k, prefilter, + exclude, } = params; if terms.is_empty() { return Ok(Vec::new()); @@ -59,16 +67,17 @@ impl InvertedIndex { .open_table(POSTINGS) .map_err(|e| inverted_err("open postings", e))?; - // Load posting list for each term. + // Load posting list for each analyzed term. let mut term_lists: Vec> = Vec::with_capacity(terms.len()); for term in terms { - let analyzed = self.analyze_for_collection(database_id, tid, collection, term)?; - let canonical = analyzed.into_iter().next().unwrap_or_else(|| term.clone()); - let postings: Vec = postings_table - .get((database_id, t, collection, canonical.as_str())) + let postings: Vec = match postings_table + .get(posting_key(database_id, t, index, term.as_str())) .map_err(|e| inverted_err("read posting", e))? - .and_then(|v| zerompk::from_msgpack(v.value()).ok()) - .unwrap_or_default(); + { + Some(v) => zerompk::from_msgpack(v.value()) + .map_err(|e| inverted_err("decode posting", e))?, + None => Vec::new(), + }; term_lists.push(postings); } @@ -79,7 +88,9 @@ impl InvertedIndex { 'outer: for posting in first { // Prefilter check. - if prefilter.is_some_and(|bm| !bm.0.contains(posting.doc_id.as_u32())) { + if prefilter.is_some_and(|bm| !bm.contains(posting.doc_id)) + || exclude.is_some_and(|bm| bm.contains(posting.doc_id)) + { continue; } @@ -121,6 +132,7 @@ impl InvertedIndex { debug!( tid = t, %collection, + field = index.field_key(), terms = terms.len(), hits = results.len(), "phrase search" @@ -132,15 +144,15 @@ impl InvertedIndex { /// /// Supports `NOT ` and `-` negation in the query string. /// Returns `Err` for invalid queries (NOT-only, unsupported parentheses). - pub fn search( + pub fn search<'a>( &self, database_id: u64, tid: TenantId, - collection: &str, + index: impl Into>, params: FtsSearchParams<'_>, ) -> crate::Result> { self.inner - .search(database_id, tid.as_u64(), collection, params) + .search(database_id, tid.as_u64(), index, params) .map_err(fts_index_err) } @@ -164,6 +176,7 @@ mod tests { use nodedb_fts::posting::QueryMode; use super::*; + use crate::engine::sparse::inverted::test_support::body; const DB: u64 = 0; const T: TenantId = TenantId::new(1); @@ -187,7 +200,7 @@ mod tests { T, "docs", Surrogate::new(1), - "The quick brown fox jumps over the lazy dog", + &body("The quick brown fox jumps over the lazy dog"), ) .unwrap(); idx.index_document( @@ -195,7 +208,7 @@ mod tests { T, "docs", Surrogate::new(2), - "A fast brown dog runs across the field", + &body("A fast brown dog runs across the field"), ) .unwrap(); idx.index_document( @@ -203,7 +216,7 @@ mod tests { T, "docs", Surrogate::new(3), - "Rust programming language for systems", + &body("Rust programming language for systems"), ) .unwrap(); @@ -233,11 +246,17 @@ mod tests { T, "docs", Surrogate::new(1), - "running distributed databases", + &body("running distributed databases"), + ) + .unwrap(); + idx.index_document( + DB, + T, + "docs", + Surrogate::new(2), + &body("the cat sat on a mat"), ) .unwrap(); - idx.index_document(DB, T, "docs", Surrogate::new(2), "the cat sat on a mat") - .unwrap(); let results = idx .search( @@ -265,7 +284,7 @@ mod tests { T, "docs", Surrogate::new(1), - "distributed database systems", + &body("distributed database systems"), ) .unwrap(); @@ -290,7 +309,7 @@ mod tests { #[test] fn empty_query() { let (idx, _dir) = open_temp(); - idx.index_document(DB, T, "docs", Surrogate::new(1), "some text here") + idx.index_document(DB, T, "docs", Surrogate::new(1), &body("some text here")) .unwrap(); let results = idx @@ -313,10 +332,22 @@ mod tests { #[test] fn collections_isolated() { let (idx, _dir) = open_temp(); - idx.index_document(DB, T, "col_a", Surrogate::new(1), "alpha bravo charlie") - .unwrap(); - idx.index_document(DB, T, "col_b", Surrogate::new(1), "delta echo foxtrot") - .unwrap(); + idx.index_document( + DB, + T, + "col_a", + Surrogate::new(1), + &body("alpha bravo charlie"), + ) + .unwrap(); + idx.index_document( + DB, + T, + "col_b", + Surrogate::new(1), + &body("delta echo foxtrot"), + ) + .unwrap(); let results = idx .search( diff --git a/nodedb/src/engine/sparse/inverted/test_support.rs b/nodedb/src/engine/sparse/inverted/test_support.rs new file mode 100644 index 000000000..0de53b493 --- /dev/null +++ b/nodedb/src/engine/sparse/inverted/test_support.rs @@ -0,0 +1,19 @@ +// SPDX-License-Identifier: BUSL-1.1 + +//! Fixtures for the inverted index's inline tests. + +use nodedb_fts::DocumentText; + +/// A document whose only string field is `body`. +pub(crate) fn body(text: &str) -> DocumentText { + fields(&[("body", text)]) +} + +/// A document with the given `(field, text)` pairs. +pub(crate) fn fields(pairs: &[(&str, &str)]) -> DocumentText { + DocumentText::from_fields( + pairs + .iter() + .map(|(f, t)| ((*f).to_string(), (*t).to_string())), + ) +} diff --git a/nodedb/src/event/wal_replay.rs b/nodedb/src/event/wal_replay.rs index fa4ea0909..388d1369f 100644 --- a/nodedb/src/event/wal_replay.rs +++ b/nodedb/src/event/wal_replay.rs @@ -1047,9 +1047,14 @@ mod tests { stream_id: 3, seq: 4, }; - let idx = FtsIndexPayload::new(prov.clone(), "articles", "doc-1", "hello world") - .to_bytes() - .expect("enc idx"); + let idx = FtsIndexPayload::new( + prov.clone(), + "articles", + "doc-1", + vec![("body".to_string(), "hello world".to_string())], + ) + .to_bytes() + .expect("enc idx"); let idx_record = make_record(RecordType::FtsIndex, &idx, 1, 0, 703); let mut seq = 0u64; assert!(record_to_events(&idx_record, &mut seq).is_empty()); diff --git a/nodedb/src/wal/manager/append.rs b/nodedb/src/wal/manager/append.rs index d64f8dde0..7a02a8dbf 100644 --- a/nodedb/src/wal/manager/append.rs +++ b/nodedb/src/wal/manager/append.rs @@ -185,7 +185,7 @@ mod tests { }, "articles", "doc-abc", - "hello world nodedb fts", + vec![("body".to_string(), "hello world nodedb fts".to_string())], ); let bytes = payload.to_bytes().unwrap(); @@ -214,6 +214,9 @@ mod tests { assert_eq!(decoded.provenance.seq, 42); assert_eq!(decoded.collection, "articles"); assert_eq!(decoded.doc_id, "doc-abc"); - assert_eq!(decoded.text, "hello world nodedb fts"); + assert_eq!( + decoded.fields, + vec![("body".to_string(), "hello world nodedb fts".to_string())] + ); } } diff --git a/nodedb/tests/inproc/cases/collection_purge_persistent.rs b/nodedb/tests/inproc/cases/collection_purge_persistent.rs index 31cdc4473..3d52b76d7 100644 --- a/nodedb/tests/inproc/cases/collection_purge_persistent.rs +++ b/nodedb/tests/inproc/cases/collection_purge_persistent.rs @@ -179,11 +179,14 @@ fn inverted_index_purge_is_scoped_to_collection() { .unwrap(); let tid = TenantId::new(TENANT); + let body = |text: &str| { + nodedb_fts::DocumentText::from_fields([("body".to_string(), text.to_string())]) + }; inverted - .index_document(DB, tid, "keep", Surrogate(1), "hello world") + .index_document(DB, tid, "keep", Surrogate(1), &body("hello world")) .unwrap(); inverted - .index_document(DB, tid, "purge_me", Surrogate(2), "hello universe") + .index_document(DB, tid, "purge_me", Surrogate(2), &body("hello universe")) .unwrap(); inverted.purge_collection(DB, tid, "purge_me").unwrap(); diff --git a/nodedb/tests/inproc/cases/fts_update_reindex.rs b/nodedb/tests/inproc/cases/fts_update_reindex.rs index 1e624a926..165bf5784 100644 --- a/nodedb/tests/inproc/cases/fts_update_reindex.rs +++ b/nodedb/tests/inproc/cases/fts_update_reindex.rs @@ -62,6 +62,45 @@ async fn point_update_moves_the_row_to_its_new_terms() { assert_eq!(text_ids(&server, "fu_point", "beta").await, ["u1"]); } +/// An update that moves a word from one field to another retracts it from +/// the old field's index and adds it to the new one, and the per-field +/// indexes hold that state across a restart. +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn update_across_fields_keeps_field_scopes_after_restart() { + let server = TestServer::start().await; + server + .exec("CREATE COLLECTION fu_fields WITH (engine='document_schemaless')") + .await + .unwrap(); + server + .exec("INSERT INTO fu_fields { id: 'f1', title: 'delta heading', body: 'plain body' }") + .await + .unwrap(); + server + .exec("UPDATE fu_fields SET title = 'plain heading', body = 'delta body' WHERE id = 'f1'") + .await + .unwrap(); + + let (server, dir) = server.take_dir(); + server.graceful_shutdown().await; + let (server, _dir) = TestServer::open_on_path(dir).await; + + assert!( + text_ids_on(&server, "fu_fields", "title", "delta") + .await + .is_empty(), + "the title index must no longer hold the moved word" + ); + assert_eq!( + text_ids_on(&server, "fu_fields", "body", "delta").await, + ["f1"] + ); + assert_eq!( + text_ids_on(&server, "fu_fields", "*", "delta").await, + ["f1"] + ); +} + #[tokio::test(flavor = "multi_thread", worker_threads = 4)] async fn bulk_update_moves_every_row_to_its_new_terms() { let server = TestServer::start().await; diff --git a/nodedb/tests/inproc/cases/mod.rs b/nodedb/tests/inproc/cases/mod.rs index f2eef5d32..6bd9f9821 100644 --- a/nodedb/tests/inproc/cases/mod.rs +++ b/nodedb/tests/inproc/cases/mod.rs @@ -154,6 +154,7 @@ mod quota_live_enforcement_apply; mod quota_three_level_denial; mod redaction_policy_ddl; mod reindex_concurrent_writes; +mod reindex_field_scopes; mod reindex_vector_concurrent; mod request_tracker_backpressure; mod resp_row_level_security; diff --git a/nodedb/tests/inproc/cases/reindex_field_scopes.rs b/nodedb/tests/inproc/cases/reindex_field_scopes.rs new file mode 100644 index 000000000..190aa277a --- /dev/null +++ b/nodedb/tests/inproc/cases/reindex_field_scopes.rs @@ -0,0 +1,66 @@ +// SPDX-License-Identifier: BUSL-1.1 + +//! A full-text REINDEX rebuilds every index of the collection: the +//! whole-document index and one index per text field. A field-scoped search +//! answers the same after the rebuild as before it. + +use nodedb_test_support::pgwire_harness::TestServer; + +/// Ids a full-text match on `field` returns, sorted. +async fn text_ids_on( + server: &TestServer, + collection: &str, + field: &str, + term: &str, +) -> Vec { + let mut ids = server + .query_text(&format!( + "SELECT id FROM {collection} WHERE text_match({field}, '{term}')" + )) + .await + .unwrap(); + ids.sort(); + ids +} + +/// The field scopes of `rx_fields`: `rust` in `d1`'s title and `d2`'s body. +async fn assert_scopes(server: &TestServer) { + assert_eq!( + text_ids_on(server, "rx_fields", "title", "rust").await, + ["d1"] + ); + assert_eq!( + text_ids_on(server, "rx_fields", "body", "rust").await, + ["d2"] + ); + assert_eq!( + text_ids_on(server, "rx_fields", "*", "rust").await, + ["d1", "d2"] + ); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn a_text_rebuild_preserves_field_scopes() { + let server = TestServer::start().await; + server + .exec("CREATE COLLECTION rx_fields WITH (engine='document_schemaless')") + .await + .unwrap(); + for (id, title, body) in [ + ("d1", "rust handbook", "an introduction"), + ("d2", "cooking", "the rust belt"), + ] { + server + .exec(&format!( + "INSERT INTO rx_fields {{ id: '{id}', title: '{title}', body: '{body}' }}" + )) + .await + .unwrap(); + } + assert_scopes(&server).await; + + // The plain form answers after the rebuild has cut over. + server.exec("REINDEX INDEX fts rx_fields").await.unwrap(); + + assert_scopes(&server).await; +} diff --git a/nodedb/tests/inproc/cases/sync_compat.rs b/nodedb/tests/inproc/cases/sync_compat.rs index fc19e1398..02fe1eee5 100644 --- a/nodedb/tests/inproc/cases/sync_compat.rs +++ b/nodedb/tests/inproc/cases/sync_compat.rs @@ -290,6 +290,54 @@ fn ping_pong_roundtrip() { assert!(!parsed.is_pong); } +fn fts_index_msg(fields: Vec<(String, String)>) -> FtsIndexMsg { + FtsIndexMsg { + lite_id: "lite-1".into(), + collection: "articles".into(), + doc_id: "doc-7".into(), + fields, + batch_id: 3, + producer_id: 9, + epoch: 1, + seq: 4, + } +} + +/// A Lite `FtsIndex` frame carries each string field by name, so Origin can +/// index every field on its own. +#[test] +fn lite_fts_index_frame_carries_fields() { + let fields = vec![ + ("body".to_string(), "fearless concurrency".to_string()), + ("title".to_string(), "Rust".to_string()), + ]; + let frame = + SyncFrame::new_msgpack(SyncMessageType::FtsIndex, &fts_index_msg(fields.clone())).unwrap(); + let parsed: FtsIndexMsg = SyncFrame::from_bytes(&frame.to_bytes()) + .unwrap() + .decode_body() + .unwrap(); + + assert_eq!(parsed.fields, fields); + assert_eq!(parsed.collection, "articles"); + assert_eq!(parsed.seq, 4); +} + +/// A document whose update stripped every string field syncs as an +/// `FtsIndex` frame with no fields: Origin applies it as a removal. +#[test] +fn lite_fts_index_frame_with_no_fields_roundtrips() { + let frame = + SyncFrame::new_msgpack(SyncMessageType::FtsIndex, &fts_index_msg(Vec::new())).unwrap(); + let parsed: FtsIndexMsg = SyncFrame::from_bytes(&frame.to_bytes()) + .unwrap() + .decode_body() + .unwrap(); + + assert!(parsed.fields.is_empty()); + assert_eq!(parsed.doc_id, "doc-7"); +} + #[test] fn all_11_message_types_valid() { for (code, expected) in [ From 820a73ebd7b331cc3d550e2bd631da08c13218e9 Mon Sep 17 00:00:00 2001 From: Farhan Syah Date: Thu, 8 Oct 2026 16:51:47 +0800 Subject: [PATCH 15/24] feat(text-search): scope text queries to a column with named options text_match and bm25_score take a column (or * for the whole document) and named options (mode, fuzzy). A column that cannot serve a search is a typed error (42703 or 42804). AND mode returns the true top-k of the documents holding every word, a word matching through its synonyms. Negative terms and residual WHERE and RLS filters apply before ranking, and equal scores order by row. A bm25_score in the SELECT list becomes a per-row score column: 0 for a row the index holds that does not match, NULL for a row it does not hold. Hybrid and three-source RRF searches take a field and mode and filter each leg before fusion. Staged transaction rows score through the same scorer, and a phrase is analyzed once with the collection's analyzer. --- nodedb-fts/src/analyzer/language/stemmer.rs | 13 +- nodedb-fts/src/analyzer/pipeline.rs | 76 +- nodedb-fts/src/index/synonym_groups.rs | 30 +- nodedb-fts/src/posting.rs | 16 + nodedb-fts/src/search/bm25_search.rs | 681 ++++++-------- nodedb-fts/src/search/bmw/heap.rs | 81 +- nodedb-fts/src/search/bmw/mod.rs | 1 - nodedb-fts/src/search/bmw/query.rs | 288 ------ nodedb-fts/src/search/bmw/scorer.rs | 241 +++-- nodedb-fts/src/search/bmw/skip_index.rs | 28 + nodedb-fts/src/search/doc_score.rs | 108 +++ nodedb-fts/src/search/doc_scorer.rs | 271 ++++++ nodedb-fts/src/search/fuzzy_search.rs | 110 +-- nodedb-fts/src/search/match_mode.rs | 147 +++ nodedb-fts/src/search/mod.rs | 5 + nodedb-fts/src/search/phrase.rs | 29 - nodedb-fts/src/search/query_terms.rs | 289 ++++++ nodedb-fts/src/search/staged.rs | 171 ++++ nodedb-physical/src/physical_plan/text.rs | 145 ++- .../src/functions/builtins/scalars/vector.rs | 5 +- nodedb-sql/src/planner/ast_helpers.rs | 12 + .../src/planner/search_scope/expressions.rs | 79 ++ nodedb-sql/src/planner/search_scope/lookup.rs | 140 +++ nodedb-sql/src/planner/search_scope/mod.rs | 23 + .../scope.rs} | 360 ++------ nodedb-sql/src/planner/select/entry.rs | 847 ------------------ .../src/planner/select/entry/fixtures.rs | 200 +++++ nodedb-sql/src/planner/select/entry/mod.rs | 12 + .../src/planner/select/entry/payload.rs | 250 ++++++ nodedb-sql/src/planner/select/entry/query.rs | 232 +++++ nodedb-sql/src/planner/select/entry/search.rs | 797 ++++++++++++++++ nodedb-sql/src/planner/select/helpers.rs | 2 +- nodedb-sql/src/planner/select/limit.rs | 57 +- nodedb-sql/src/planner/select/mod.rs | 2 + .../src/planner/select/order_by/apply.rs | 59 +- .../planner/select/order_by/graph_score.rs | 218 +++++ .../src/planner/select/order_by/hybrid.rs | 312 ++++--- nodedb-sql/src/planner/select/order_by/mod.rs | 4 + .../src/planner/select/order_by/projection.rs | 171 +--- .../src/planner/select/order_by/text_score.rs | 94 ++ .../src/planner/select/order_by/triggers.rs | 162 +++- nodedb-sql/src/planner/select/text_call.rs | 113 +++ nodedb-sql/src/planner/select/text_options.rs | 252 ++++++ nodedb-sql/src/planner/select/where_search.rs | 175 ++-- nodedb-sql/src/resolver/mod.rs | 1 + nodedb-sql/src/resolver/text_column.rs | 116 +++ nodedb-sql/src/types/plan/variants/text.rs | 63 ++ nodedb-sql/src/visitor/plan_visitor/args.rs | 28 +- .../src/visitor/plan_visitor/trait_def.rs | 15 +- nodedb-types/src/text_search.rs | 107 ++- nodedb/src/control/clone/resolver/rewrite.rs | 25 +- .../control/planner/redaction_refusal/mod.rs | 1 + .../control/planner/redaction_refusal/text.rs | 243 +++++ .../rls_injection/permission_tree/text.rs | 64 +- .../src/control/planner/rls_injection/text.rs | 146 ++- .../planner/sql_plan_convert/filter/mod.rs | 3 +- .../planner/sql_plan_convert/scan/mod.rs | 8 +- .../planner/sql_plan_convert/scan/search.rs | 135 +-- .../sql_plan_convert/scan/text_score_bound.rs | 355 ++++++++ .../sql_plan_convert/scan/text_search.rs | 96 ++ .../planner/sql_plan_convert/scan_params.rs | 26 + .../planner/sql_plan_convert/set_ops.rs | 82 +- .../visitor/arms_scan_search.rs | 44 +- .../native/dispatch/plan_builder/text.rs | 63 +- nodedb/src/control/server/pgwire/types/mod.rs | 7 +- .../server/response_translate/dispatch.rs | 6 +- .../shared/ddl/neutral/column_default.rs | 5 +- .../predicate/txn_buffering/classify.rs | 47 +- nodedb/src/data/executor/dispatch/text.rs | 69 +- .../src/data/executor/handlers/hybrid_key.rs | 23 + .../data/executor/handlers/hybrid_overlay.rs | 131 +-- .../src/data/executor/handlers/text_rows.rs | 280 ++++++ .../executor/handlers/text_score_columns.rs | 128 +++ .../data/executor/handlers/text_score_sink.rs | 277 ++++++ .../src/data/executor/handlers/text_search.rs | 249 +++-- .../executor/handlers/text_search_hybrid.rs | 111 +-- .../executor/handlers/text_search_scan.rs | 403 +++++---- .../executor/handlers/text_search_triple.rs | 116 +-- .../handlers/transaction/overlay/fts_merge.rs | 500 ++--------- .../handlers/transaction/overlay/fts_score.rs | 232 ++--- .../handlers/transaction/overlay/mod.rs | 2 +- nodedb/src/engine/sparse/btree_scan_while.rs | 52 ++ .../engine/sparse/inverted/staged_search.rs | 177 ++++ nodedb/src/engine/sparse/mod.rs | 1 + .../cross_engine_three_way_fts_vector_doc.rs | 4 + .../test_cross_engine_validation.rs | 22 +- .../test_tenant_isolation_fulltext.rs | 8 + ...test_tenant_isolation_fulltext_negative.rs | 8 + nodedb/tests/wire/cases/engine_surface_fts.rs | 332 ++++++- .../wire/cases/engine_surface_fts_options.rs | 175 ++++ .../tests/wire/cases/fts_query_semantics.rs | 379 ++++++++ nodedb/tests/wire/cases/sql_fts_strict.rs | 71 ++ nodedb/tests/wire/cases/sql_hybrid_search.rs | 141 +++ .../cases/sql_search_subquery_composition.rs | 6 +- .../wire/cases/sql_three_source_rrf_scoped.rs | 118 +++ .../sql_transactions_fts_analyzer_overlay.rs | 54 ++ .../cases/sql_transactions_fts_overlay.rs | 99 ++ .../cases/sql_transactions_hybrid_overlay.rs | 88 ++ 98 files changed, 9464 insertions(+), 3786 deletions(-) delete mode 100644 nodedb-fts/src/search/bmw/query.rs create mode 100644 nodedb-fts/src/search/doc_score.rs create mode 100644 nodedb-fts/src/search/doc_scorer.rs create mode 100644 nodedb-fts/src/search/match_mode.rs create mode 100644 nodedb-fts/src/search/query_terms.rs create mode 100644 nodedb-fts/src/search/staged.rs create mode 100644 nodedb-sql/src/planner/search_scope/expressions.rs create mode 100644 nodedb-sql/src/planner/search_scope/lookup.rs create mode 100644 nodedb-sql/src/planner/search_scope/mod.rs rename nodedb-sql/src/planner/{search_scope.rs => search_scope/scope.rs} (50%) delete mode 100644 nodedb-sql/src/planner/select/entry.rs create mode 100644 nodedb-sql/src/planner/select/entry/fixtures.rs create mode 100644 nodedb-sql/src/planner/select/entry/mod.rs create mode 100644 nodedb-sql/src/planner/select/entry/payload.rs create mode 100644 nodedb-sql/src/planner/select/entry/query.rs create mode 100644 nodedb-sql/src/planner/select/entry/search.rs create mode 100644 nodedb-sql/src/planner/select/order_by/graph_score.rs create mode 100644 nodedb-sql/src/planner/select/order_by/text_score.rs create mode 100644 nodedb-sql/src/planner/select/text_call.rs create mode 100644 nodedb-sql/src/planner/select/text_options.rs create mode 100644 nodedb-sql/src/resolver/text_column.rs create mode 100644 nodedb-sql/src/types/plan/variants/text.rs create mode 100644 nodedb/src/control/planner/redaction_refusal/text.rs create mode 100644 nodedb/src/control/planner/sql_plan_convert/scan/text_score_bound.rs create mode 100644 nodedb/src/control/planner/sql_plan_convert/scan/text_search.rs create mode 100644 nodedb/src/data/executor/handlers/text_rows.rs create mode 100644 nodedb/src/data/executor/handlers/text_score_columns.rs create mode 100644 nodedb/src/data/executor/handlers/text_score_sink.rs create mode 100644 nodedb/src/engine/sparse/btree_scan_while.rs create mode 100644 nodedb/src/engine/sparse/inverted/staged_search.rs create mode 100644 nodedb/tests/wire/cases/engine_surface_fts_options.rs create mode 100644 nodedb/tests/wire/cases/fts_query_semantics.rs create mode 100644 nodedb/tests/wire/cases/sql_three_source_rrf_scoped.rs 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/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/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/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-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-sql/src/functions/builtins/scalars/vector.rs b/nodedb-sql/src/functions/builtins/scalars/vector.rs index ef1571f40..3aa462f65 100644 --- a/nodedb-sql/src/functions/builtins/scalars/vector.rs +++ b/nodedb-sql/src/functions/builtins/scalars/vector.rs @@ -83,11 +83,12 @@ pub(super) fn vector_functions() -> Vec { Some(ColumnType::Float64), arg_types::SEARCH_SCORE_ARGS, ), + // Options are named (`mode => 'or'`, `fuzzy => true`), not positional. m( "text_match", Scalar, 2, - 3, + 2, SearchTrigger::TextMatch, None, arg_types::TEXT_MATCH_ARGS, @@ -98,7 +99,7 @@ pub(super) fn vector_functions() -> Vec { "search", Scalar, 2, - 3, + 2, SearchTrigger::TextMatch, None, arg_types::TEXT_MATCH_ARGS, diff --git a/nodedb-sql/src/planner/ast_helpers.rs b/nodedb-sql/src/planner/ast_helpers.rs index d6723e1ee..b6de6966f 100644 --- a/nodedb-sql/src/planner/ast_helpers.rs +++ b/nodedb-sql/src/planner/ast_helpers.rs @@ -7,6 +7,7 @@ use sqlparser::ast; use crate::error::{Result, SqlError}; use crate::parser::normalize::{normalize_ident, normalize_object_name_checked}; use crate::planner::select::convert_where_to_filters; +use crate::resolver::columns::ResolvedTable; use crate::types::Filter; /// Return `(table, column)` for a `table.col` compound identifier, or `None`. @@ -186,6 +187,17 @@ pub fn strip_single_table_qualifiers( Ok(out) } +/// The qualifiers a single-table query's column references may carry: the +/// table's name, and its alias when it has one. +pub(crate) fn single_table_qualifiers(table: &ResolvedTable) -> Vec<&str> { + let ref_name = table.ref_name(); + if ref_name == table.name { + vec![table.name.as_str()] + } else { + vec![table.name.as_str(), ref_name] + } +} + /// Strip `qualifier.` from all compound identifiers in `expr`, then convert /// the result to `Vec` via `convert_where_to_filters`. pub fn strip_and_convert_filters( diff --git a/nodedb-sql/src/planner/search_scope/expressions.rs b/nodedb-sql/src/planner/search_scope/expressions.rs new file mode 100644 index 000000000..ca4722fbf --- /dev/null +++ b/nodedb-sql/src/planner/search_scope/expressions.rs @@ -0,0 +1,79 @@ +// SPDX-License-Identifier: Apache-2.0 + +//! Row-expression search-function checks. + +use super::lookup::first_search_function; +use super::scope::Scope; +use crate::error::{Result, SqlError}; +use crate::types::query::{AggregateExpr, Projection, SortKey, WindowSpec}; +use crate::types::{Filter, FilterExpr}; +use crate::types_expr::SqlExpr; + +impl Scope<'_> { + pub(super) fn filters(&self, filters: &[Filter]) -> Result<()> { + filters + .iter() + .try_for_each(|filter| self.filter(&filter.expr)) + } + + pub(super) fn filter(&self, expr: &FilterExpr) -> Result<()> { + match expr { + FilterExpr::Expr(expr) => self.expr(expr), + FilterExpr::And(children) | FilterExpr::Or(children) => self.filters(children), + FilterExpr::Not(child) => self.filter(&child.expr), + FilterExpr::Comparison { .. } + | FilterExpr::InList { .. } + | FilterExpr::Between { .. } + | FilterExpr::IsNull { .. } + | FilterExpr::IsNotNull { .. } => Ok(()), + } + } + + pub(super) fn projection(&self, projection: &[Projection]) -> Result<()> { + for item in projection { + match item { + Projection::Computed { expr, .. } | Projection::CpComputed { expr, .. } => { + self.expr(expr)? + } + Projection::Column(_) | Projection::Star | Projection::QualifiedStar(_) => {} + } + } + Ok(()) + } + + pub(super) fn sort_keys(&self, sort_keys: &[SortKey]) -> Result<()> { + sort_keys.iter().try_for_each(|key| self.expr(&key.expr)) + } + + pub(super) fn windows(&self, windows: &[WindowSpec]) -> Result<()> { + for window in windows { + self.exprs(&window.args)?; + self.exprs(&window.partition_by)?; + self.sort_keys(&window.order_by)?; + } + Ok(()) + } + + pub(super) fn aggregates(&self, aggregates: &[AggregateExpr]) -> Result<()> { + aggregates + .iter() + .try_for_each(|aggregate| self.exprs(&aggregate.args)) + } + + pub(super) fn assignments(&self, assignments: &[(String, SqlExpr)]) -> Result<()> { + assignments.iter().try_for_each(|(_, expr)| self.expr(expr)) + } + + pub(super) fn exprs(&self, exprs: &[SqlExpr]) -> Result<()> { + exprs.iter().try_for_each(|expr| self.expr(expr)) + } + + pub(super) fn expr(&self, expr: &SqlExpr) -> Result<()> { + match first_search_function(expr, self.functions) { + Some(name) => Err(SqlError::SearchFunctionOutsideSearch { + name: name.to_owned(), + }), + None => Ok(()), + } + } +} diff --git a/nodedb-sql/src/planner/search_scope/lookup.rs b/nodedb-sql/src/planner/search_scope/lookup.rs new file mode 100644 index 000000000..e57e91261 --- /dev/null +++ b/nodedb-sql/src/planner/search_scope/lookup.rs @@ -0,0 +1,140 @@ +// SPDX-License-Identifier: Apache-2.0 + +//! Which functions are index-owned, and which plans serve a score column. + +use crate::functions::registry::{FunctionRegistry, SearchTrigger}; +use crate::types::SqlPlan; +use crate::types::{CtePlan, LateralLoopPlan, LateralTopKPlan}; +use crate::types_expr::SqlExpr; + +/// Whether the search trigger `trigger` names a function that reads an +/// index and has no per-row value. +pub(super) fn is_index_owned(trigger: SearchTrigger) -> bool { + match trigger { + SearchTrigger::MultiVectorSearch + | SearchTrigger::SparseSearch + | SearchTrigger::TextSearch + | SearchTrigger::HybridSearch + | SearchTrigger::TextMatch + | SearchTrigger::GraphSearch => true, + // The vector distances and the spatial predicates evaluate per row; + // the time bucket is a scalar; the array functions are table-valued + // and planned from FROM. + SearchTrigger::None + | SearchTrigger::VectorSearch + | SearchTrigger::SpatialDWithin + | SearchTrigger::SpatialContains + | SearchTrigger::SpatialIntersects + | SearchTrigger::SpatialWithin + | SearchTrigger::TimeBucket + | SearchTrigger::ArraySlice + | SearchTrigger::ArrayProject + | SearchTrigger::ArrayAgg + | SearchTrigger::ArrayElementwise + | SearchTrigger::ArrayFlush + | SearchTrigger::ArrayCompact => false, + } +} + +/// The first index-owned search function `expr` calls outside a subquery. +pub(super) fn first_search_function<'e>( + expr: &'e SqlExpr, + functions: &FunctionRegistry, +) -> Option<&'e str> { + let find = |e: &'e SqlExpr| first_search_function(e, functions); + match expr { + SqlExpr::Function { name, args, .. } => { + if is_index_owned(functions.search_trigger(name)) { + return Some(name.as_str()); + } + args.iter().find_map(find) + } + SqlExpr::BinaryOp { left, right, .. } => find(left).or_else(|| find(right)), + SqlExpr::UnaryOp { expr, .. } + | SqlExpr::Cast { expr, .. } + | SqlExpr::IsNull { expr, .. } => find(expr), + SqlExpr::Case { + operand, + when_then, + else_expr, + } => operand + .as_deref() + .and_then(find) + .or_else(|| { + when_then + .iter() + .find_map(|(when, then)| find(when).or_else(|| find(then))) + }) + .or_else(|| else_expr.as_deref().and_then(find)), + SqlExpr::InList { expr, list, .. } => find(expr).or_else(|| list.iter().find_map(find)), + SqlExpr::Between { + expr, low, high, .. + } => find(expr).or_else(|| find(low)).or_else(|| find(high)), + SqlExpr::Like { expr, pattern, .. } => find(expr).or_else(|| find(pattern)), + SqlExpr::ArrayLiteral(items) => items.iter().find_map(find), + SqlExpr::Column { .. } | SqlExpr::Literal(_) | SqlExpr::Subquery(_) | SqlExpr::Wildcard => { + None + } + } +} + +/// Whether `plan` is, or wraps, a search plan that serves a score column. +pub(super) fn has_search_plan(plan: &SqlPlan) -> bool { + match plan { + SqlPlan::VectorSearch { .. } + | SqlPlan::MultiVectorSearch { .. } + | SqlPlan::SparseSearch { .. } + | SqlPlan::TextSearch { .. } + | SqlPlan::HybridSearch { .. } + | SqlPlan::HybridSearchTriple { .. } => true, + SqlPlan::Subquery { input, .. } | SqlPlan::Aggregate { input, .. } => { + has_search_plan(input) + } + SqlPlan::Join { left, right, .. } => has_search_plan(left) || has_search_plan(right), + SqlPlan::LateralTopK(LateralTopKPlan { outer, .. }) => has_search_plan(outer), + SqlPlan::LateralLoop(LateralLoopPlan { outer, inner, .. }) => { + has_search_plan(outer) || has_search_plan(inner) + } + SqlPlan::Union { inputs, .. } => inputs.iter().any(has_search_plan), + SqlPlan::Intersect { left, right, .. } | SqlPlan::Except { left, right, .. } => { + has_search_plan(left) || has_search_plan(right) + } + SqlPlan::Cte(CtePlan { outer, .. }) => has_search_plan(outer), + SqlPlan::ConstantResult { .. } + | SqlPlan::Scan { .. } + | SqlPlan::PointGet { .. } + | SqlPlan::DocumentIndexLookup { .. } + | SqlPlan::RangeScan { .. } + | SqlPlan::Insert { .. } + | SqlPlan::KvInsert { .. } + | SqlPlan::Upsert { .. } + | SqlPlan::InsertSelect { .. } + | SqlPlan::Update { .. } + | SqlPlan::UpdateFrom { .. } + | SqlPlan::Delete { .. } + | SqlPlan::Truncate { .. } + | SqlPlan::TimeseriesScan { .. } + | SqlPlan::TimeseriesIngest { .. } + | SqlPlan::SpatialScan { .. } + | SqlPlan::RecursiveScan { .. } + | SqlPlan::RecursiveValue { .. } + | SqlPlan::CreateArray { .. } + | SqlPlan::DropArray { .. } + | SqlPlan::AlterArray { .. } + | SqlPlan::InsertArray { .. } + | SqlPlan::DeleteArray { .. } + | SqlPlan::ArraySlice { .. } + | SqlPlan::ArrayProject { .. } + | SqlPlan::ArrayAgg { .. } + | SqlPlan::ArrayElementwise { .. } + | SqlPlan::ArrayFlush { .. } + | SqlPlan::ArrayCompact { .. } + | SqlPlan::Merge { .. } + | SqlPlan::VectorPrimaryInsert { .. } + | SqlPlan::VectorPrimaryDelete { .. } + | SqlPlan::VectorPrimaryTruncate { .. } + | SqlPlan::VectorPrimaryUpdate { .. } + | SqlPlan::CreateIndex { .. } + | SqlPlan::DropIndex { .. } => false, + } +} diff --git a/nodedb-sql/src/planner/search_scope/mod.rs b/nodedb-sql/src/planner/search_scope/mod.rs new file mode 100644 index 000000000..768d487c9 --- /dev/null +++ b/nodedb-sql/src/planner/search_scope/mod.rs @@ -0,0 +1,23 @@ +// SPDX-License-Identifier: Apache-2.0 + +//! Plan-time refusal of index-owned search functions in row-evaluated +//! positions. +//! +//! `bm25_score`, `search_score`, `text_match`, `search`, `rrf_score`, +//! `sparse_score`, `graph_score`, `multi_vector_score` and +//! `multi_vector_search` read a search index. The planner lowers each call it +//! recognises into its search plan, which serves the score as a column. A +//! call the planner could not lower stays in a filter, projection, sort key, +//! assignment or aggregate argument. The row evaluator has no index and no +//! value for it, so this pass refuses the statement at plan time. The +//! refusal does not depend on whether the collection holds rows. +//! +//! A wrapper plan (a subquery tail, aggregate, join or lateral join) over a +//! search plan is not checked for its own expressions: its projection names +//! the score column the search plan serves. + +mod expressions; +mod lookup; +mod scope; + +pub use scope::refuse_row_scoped_search_functions; diff --git a/nodedb-sql/src/planner/search_scope.rs b/nodedb-sql/src/planner/search_scope/scope.rs similarity index 50% rename from nodedb-sql/src/planner/search_scope.rs rename to nodedb-sql/src/planner/search_scope/scope.rs index 76a595276..8be05e38a 100644 --- a/nodedb-sql/src/planner/search_scope.rs +++ b/nodedb-sql/src/planner/search_scope/scope.rs @@ -1,26 +1,19 @@ // SPDX-License-Identifier: Apache-2.0 -//! Plan-time refusal of index-owned search functions in row-evaluated -//! positions. -//! -//! `bm25_score`, `search_score`, `text_match`, `search`, `rrf_score`, -//! `sparse_score`, `graph_score`, `multi_vector_score` and -//! `multi_vector_search` read a search index. The planner lowers each call it -//! recognises into its search plan, which serves the score as a column. A -//! call the planner could not lower stays in a filter, projection, sort key, -//! assignment or aggregate argument. The row evaluator has no index and no -//! value for it, so this pass refuses the statement at plan time. The -//! refusal does not depend on whether the collection holds rows. -//! -//! A wrapper plan (a subquery tail, aggregate, join or lateral join) over a -//! search plan is not checked for its own expressions: its projection names -//! the score column the search plan serves. +//! The pass that walks a plan and refuses a search function in a +//! row-evaluated position. -use crate::error::{Result, SqlError}; -use crate::functions::registry::{FunctionRegistry, SearchTrigger}; -use crate::types::query::{AggregateExpr, Projection, SortKey, WindowSpec}; -use crate::types::{Filter, FilterExpr, MergePlanAction, SqlPlan}; -use crate::types_expr::SqlExpr; +use crate::error::Result; +use crate::functions::registry::FunctionRegistry; +use crate::types::{ + CtePlan, DocumentIndexLookupPlan, HybridSearchPlan, HybridSearchTriplePlan, KvInsertPlan, + LateralLoopPlan, LateralTopKPlan, MergePlan, RangeScanPlan, RecursiveScanPlan, TextSearchPlan, + TimeseriesScanPlan, UpsertPlan, VectorPrimaryDeletePlan, VectorPrimaryInsertPlan, + VectorPrimaryUpdatePlan, +}; +use crate::types::{MergePlanAction, SqlPlan}; + +use super::lookup::has_search_plan; /// Refuse `plan` when an index-owned search function sits where the row /// evaluator runs it. @@ -31,8 +24,8 @@ pub fn refuse_row_scoped_search_functions( Scope { functions }.plan(plan) } -struct Scope<'a> { - functions: &'a FunctionRegistry, +pub(super) struct Scope<'a> { + pub(super) functions: &'a FunctionRegistry, } impl Scope<'_> { @@ -45,33 +38,32 @@ impl Scope<'_> { window_functions, .. } - | SqlPlan::DocumentIndexLookup { + | SqlPlan::DocumentIndexLookup(DocumentIndexLookupPlan { filters, projection, sort_keys, window_functions, .. - } => { + }) => { self.filters(filters)?; self.projection(projection)?; self.sort_keys(sort_keys)?; self.windows(window_functions) } - SqlPlan::PointGet { projection, .. } | SqlPlan::RangeScan { projection, .. } => { - self.projection(projection) - } - SqlPlan::KvInsert { + SqlPlan::PointGet { projection, .. } + | SqlPlan::RangeScan(RangeScanPlan { projection, .. }) => self.projection(projection), + SqlPlan::KvInsert(KvInsertPlan { on_conflict_updates, .. - } - | SqlPlan::Upsert { + }) + | SqlPlan::Upsert(UpsertPlan { on_conflict_updates, .. - } - | SqlPlan::VectorPrimaryInsert { + }) + | SqlPlan::VectorPrimaryInsert(VectorPrimaryInsertPlan { on_conflict_updates, .. - } => self.assignments(on_conflict_updates), + }) => self.assignments(on_conflict_updates), SqlPlan::InsertSelect { source, column_map, .. } => { @@ -83,11 +75,11 @@ impl Scope<'_> { filters, .. } - | SqlPlan::VectorPrimaryUpdate { + | SqlPlan::VectorPrimaryUpdate(VectorPrimaryUpdatePlan { assignments, filters, .. - } => { + }) => { self.assignments(assignments)?; self.filters(filters) } @@ -101,7 +93,8 @@ impl Scope<'_> { self.assignments(assignments)?; self.filters(target_filters) } - SqlPlan::Delete { filters, .. } | SqlPlan::VectorPrimaryDelete { filters, .. } => { + SqlPlan::Delete { filters, .. } + | SqlPlan::VectorPrimaryDelete(VectorPrimaryDeletePlan { filters, .. }) => { self.filters(filters) } SqlPlan::Join { @@ -140,13 +133,13 @@ impl Scope<'_> { self.filters(having)?; self.sort_keys(sort_keys) } - SqlPlan::TimeseriesScan { + SqlPlan::TimeseriesScan(TimeseriesScanPlan { aggregates, filters, projection, sort_keys, .. - } => { + }) => { self.aggregates(aggregates)?; self.filters(filters)?; self.projection(projection)?; @@ -154,17 +147,20 @@ impl Scope<'_> { } // A search plan serves its own score call as a column; only its // residual filters run on the row evaluator. - SqlPlan::VectorSearch { filters, .. } | SqlPlan::TextSearch { filters, .. } => { + SqlPlan::VectorSearch { filters, .. } + | SqlPlan::TextSearch(TextSearchPlan { filters, .. }) + | SqlPlan::HybridSearch(HybridSearchPlan { filters, .. }) + | SqlPlan::HybridSearchTriple(HybridSearchTriplePlan { filters, .. }) => { self.filters(filters) } SqlPlan::SpatialScan { attribute_filters, .. } => self.filters(attribute_filters), - SqlPlan::RecursiveScan { + SqlPlan::RecursiveScan(RecursiveScanPlan { base_filters, recursive_filters, .. - } => { + }) => { self.filters(base_filters)?; self.filters(recursive_filters) } @@ -173,7 +169,7 @@ impl Scope<'_> { self.plan(left)?; self.plan(right) } - SqlPlan::Cte { definitions, outer } => { + SqlPlan::Cte(CtePlan { definitions, outer }) => { for (_, definition) in definitions { self.plan(definition)?; } @@ -196,9 +192,9 @@ impl Scope<'_> { self.windows(window_functions)?; self.sort_keys(sort_keys) } - SqlPlan::Merge { + SqlPlan::Merge(MergePlan { source, clauses, .. - } => { + }) => { self.plan(source)?; for clause in clauses { self.filters(&clause.extra_predicate)?; @@ -210,13 +206,13 @@ impl Scope<'_> { } Ok(()) } - SqlPlan::LateralTopK { + SqlPlan::LateralTopK(LateralTopKPlan { outer, inner_filters, inner_order_by, projection, .. - } => { + }) => { self.plan(outer)?; self.filters(inner_filters)?; self.sort_keys(inner_order_by)?; @@ -225,12 +221,12 @@ impl Scope<'_> { } self.projection(projection) } - SqlPlan::LateralLoop { + SqlPlan::LateralLoop(LateralLoopPlan { outer, inner, projection, .. - } => { + }) => { self.plan(outer)?; self.plan(inner)?; if has_search_plan(outer) || has_search_plan(inner) { @@ -247,8 +243,6 @@ impl Scope<'_> { | SqlPlan::TimeseriesIngest { .. } | SqlPlan::MultiVectorSearch { .. } | SqlPlan::SparseSearch { .. } - | SqlPlan::HybridSearch { .. } - | SqlPlan::HybridSearchTriple { .. } | SqlPlan::RecursiveValue { .. } | SqlPlan::CreateArray { .. } | SqlPlan::DropArray { .. } @@ -266,208 +260,17 @@ impl Scope<'_> { | SqlPlan::DropIndex { .. } => Ok(()), } } - - fn filters(&self, filters: &[Filter]) -> Result<()> { - filters - .iter() - .try_for_each(|filter| self.filter(&filter.expr)) - } - - fn filter(&self, expr: &FilterExpr) -> Result<()> { - match expr { - FilterExpr::Expr(expr) => self.expr(expr), - FilterExpr::And(children) | FilterExpr::Or(children) => self.filters(children), - FilterExpr::Not(child) => self.filter(&child.expr), - FilterExpr::Comparison { .. } - | FilterExpr::InList { .. } - | FilterExpr::Between { .. } - | FilterExpr::IsNull { .. } - | FilterExpr::IsNotNull { .. } => Ok(()), - } - } - - fn projection(&self, projection: &[Projection]) -> Result<()> { - for item in projection { - match item { - Projection::Computed { expr, .. } | Projection::CpComputed { expr, .. } => { - self.expr(expr)? - } - Projection::Column(_) | Projection::Star | Projection::QualifiedStar(_) => {} - } - } - Ok(()) - } - - fn sort_keys(&self, sort_keys: &[SortKey]) -> Result<()> { - sort_keys.iter().try_for_each(|key| self.expr(&key.expr)) - } - - fn windows(&self, windows: &[WindowSpec]) -> Result<()> { - for window in windows { - self.exprs(&window.args)?; - self.exprs(&window.partition_by)?; - self.sort_keys(&window.order_by)?; - } - Ok(()) - } - - fn aggregates(&self, aggregates: &[AggregateExpr]) -> Result<()> { - aggregates - .iter() - .try_for_each(|aggregate| self.exprs(&aggregate.args)) - } - - fn assignments(&self, assignments: &[(String, SqlExpr)]) -> Result<()> { - assignments.iter().try_for_each(|(_, expr)| self.expr(expr)) - } - - fn exprs(&self, exprs: &[SqlExpr]) -> Result<()> { - exprs.iter().try_for_each(|expr| self.expr(expr)) - } - - fn expr(&self, expr: &SqlExpr) -> Result<()> { - match first_search_function(expr, self.functions) { - Some(name) => Err(SqlError::SearchFunctionOutsideSearch { - name: name.to_owned(), - }), - None => Ok(()), - } - } -} - -/// Whether the search trigger `trigger` names a function that reads an -/// index and has no per-row value. -fn is_index_owned(trigger: SearchTrigger) -> bool { - match trigger { - SearchTrigger::MultiVectorSearch - | SearchTrigger::SparseSearch - | SearchTrigger::TextSearch - | SearchTrigger::HybridSearch - | SearchTrigger::TextMatch - | SearchTrigger::GraphSearch => true, - // The vector distances and the spatial predicates evaluate per row; - // the time bucket is a scalar; the array functions are table-valued - // and planned from FROM. - SearchTrigger::None - | SearchTrigger::VectorSearch - | SearchTrigger::SpatialDWithin - | SearchTrigger::SpatialContains - | SearchTrigger::SpatialIntersects - | SearchTrigger::SpatialWithin - | SearchTrigger::TimeBucket - | SearchTrigger::ArraySlice - | SearchTrigger::ArrayProject - | SearchTrigger::ArrayAgg - | SearchTrigger::ArrayElementwise - | SearchTrigger::ArrayFlush - | SearchTrigger::ArrayCompact => false, - } -} - -/// The first index-owned search function `expr` calls outside a subquery. -fn first_search_function<'e>(expr: &'e SqlExpr, functions: &FunctionRegistry) -> Option<&'e str> { - let find = |e: &'e SqlExpr| first_search_function(e, functions); - match expr { - SqlExpr::Function { name, args, .. } => { - if is_index_owned(functions.search_trigger(name)) { - return Some(name.as_str()); - } - args.iter().find_map(find) - } - SqlExpr::BinaryOp { left, right, .. } => find(left).or_else(|| find(right)), - SqlExpr::UnaryOp { expr, .. } - | SqlExpr::Cast { expr, .. } - | SqlExpr::IsNull { expr, .. } => find(expr), - SqlExpr::Case { - operand, - when_then, - else_expr, - } => operand - .as_deref() - .and_then(find) - .or_else(|| { - when_then - .iter() - .find_map(|(when, then)| find(when).or_else(|| find(then))) - }) - .or_else(|| else_expr.as_deref().and_then(find)), - SqlExpr::InList { expr, list, .. } => find(expr).or_else(|| list.iter().find_map(find)), - SqlExpr::Between { - expr, low, high, .. - } => find(expr).or_else(|| find(low)).or_else(|| find(high)), - SqlExpr::Like { expr, pattern, .. } => find(expr).or_else(|| find(pattern)), - SqlExpr::ArrayLiteral(items) => items.iter().find_map(find), - SqlExpr::Column { .. } | SqlExpr::Literal(_) | SqlExpr::Subquery(_) | SqlExpr::Wildcard => { - None - } - } -} - -/// Whether `plan` is, or wraps, a search plan that serves a score column. -fn has_search_plan(plan: &SqlPlan) -> bool { - match plan { - SqlPlan::VectorSearch { .. } - | SqlPlan::MultiVectorSearch { .. } - | SqlPlan::SparseSearch { .. } - | SqlPlan::TextSearch { .. } - | SqlPlan::HybridSearch { .. } - | SqlPlan::HybridSearchTriple { .. } => true, - SqlPlan::Subquery { input, .. } | SqlPlan::Aggregate { input, .. } => { - has_search_plan(input) - } - SqlPlan::Join { left, right, .. } => has_search_plan(left) || has_search_plan(right), - SqlPlan::LateralTopK { outer, .. } => has_search_plan(outer), - SqlPlan::LateralLoop { outer, inner, .. } => { - has_search_plan(outer) || has_search_plan(inner) - } - SqlPlan::Union { inputs, .. } => inputs.iter().any(has_search_plan), - SqlPlan::Intersect { left, right, .. } | SqlPlan::Except { left, right, .. } => { - has_search_plan(left) || has_search_plan(right) - } - SqlPlan::Cte { outer, .. } => has_search_plan(outer), - SqlPlan::ConstantResult { .. } - | SqlPlan::Scan { .. } - | SqlPlan::PointGet { .. } - | SqlPlan::DocumentIndexLookup { .. } - | SqlPlan::RangeScan { .. } - | SqlPlan::Insert { .. } - | SqlPlan::KvInsert { .. } - | SqlPlan::Upsert { .. } - | SqlPlan::InsertSelect { .. } - | SqlPlan::Update { .. } - | SqlPlan::UpdateFrom { .. } - | SqlPlan::Delete { .. } - | SqlPlan::Truncate { .. } - | SqlPlan::TimeseriesScan { .. } - | SqlPlan::TimeseriesIngest { .. } - | SqlPlan::SpatialScan { .. } - | SqlPlan::RecursiveScan { .. } - | SqlPlan::RecursiveValue { .. } - | SqlPlan::CreateArray { .. } - | SqlPlan::DropArray { .. } - | SqlPlan::AlterArray { .. } - | SqlPlan::InsertArray { .. } - | SqlPlan::DeleteArray { .. } - | SqlPlan::ArraySlice { .. } - | SqlPlan::ArrayProject { .. } - | SqlPlan::ArrayAgg { .. } - | SqlPlan::ArrayElementwise { .. } - | SqlPlan::ArrayFlush { .. } - | SqlPlan::ArrayCompact { .. } - | SqlPlan::Merge { .. } - | SqlPlan::VectorPrimaryInsert { .. } - | SqlPlan::VectorPrimaryDelete { .. } - | SqlPlan::VectorPrimaryTruncate { .. } - | SqlPlan::VectorPrimaryUpdate { .. } - | SqlPlan::CreateIndex { .. } - | SqlPlan::DropIndex { .. } => false, - } } #[cfg(test)] mod tests { - use super::*; + use super::refuse_row_scoped_search_functions; + use crate::error::{Result, SqlError}; + use crate::functions::registry::FunctionRegistry; use crate::types::query::EngineType; + use crate::types::query::Projection; + use crate::types::{Filter, FilterExpr, SqlPlan}; + use crate::types_expr::SqlExpr; use crate::types_expr::SqlValue; fn call(name: &str) -> SqlExpr { @@ -558,38 +361,57 @@ mod tests { assert_eq!(check(&plan), Ok(())); } - #[test] - fn a_search_plan_projection_serves_its_score() { - let plan = SqlPlan::TextSearch { + /// A text search that scores `rust` in every document under alias `s`. + fn text_search(projection: Vec) -> SqlPlan { + SqlPlan::TextSearch(crate::types::TextSearchPlan { collection: "docs".into(), - query: crate::fts_types::FtsQuery::Plain { - text: "rust".into(), - fuzzy: true, + shape: crate::types::TextSearchShape::Match { + field: None, + query: crate::fts_types::FtsQuery::Plain { + text: "rust".into(), + fuzzy: true, + }, + mode: nodedb_types::text_search::QueryMode::And, + top_k: Some(10), }, - top_k: 10, filters: Vec::new(), - score_alias: Some("s".into()), - projection: vec![Projection::Computed { - expr: call("bm25_score"), + scores: vec![crate::types::TextScoreColumn { + field: None, + query: "rust".into(), + mode: nodedb_types::text_search::QueryMode::And, + fuzzy: true, alias: "s".into(), }], - }; + projection, + }) + } + + #[test] + fn a_search_plan_projection_serves_its_score() { + let plan = text_search(vec![Projection::Computed { + expr: call("bm25_score"), + alias: "s".into(), + }]); assert_eq!(check(&plan), Ok(())); } #[test] - fn a_subquery_tail_over_a_search_plan_is_not_checked() { - let search = SqlPlan::TextSearch { - collection: "docs".into(), - query: crate::fts_types::FtsQuery::Plain { - text: "rust".into(), - fuzzy: true, - }, - top_k: 10, - filters: Vec::new(), - score_alias: Some("s".into()), - projection: Vec::new(), + fn a_match_in_a_text_search_filter_is_refused() { + let SqlPlan::TextSearch(mut search) = text_search(Vec::new()) else { + unreachable!("text_search builds a TextSearch plan"); }; + search.filters.push(Filter { + expr: FilterExpr::Expr(call("text_match")), + }); + assert!(matches!( + check(&SqlPlan::TextSearch(search)), + Err(SqlError::SearchFunctionOutsideSearch { .. }) + )); + } + + #[test] + fn a_subquery_tail_over_a_search_plan_is_not_checked() { + let search = text_search(Vec::new()); let plan = SqlPlan::Subquery { input: Box::new(search), filters: Vec::new(), diff --git a/nodedb-sql/src/planner/select/entry.rs b/nodedb-sql/src/planner/select/entry.rs deleted file mode 100644 index 8e00878c6..000000000 --- a/nodedb-sql/src/planner/select/entry.rs +++ /dev/null @@ -1,847 +0,0 @@ -// SPDX-License-Identifier: Apache-2.0 - -//! Top-level query entry: CTE handling and UNION dispatch. ORDER BY and -//! search-trigger detection live in `order_by.rs`; LIMIT / OFFSET application -//! lives in `limit.rs`. - -use nodedb_types::DatabaseId; -use sqlparser::ast::{Query, SetExpr}; - -use super::cte_catalog::CteCatalog; -use super::limit::apply_limit; -use super::order_by::{apply_order_by, try_hybrid_from_projection}; -use super::query_tail::QueryTail; -use super::select_stmt::{has_aggregation, plan_select}; -use crate::error::{Result, SqlError}; -use crate::functions::registry::FunctionRegistry; -use crate::reserved::check_ast_identifier; -use crate::resolver::derived::{infer_subquery_relation, rename_output_columns}; -use crate::temporal::TemporalScope; -use crate::types::{Projection, SqlExpr, *}; - -/// Returns `true` when every projection item is either: -/// - a plain column reference to the surrogate/PK column (`id` or `document_id`), or -/// - a `vector_distance(...)` function call (any alias). -/// -/// Anything else — a payload field, `*`, or an unrecognised expression — returns `false`. -fn is_pure_vector_projection(projection: &[Projection]) -> bool { - if projection.is_empty() { - return false; - } - for item in projection { - match item { - Projection::Column(name) => { - let lower = name.to_ascii_lowercase(); - if lower != "id" && lower != "document_id" { - return false; - } - } - Projection::Computed { expr, .. } => { - // Accept any of the three vector distance function names. - let SqlExpr::Function { name, .. } = expr else { - return false; - }; - if !name.eq_ignore_ascii_case("vector_distance") - && !name.eq_ignore_ascii_case("vector_cosine_distance") - && !name.eq_ignore_ascii_case("vector_neg_inner_product") - { - return false; - } - } - // A Control-Plane-computed item is evaluated over the fetched - // payload, so the payload must be fetched. - Projection::CpComputed { .. } | Projection::Star | Projection::QualifiedStar(_) => { - return false; - } - } - } - true -} - -/// Plan a SELECT query that produces the statement's result rows. -/// -/// Only this SELECT list may hold a Control-Plane-computed item (a sequence -/// accessor over a relation): the Control Plane evaluates the statement's -/// output rows and nothing deeper. -pub fn plan_statement_query( - query: &Query, - catalog: &dyn SqlCatalog, - functions: &FunctionRegistry, - temporal: TemporalScope, -) -> Result { - plan_query_at(query, catalog, functions, temporal, true) -} - -/// Plan a nested SELECT query: a subquery, CTE body, UNION branch, derived -/// table, INSERT source, or MERGE source. Its SELECT list refuses a -/// sequence accessor over a relation. -pub fn plan_query( - query: &Query, - catalog: &dyn SqlCatalog, - functions: &FunctionRegistry, - temporal: TemporalScope, -) -> Result { - plan_query_at(query, catalog, functions, temporal, false) -} - -/// Plan a SELECT query. `statement_output` says whether its rows are the -/// statement's result; a WITH clause passes it to the outer query only. -fn plan_query_at( - query: &Query, - catalog: &dyn SqlCatalog, - functions: &FunctionRegistry, - temporal: TemporalScope, - statement_output: bool, -) -> Result { - // Handle CTEs (WITH clause). - if let Some(with) = &query.with - && with.recursive - { - return crate::planner::cte::plan_recursive_cte(query, catalog, functions, temporal); - } - // Non-recursive CTEs: plan each CTE subquery and the outer query. - if let Some(with) = &query.with - && !with.cte_tables.is_empty() - { - let inner_query = Query { - with: None, - body: query.body.clone(), - order_by: query.order_by.clone(), - limit_clause: query.limit_clause.clone(), - fetch: query.fetch.clone(), - locks: query.locks.clone(), - for_clause: query.for_clause.clone(), - settings: query.settings.clone(), - format_clause: query.format_clause.clone(), - pipe_operators: query.pipe_operators.clone(), - }; - - // Plan each CTE subquery and infer the relation it exposes. - let mut definitions = Vec::new(); - let mut relations = Vec::new(); - for cte in &with.cte_tables { - let name = check_ast_identifier(&cte.alias.name)?; - let declared: Vec = cte - .alias - .columns - .iter() - .map(|column| check_ast_identifier(&column.name)) - .collect::>()?; - let cte_plan = plan_query(&cte.query, catalog, functions, temporal)?; - let info = infer_subquery_relation( - catalog, - &name, - &cte.query, - Some(&cte_plan), - functions, - temporal, - )?; - definitions.push((name.clone(), cte_plan)); - relations.push((name, rename_output_columns(info, &declared))); - } - - // Build CTE-aware catalog so the outer query can reference CTE names. - let cte_catalog = CteCatalog { - inner: catalog, - relations, - }; - let outer = plan_query_at( - &inner_query, - &cte_catalog, - functions, - temporal, - statement_output, - )?; - - return Ok(SqlPlan::Cte { - definitions, - outer: Box::new(outer), - }); - } - - // Handle UNION. - match &*query.body { - SetExpr::Select(select) => { - // ORDER BY / LIMIT belong to the query, not to its SELECT body, - // but the scan planner needs them to pick an access path that can - // honour them — so they travel down with the SELECT. - let tail = QueryTail { - order_by: query.order_by.as_ref(), - limit_clause: &query.limit_clause, - fetch: query.fetch.as_ref(), - }; - let planned = plan_select( - select, - catalog, - functions, - temporal, - &tail, - statement_output, - )?; - let scope = planned.scope; - let mut plan = planned.plan; - // Snapshot the projection before ORDER BY transforms the plan, - // in case `apply_order_by` converts a Scan into VectorSearch. - let pre_order_by_projection: Option> = match &plan { - SqlPlan::Scan { projection, .. } => Some(projection.clone()), - _ => None, - }; - let pre_order_by_collection: Option = match &plan { - SqlPlan::Scan { collection, .. } => Some(collection.clone()), - _ => None, - }; - if let Some(order_by) = &query.order_by { - plan = apply_order_by(&plan, order_by, functions, &select.projection, &scope)?; - } - // Fall back to a SELECT-projection scan for hybrid-search and - // text-search triggers. The `SELECT id, rrf_score(...) AS score - // FROM c WHERE ... LIMIT N` shape has no ORDER BY, so - // `apply_order_by` cannot fire. The same applies to - // `SELECT id, bm25_score(field, term) FROM c ORDER BY id` where - // ORDER BY does not contain a search trigger. - // - // Also fires when the plan is already `TextSearch` (set by the - // WHERE `text_match(...)` path) and the SELECT list additionally - // contains `bm25_score(...)` — in that case we attach the - // `score_alias` so the executor knows to inject the score column. - // - // `apply_order_by` may have wrapped a search plan in a - // post-processing tail to carry the sort, so the upgrade inspects - // the body and is re-wrapped in place — otherwise the score column - // the SELECT list asked for would never be attached. - let upgrade = { - let leaf = match &plan { - SqlPlan::Subquery { input, .. } => input.as_ref(), - other => other, - }; - if matches!(leaf, SqlPlan::Scan { .. } | SqlPlan::TextSearch { .. }) { - try_hybrid_from_projection(leaf, &select.projection, functions)? - } else { - None - } - }; - if let Some(upgraded_leaf) = upgrade { - plan = match plan { - SqlPlan::Subquery { - filters, - projection, - window_functions, - sort_keys, - offset, - distinct, - limit, - .. - } => SqlPlan::Subquery { - input: Box::new(upgraded_leaf), - filters, - projection, - window_functions, - sort_keys, - offset, - distinct, - limit, - }, - _ => upgraded_leaf, - }; - } - // After ORDER BY: if we now have a VectorSearch, check whether - // the collection is vector-primary and the projection is - // payload-free. If so, set `skip_payload_fetch`. - if let SqlPlan::VectorSearch { - ref collection, - ref mut skip_payload_fetch, - ref mut filters, - ref mut payload_filters, - .. - } = plan - { - let info = catalog - .get_collection(DatabaseId::DEFAULT, collection) - .ok() - .flatten(); - let is_vector_primary = info - .as_ref() - .map(|c| c.primary == nodedb_types::PrimaryEngine::Vector) - .unwrap_or(false); - if is_vector_primary { - if let Some(ref proj) = pre_order_by_projection - && pre_order_by_collection.as_deref() == Some(collection.as_str()) - { - *skip_payload_fetch = is_pure_vector_projection(proj); - } - if let Some(vp) = info.as_ref().and_then(|c| c.vector_primary.as_ref()) { - let mut peeled: Vec = Vec::new(); - let is_indexed = |name: &str| { - vp.payload_indexes - .iter() - .any(|(p, _)| p.eq_ignore_ascii_case(name)) - }; - filters.retain(|f| match &f.expr { - FilterExpr::Comparison { - field, - op: CompareOp::Eq, - value, - } if is_indexed(field) => { - peeled.push(SqlPayloadAtom::Eq(field.clone(), value.clone())); - false - } - FilterExpr::InList { field, values } if is_indexed(field) => { - peeled.push(SqlPayloadAtom::In(field.clone(), values.clone())); - false - } - FilterExpr::Between { field, low, high } if is_indexed(field) => { - peeled.push(SqlPayloadAtom::Range { - field: field.clone(), - low: Some(low.clone()), - low_inclusive: true, - high: Some(high.clone()), - high_inclusive: true, - }); - false - } - FilterExpr::Comparison { field, op, value } - if matches!( - op, - CompareOp::Lt | CompareOp::Le | CompareOp::Gt | CompareOp::Ge - ) && is_indexed(field) => - { - let inclusive = matches!(op, CompareOp::Le | CompareOp::Ge); - let upper = matches!(op, CompareOp::Lt | CompareOp::Le); - peeled.push(SqlPayloadAtom::Range { - field: field.clone(), - low: if upper { None } else { Some(value.clone()) }, - low_inclusive: !upper && inclusive, - high: if upper { Some(value.clone()) } else { None }, - high_inclusive: upper && inclusive, - }); - false - } - FilterExpr::Expr(SqlExpr::BinaryOp { - left, - op: BinaryOp::Eq, - right, - }) => match (&**left, &**right) { - (SqlExpr::Column { name, .. }, SqlExpr::Literal(v)) - if is_indexed(name) => - { - peeled.push(SqlPayloadAtom::Eq(name.clone(), v.clone())); - false - } - (SqlExpr::Literal(v), SqlExpr::Column { name, .. }) - if is_indexed(name) => - { - peeled.push(SqlPayloadAtom::Eq(name.clone(), v.clone())); - false - } - _ => true, - }, - FilterExpr::Expr(SqlExpr::InList { - expr, - list, - negated: false, - }) => match &**expr { - SqlExpr::Column { name, .. } if is_indexed(name) => { - let mut lits = Vec::with_capacity(list.len()); - let all_lit = list.iter().all(|e| { - if let SqlExpr::Literal(v) = e { - lits.push(v.clone()); - true - } else { - false - } - }); - if all_lit { - peeled.push(SqlPayloadAtom::In(name.clone(), lits)); - false - } else { - true - } - } - _ => true, - }, - FilterExpr::Expr(SqlExpr::Between { - expr, - low, - high, - negated: false, - }) => match (&**expr, &**low, &**high) { - ( - SqlExpr::Column { name, .. }, - SqlExpr::Literal(lo), - SqlExpr::Literal(hi), - ) if is_indexed(name) => { - peeled.push(SqlPayloadAtom::Range { - field: name.clone(), - low: Some(lo.clone()), - low_inclusive: true, - high: Some(hi.clone()), - high_inclusive: true, - }); - false - } - _ => true, - }, - FilterExpr::Expr(SqlExpr::BinaryOp { left, op, right }) - if matches!( - op, - BinaryOp::Lt | BinaryOp::Le | BinaryOp::Gt | BinaryOp::Ge - ) => - { - match (&**left, &**right) { - (SqlExpr::Column { name, .. }, SqlExpr::Literal(v)) - if is_indexed(name) => - { - let inclusive = matches!(op, BinaryOp::Le | BinaryOp::Ge); - let upper = matches!(op, BinaryOp::Lt | BinaryOp::Le); - peeled.push(SqlPayloadAtom::Range { - field: name.clone(), - low: if upper { None } else { Some(v.clone()) }, - low_inclusive: !upper && inclusive, - high: if upper { Some(v.clone()) } else { None }, - high_inclusive: upper && inclusive, - }); - false - } - _ => true, - } - } - _ => true, - }); - *payload_filters = peeled; - } - } - } - let plan = apply_limit(plan, &tail)?; - // ORDER BY and LIMIT sit on the aggregate by now, so the wrap - // only restates the output columns around it. - crate::planner::aggregate_cp_wrap::wrap_aggregate_cp_items( - plan, - &select.projection, - has_aggregation(select, functions), - functions, - &scope, - ) - } - SetExpr::SetOperation { - op, - left, - right, - set_quantifier, - } => crate::planner::union::plan_set_operation( - op, - left, - right, - set_quantifier, - catalog, - functions, - temporal, - ), - _ => Err(SqlError::Unsupported { - detail: format!("query body type: {}", query.body), - }), - } -} - -/// Unit tests for SELECT query planning. -#[cfg(test)] -mod tests { - use super::*; - use crate::parser::preprocess::pipeline::preprocess; - use crate::parser::statement::parse_sql; - use sqlparser::ast::Statement; - - struct TestCatalog; - - impl SqlCatalog for TestCatalog { - fn get_collection( - &self, - _: nodedb_types::DatabaseId, - name: &str, - ) -> std::result::Result, SqlCatalogError> { - let info = match name { - "products" => Some(CollectionInfo { - name: "products".into(), - engine: EngineType::DocumentSchemaless, - columns: Vec::new(), - primary_key: Some("id".into()), - has_auto_tier: false, - indexes: Vec::new(), - bitemporal: false, - primary: nodedb_types::PrimaryEngine::Document, - vector_primary: None, - partition_strategy: nodedb_types::PartitionStrategy::CollectionHomed, - open_schema: CollectionInfo::open_schema_for(EngineType::DocumentSchemaless), - }), - "users" => Some(CollectionInfo { - name: "users".into(), - engine: EngineType::DocumentSchemaless, - columns: Vec::new(), - primary_key: Some("id".into()), - has_auto_tier: false, - indexes: Vec::new(), - bitemporal: false, - primary: nodedb_types::PrimaryEngine::Document, - vector_primary: None, - partition_strategy: nodedb_types::PartitionStrategy::CollectionHomed, - open_schema: CollectionInfo::open_schema_for(EngineType::DocumentSchemaless), - }), - "orders" => Some(CollectionInfo { - name: "orders".into(), - engine: EngineType::DocumentSchemaless, - columns: Vec::new(), - primary_key: Some("id".into()), - has_auto_tier: false, - indexes: Vec::new(), - bitemporal: false, - primary: nodedb_types::PrimaryEngine::Document, - vector_primary: None, - partition_strategy: nodedb_types::PartitionStrategy::CollectionHomed, - open_schema: CollectionInfo::open_schema_for(EngineType::DocumentSchemaless), - }), - "docs" => Some(CollectionInfo { - name: "docs".into(), - engine: EngineType::DocumentSchemaless, - columns: Vec::new(), - primary_key: Some("id".into()), - has_auto_tier: false, - indexes: Vec::new(), - bitemporal: false, - primary: nodedb_types::PrimaryEngine::Document, - vector_primary: None, - partition_strategy: nodedb_types::PartitionStrategy::CollectionHomed, - open_schema: CollectionInfo::open_schema_for(EngineType::DocumentSchemaless), - }), - "tags" => Some(CollectionInfo { - name: "tags".into(), - engine: EngineType::DocumentSchemaless, - columns: Vec::new(), - primary_key: Some("id".into()), - has_auto_tier: false, - indexes: Vec::new(), - bitemporal: false, - primary: nodedb_types::PrimaryEngine::Document, - vector_primary: None, - partition_strategy: nodedb_types::PartitionStrategy::CollectionHomed, - open_schema: CollectionInfo::open_schema_for(EngineType::DocumentSchemaless), - }), - "user_prefs" => Some(CollectionInfo { - name: "user_prefs".into(), - engine: EngineType::KeyValue, - columns: Vec::new(), - primary_key: Some("key".into()), - has_auto_tier: false, - indexes: Vec::new(), - bitemporal: false, - primary: nodedb_types::PrimaryEngine::Document, - vector_primary: None, - partition_strategy: nodedb_types::PartitionStrategy::CollectionHomed, - open_schema: CollectionInfo::open_schema_for(EngineType::KeyValue), - }), - "embeddings" => Some(CollectionInfo { - name: "embeddings".into(), - engine: EngineType::DocumentSchemaless, - columns: Vec::new(), - primary_key: Some("id".into()), - has_auto_tier: false, - indexes: Vec::new(), - bitemporal: false, - primary: nodedb_types::PrimaryEngine::Document, - vector_primary: None, - partition_strategy: nodedb_types::PartitionStrategy::CollectionHomed, - open_schema: CollectionInfo::open_schema_for(EngineType::DocumentSchemaless), - }), - _ => None, - }; - Ok(info) - } - - fn lookup_array(&self, name: &str) -> Option { - if name == "genome" { - Some(crate::types::ArrayCatalogView { - name: "genome".into(), - dims: vec![ - crate::types_array::ArrayDimAst { - name: "chrom".into(), - dtype: crate::types_array::ArrayDimType::Int64, - lo: crate::types_array::ArrayDomainBound::Int64(1), - hi: crate::types_array::ArrayDomainBound::Int64(23), - }, - crate::types_array::ArrayDimAst { - name: "pos".into(), - dtype: crate::types_array::ArrayDimType::Int64, - lo: crate::types_array::ArrayDomainBound::Int64(0), - hi: crate::types_array::ArrayDomainBound::Int64(1_000_000), - }, - ], - attrs: vec![crate::types_array::ArrayAttrAst { - name: "qual".into(), - dtype: crate::types_array::ArrayAttrType::Float64, - nullable: true, - }], - tile_extents: vec![1, 1_000_000], - }) - } else { - None - } - } - } - - fn plan_select_sql(sql: &str) -> SqlPlan { - // Run preprocessor so operator rewrites (`<->`, `<=>`, `<#>`) are applied - // before sqlparser sees the SQL. - let (preprocessed_sql, temporal) = match preprocess(sql).unwrap() { - Some(p) => (p.sql, p.temporal), - None => (sql.to_string(), crate::TemporalScope::default()), - }; - let statements = parse_sql(&preprocessed_sql).unwrap(); - let Statement::Query(query) = &statements[0] else { - panic!("expected query statement"); - }; - plan_query(query, &TestCatalog, &FunctionRegistry::new(), temporal).unwrap() - } - - #[test] - fn aggregate_subquery_join_filters_input_before_aggregation() { - let plan = plan_select_sql( - "SELECT AVG(price) FROM products WHERE category IN (SELECT DISTINCT category FROM products WHERE qty > 100)", - ); - - let SqlPlan::Aggregate { input, .. } = plan else { - panic!("expected aggregate plan"); - }; - - let SqlPlan::Join { - left, - join_type, - on, - .. - } = *input - else { - panic!("expected semi-join below aggregate"); - }; - - assert_eq!(join_type, JoinType::Semi); - assert_eq!(on, vec![("category".into(), "category".into())]); - assert!(matches!(*left, SqlPlan::Scan { .. })); - } - - #[test] - fn scalar_subquery_defers_projection_until_after_join_filter() { - let plan = plan_select_sql( - "SELECT user_id FROM orders WHERE amount > (SELECT AVG(amount) FROM orders)", - ); - - let SqlPlan::Join { - left, - projection, - filters, - .. - } = plan - else { - panic!("expected join plan"); - }; - - let SqlPlan::Scan { - projection: scan_projection, - .. - } = *left - else { - panic!("expected scan on join left"); - }; - - assert!(scan_projection.is_empty(), "scan projected too early"); - assert_eq!(projection.len(), 1); - match &projection[0] { - Projection::Column(name) => assert_eq!(name, "user_id"), - other => panic!("expected user_id projection, got {other:?}"), - } - assert!( - !filters.is_empty(), - "scalar comparison should stay post-join" - ); - } - - #[test] - fn chained_join_preserves_qualified_on_keys() { - let plan = plan_select_sql( - "SELECT d.name, t.tag, p.theme \ - FROM docs d \ - LEFT JOIN tags t ON d.id = t.doc_id \ - INNER JOIN user_prefs p ON d.id = p.key", - ); - - let SqlPlan::Join { left, on, .. } = plan else { - panic!("expected outer join plan"); - }; - assert_eq!(on, vec![("d.id".into(), "p.key".into())]); - - let SqlPlan::Join { on: inner_on, .. } = *left else { - panic!("expected nested left join"); - }; - assert_eq!(inner_on, vec![("d.id".into(), "t.doc_id".into())]); - } - - #[test] - fn order_by_vector_distance_with_array_join_fuses_into_vector_search() { - let plan = plan_select_sql( - "SELECT v.id FROM embeddings v \ - JOIN ARRAY_SLICE('genome', '{chrom: [1, 1], pos: [0, 50000]}') AS s \ - ON v.id = s.qual \ - ORDER BY vector_distance(v.embedding, [1.0, 0.0, 0.0]) \ - LIMIT 10", - ); - - let SqlPlan::VectorSearch { - collection, - top_k, - array_prefilter, - .. - } = plan - else { - panic!("expected fused VectorSearch plan"); - }; - assert_eq!(collection, "embeddings"); - assert_eq!(top_k, 10); - let prefilter = array_prefilter.expect("array_prefilter must be set on fused plan"); - assert_eq!(prefilter.array_name, "genome"); - assert_eq!(prefilter.slice.dim_ranges.len(), 2); - } - - #[test] - fn vector_distance_two_args_produces_default_ann_options() { - let plan = plan_select_sql( - "SELECT id FROM embeddings ORDER BY vector_distance(embedding, [1.0, 0.0]) LIMIT 5", - ); - let SqlPlan::VectorSearch { ann_options, .. } = plan else { - panic!("expected VectorSearch plan"); - }; - assert_eq!(ann_options, VectorAnnOptions::default()); - } - - #[test] - fn order_by_sparse_score_desc_routes_to_sparse_search() { - // `ORDER BY sparse_score(field, '{dim: weight, ...}') DESC LIMIT k` must - // route to `SqlPlan::SparseSearch` exactly as `vector_distance(...)` routes - // to `SqlPlan::VectorSearch`. The query literal is parsed into sorted - // `(dimension, weight)` entries and `top_k` tracks the LIMIT. - let plan = plan_select_sql( - "SELECT id FROM embeddings \ - ORDER BY sparse_score(terms, '{3: 1.0, 7: 0.5}') DESC LIMIT 5", - ); - let SqlPlan::SparseSearch { - collection, - field, - query_entries, - top_k, - .. - } = plan - else { - panic!("expected SparseSearch plan"); - }; - assert_eq!(collection, "embeddings"); - assert_eq!(field, "terms"); - assert_eq!(top_k, 5); - assert_eq!(query_entries, vec![(3, 1.0), (7, 0.5)]); - } - - #[test] - fn vector_distance_named_args_parses_ann_options() { - let plan = plan_select_sql( - "SELECT id FROM embeddings ORDER BY vector_distance(embedding, [1.0, 0.0], quantization => 'rabitq', oversample => 3) LIMIT 5", - ); - let SqlPlan::VectorSearch { - ann_options, - ef_search, - top_k, - .. - } = plan - else { - panic!("expected VectorSearch plan"); - }; - assert_eq!(ann_options.quantization, Some(VectorQuantization::RaBitQ)); - assert_eq!(ann_options.oversample, Some(3)); - // ef_search falls back to top_k * 2 (no ef_search_override supplied). - assert_eq!(ef_search, top_k * 2); - } - - #[test] - fn vector_distance_ef_search_override_applied() { - let plan = plan_select_sql( - "SELECT id FROM embeddings ORDER BY vector_distance(embedding, [1.0], ef_search => 150) LIMIT 5", - ); - let SqlPlan::VectorSearch { ef_search, .. } = plan else { - panic!("expected VectorSearch plan"); - }; - assert_eq!(ef_search, 150); - } - - #[test] - fn arrow_distance_operator_yields_l2_metric() { - // The <-> operator rewrites to vector_distance(...) via the preprocessor. - // Use the function form here since sqlparser handles bracket-array syntax. - let plan = plan_select_sql( - "SELECT id FROM embeddings ORDER BY vector_distance(embedding, [1.0, 0.0, 0.0]) LIMIT 5", - ); - let SqlPlan::VectorSearch { metric, .. } = plan else { - panic!("expected VectorSearch plan"); - }; - assert_eq!(metric, DistanceMetric::L2); - } - - #[test] - fn cosine_distance_operator_yields_cosine_metric() { - // The <=> operator rewrites to vector_cosine_distance(...). - let plan = plan_select_sql( - "SELECT id FROM embeddings ORDER BY vector_cosine_distance(embedding, [1.0, 0.0, 0.0]) LIMIT 5", - ); - let SqlPlan::VectorSearch { metric, .. } = plan else { - panic!("expected VectorSearch plan"); - }; - assert_eq!(metric, DistanceMetric::Cosine); - } - - #[test] - fn neg_inner_product_operator_yields_inner_product_metric() { - // The <#> operator rewrites to vector_neg_inner_product(...). - let plan = plan_select_sql( - "SELECT id FROM embeddings ORDER BY vector_neg_inner_product(embedding, [1.0, 0.0, 0.0]) LIMIT 5", - ); - let SqlPlan::VectorSearch { metric, .. } = plan else { - panic!("expected VectorSearch plan"); - }; - assert_eq!(metric, DistanceMetric::InnerProduct); - } - - #[test] - fn vector_distance_function_yields_l2_metric() { - let plan = plan_select_sql( - "SELECT id FROM embeddings ORDER BY vector_distance(embedding, [1.0, 0.0]) LIMIT 5", - ); - let SqlPlan::VectorSearch { metric, .. } = plan else { - panic!("expected VectorSearch plan"); - }; - assert_eq!(metric, DistanceMetric::L2); - } - - #[test] - fn vector_cosine_distance_function_yields_cosine_metric() { - let plan = plan_select_sql( - "SELECT id FROM embeddings ORDER BY vector_cosine_distance(embedding, [1.0, 0.0]) LIMIT 5", - ); - let SqlPlan::VectorSearch { metric, .. } = plan else { - panic!("expected VectorSearch plan"); - }; - assert_eq!(metric, DistanceMetric::Cosine); - } - - #[test] - fn vector_neg_inner_product_function_yields_inner_product_metric() { - let plan = plan_select_sql( - "SELECT id FROM embeddings ORDER BY vector_neg_inner_product(embedding, [1.0, 0.0]) LIMIT 5", - ); - let SqlPlan::VectorSearch { metric, .. } = plan else { - panic!("expected VectorSearch plan"); - }; - assert_eq!(metric, DistanceMetric::InnerProduct); - } -} diff --git a/nodedb-sql/src/planner/select/entry/fixtures.rs b/nodedb-sql/src/planner/select/entry/fixtures.rs new file mode 100644 index 000000000..d3d6cd336 --- /dev/null +++ b/nodedb-sql/src/planner/select/entry/fixtures.rs @@ -0,0 +1,200 @@ +// SPDX-License-Identifier: Apache-2.0 + +//! Shared query-planner test catalog and SQL preparation. + +use super::query::plan_query; +use crate::functions::registry::FunctionRegistry; +use crate::parser::preprocess::pipeline::preprocess; +use crate::parser::statement::parse_sql; +use crate::types::*; +use sqlparser::ast::Statement; + +struct TestCatalog; + +impl SqlCatalog for TestCatalog { + fn get_collection( + &self, + _: nodedb_types::DatabaseId, + name: &str, + ) -> std::result::Result, SqlCatalogError> { + let info = match name { + "products" => Some(CollectionInfo { + name: "products".into(), + engine: EngineType::DocumentSchemaless, + columns: Vec::new(), + primary_key: Some("id".into()), + has_auto_tier: false, + indexes: Vec::new(), + bitemporal: false, + primary: nodedb_types::PrimaryEngine::Document, + vector_primary: None, + partition_strategy: nodedb_types::PartitionStrategy::CollectionHomed, + open_schema: CollectionInfo::open_schema_for(EngineType::DocumentSchemaless), + }), + "users" => Some(CollectionInfo { + name: "users".into(), + engine: EngineType::DocumentSchemaless, + columns: Vec::new(), + primary_key: Some("id".into()), + has_auto_tier: false, + indexes: Vec::new(), + bitemporal: false, + primary: nodedb_types::PrimaryEngine::Document, + vector_primary: None, + partition_strategy: nodedb_types::PartitionStrategy::CollectionHomed, + open_schema: CollectionInfo::open_schema_for(EngineType::DocumentSchemaless), + }), + "orders" => Some(CollectionInfo { + name: "orders".into(), + engine: EngineType::DocumentSchemaless, + columns: Vec::new(), + primary_key: Some("id".into()), + has_auto_tier: false, + indexes: Vec::new(), + bitemporal: false, + primary: nodedb_types::PrimaryEngine::Document, + vector_primary: None, + partition_strategy: nodedb_types::PartitionStrategy::CollectionHomed, + open_schema: CollectionInfo::open_schema_for(EngineType::DocumentSchemaless), + }), + "docs" => Some(CollectionInfo { + name: "docs".into(), + engine: EngineType::DocumentSchemaless, + columns: Vec::new(), + primary_key: Some("id".into()), + has_auto_tier: false, + indexes: Vec::new(), + bitemporal: false, + primary: nodedb_types::PrimaryEngine::Document, + vector_primary: None, + partition_strategy: nodedb_types::PartitionStrategy::CollectionHomed, + open_schema: CollectionInfo::open_schema_for(EngineType::DocumentSchemaless), + }), + "tags" => Some(CollectionInfo { + name: "tags".into(), + engine: EngineType::DocumentSchemaless, + columns: Vec::new(), + primary_key: Some("id".into()), + has_auto_tier: false, + indexes: Vec::new(), + bitemporal: false, + primary: nodedb_types::PrimaryEngine::Document, + vector_primary: None, + partition_strategy: nodedb_types::PartitionStrategy::CollectionHomed, + open_schema: CollectionInfo::open_schema_for(EngineType::DocumentSchemaless), + }), + "user_prefs" => Some(CollectionInfo { + name: "user_prefs".into(), + engine: EngineType::KeyValue, + columns: Vec::new(), + primary_key: Some("key".into()), + has_auto_tier: false, + indexes: Vec::new(), + bitemporal: false, + primary: nodedb_types::PrimaryEngine::Document, + vector_primary: None, + partition_strategy: nodedb_types::PartitionStrategy::CollectionHomed, + open_schema: CollectionInfo::open_schema_for(EngineType::KeyValue), + }), + "embeddings" => Some(CollectionInfo { + name: "embeddings".into(), + engine: EngineType::DocumentSchemaless, + columns: Vec::new(), + primary_key: Some("id".into()), + has_auto_tier: false, + indexes: Vec::new(), + bitemporal: false, + primary: nodedb_types::PrimaryEngine::Document, + vector_primary: None, + partition_strategy: nodedb_types::PartitionStrategy::CollectionHomed, + open_schema: CollectionInfo::open_schema_for(EngineType::DocumentSchemaless), + }), + "articles" => Some(CollectionInfo { + name: "articles".into(), + engine: EngineType::DocumentStrict, + columns: vec![ + strict_column("id", SqlDataType::String, "TEXT", true), + strict_column("title", SqlDataType::String, "TEXT", false), + strict_column("views", SqlDataType::Int64, "INT", false), + ], + primary_key: Some("id".into()), + has_auto_tier: false, + indexes: Vec::new(), + bitemporal: false, + primary: nodedb_types::PrimaryEngine::Document, + vector_primary: None, + partition_strategy: nodedb_types::PartitionStrategy::CollectionHomed, + open_schema: CollectionInfo::open_schema_for(EngineType::DocumentStrict), + }), + _ => None, + }; + Ok(info) + } + + fn lookup_array(&self, name: &str) -> Option { + if name == "genome" { + Some(crate::types::ArrayCatalogView { + name: "genome".into(), + dims: vec![ + crate::types_array::ArrayDimAst { + name: "chrom".into(), + dtype: crate::types_array::ArrayDimType::Int64, + lo: crate::types_array::ArrayDomainBound::Int64(1), + hi: crate::types_array::ArrayDomainBound::Int64(23), + }, + crate::types_array::ArrayDimAst { + name: "pos".into(), + dtype: crate::types_array::ArrayDimType::Int64, + lo: crate::types_array::ArrayDomainBound::Int64(0), + hi: crate::types_array::ArrayDomainBound::Int64(1_000_000), + }, + ], + attrs: vec![crate::types_array::ArrayAttrAst { + name: "qual".into(), + dtype: crate::types_array::ArrayAttrType::Float64, + nullable: true, + }], + tile_extents: vec![1, 1_000_000], + }) + } else { + None + } + } +} + +fn strict_column( + name: &str, + data_type: SqlDataType, + raw: &str, + is_primary_key: bool, +) -> ColumnInfo { + ColumnInfo { + name: name.into(), + data_type, + nullable: !is_primary_key, + is_primary_key, + default: None, + raw_type: Some(raw.into()), + int_width: None, + float_width: None, + } +} + +pub(super) fn plan_select_sql(sql: &str) -> SqlPlan { + try_plan_select_sql(sql).unwrap() +} + +/// Plan `sql` against the test catalog, keeping the planner's error. +pub(super) fn try_plan_select_sql(sql: &str) -> crate::error::Result { + // Run preprocessor so operator rewrites (`<->`, `<=>`, `<#>`) are applied + // before sqlparser sees the SQL. + let (preprocessed_sql, temporal) = match preprocess(sql).unwrap() { + Some(p) => (p.sql, p.temporal), + None => (sql.to_string(), crate::TemporalScope::default()), + }; + let statements = parse_sql(&preprocessed_sql).unwrap(); + let Statement::Query(query) = &statements[0] else { + panic!("expected query statement"); + }; + plan_query(query, &TestCatalog, &FunctionRegistry::new(), temporal) +} diff --git a/nodedb-sql/src/planner/select/entry/mod.rs b/nodedb-sql/src/planner/select/entry/mod.rs new file mode 100644 index 000000000..a19ceec16 --- /dev/null +++ b/nodedb-sql/src/planner/select/entry/mod.rs @@ -0,0 +1,12 @@ +// SPDX-License-Identifier: Apache-2.0 + +//! SELECT query routing and search lowering. + +#[cfg(test)] +mod fixtures; +mod payload; +mod pk_prefilter; +mod query; +mod search; + +pub use query::{plan_query, plan_statement_query}; diff --git a/nodedb-sql/src/planner/select/entry/payload.rs b/nodedb-sql/src/planner/select/entry/payload.rs new file mode 100644 index 000000000..c2c975c61 --- /dev/null +++ b/nodedb-sql/src/planner/select/entry/payload.rs @@ -0,0 +1,250 @@ +// SPDX-License-Identifier: Apache-2.0 + +//! Vector-primary payload projection and indexed-filter extraction. + +use crate::error::Result; +use crate::types::*; +use nodedb_types::DatabaseId; + +/// Returns `true` when every projection item is either: +/// - a plain column reference to the surrogate/PK column (`id` or `document_id`), or +/// - a `vector_distance(...)` function call (any alias). +/// +/// Anything else — a payload field, `*`, or an unrecognised expression — returns `false`. +fn is_pure_vector_projection(projection: &[Projection]) -> bool { + if projection.is_empty() { + return false; + } + for item in projection { + match item { + Projection::Column(name) => { + if !name.eq_ignore_ascii_case("id") && !name.eq_ignore_ascii_case("document_id") { + return false; + } + } + Projection::Computed { expr, .. } => { + // Accept any of the three vector distance function names. + let SqlExpr::Function { name, .. } = expr else { + return false; + }; + if !name.eq_ignore_ascii_case("vector_distance") + && !name.eq_ignore_ascii_case("vector_cosine_distance") + && !name.eq_ignore_ascii_case("vector_neg_inner_product") + { + return false; + } + } + // A Control-Plane-computed item is evaluated over the fetched + // payload, so the payload must be fetched. + Projection::CpComputed { .. } | Projection::Star | Projection::QualifiedStar(_) => { + return false; + } + } + } + true +} + +pub(super) fn apply_vector_payload( + plan: &mut SqlPlan, + catalog: &dyn SqlCatalog, + pre_order_by_projection: Option<&[Projection]>, + pre_order_by_collection: Option<&str>, +) -> Result<()> { + // After ORDER BY: if we now have a VectorSearch, check whether + // the collection is vector-primary and the projection is + // payload-free. If so, set `skip_payload_fetch`. + if let SqlPlan::VectorSearch { + collection, + skip_payload_fetch, + filters, + payload_filters, + .. + } = plan + { + let info = catalog.get_collection(DatabaseId::DEFAULT, collection)?; + let is_vector_primary = info + .as_ref() + .map(|c| c.primary == nodedb_types::PrimaryEngine::Vector) + .unwrap_or(false); + if is_vector_primary { + if let Some(proj) = pre_order_by_projection + && pre_order_by_collection == Some(collection.as_str()) + { + *skip_payload_fetch = is_pure_vector_projection(proj); + } + if let Some(vp) = info.as_ref().and_then(|c| c.vector_primary.as_ref()) { + let mut peeled: Vec = Vec::new(); + let is_indexed = |name: &str| { + vp.payload_indexes + .iter() + .any(|(p, _)| p.eq_ignore_ascii_case(name)) + }; + filters.retain(|f| match &f.expr { + FilterExpr::Comparison { + field, + op: CompareOp::Eq, + value, + } if is_indexed(field) => { + peeled.push(SqlPayloadAtom::Eq(field.clone(), value.clone())); + false + } + FilterExpr::InList { field, values } if is_indexed(field) => { + peeled.push(SqlPayloadAtom::In(field.clone(), values.clone())); + false + } + FilterExpr::Between { field, low, high } if is_indexed(field) => { + peeled.push(SqlPayloadAtom::Range { + field: field.clone(), + low: Some(low.clone()), + low_inclusive: true, + high: Some(high.clone()), + high_inclusive: true, + }); + false + } + FilterExpr::Comparison { field, op, value } + if matches!( + op, + CompareOp::Lt | CompareOp::Le | CompareOp::Gt | CompareOp::Ge + ) && is_indexed(field) => + { + let inclusive = matches!(op, CompareOp::Le | CompareOp::Ge); + let upper = matches!(op, CompareOp::Lt | CompareOp::Le); + peeled.push(SqlPayloadAtom::Range { + field: field.clone(), + low: if upper { None } else { Some(value.clone()) }, + low_inclusive: !upper && inclusive, + high: if upper { Some(value.clone()) } else { None }, + high_inclusive: upper && inclusive, + }); + false + } + FilterExpr::Expr(SqlExpr::BinaryOp { + left, + op: BinaryOp::Eq, + right, + }) => match (&**left, &**right) { + (SqlExpr::Column { name, .. }, SqlExpr::Literal(v)) if is_indexed(name) => { + peeled.push(SqlPayloadAtom::Eq(name.clone(), v.clone())); + false + } + (SqlExpr::Literal(v), SqlExpr::Column { name, .. }) if is_indexed(name) => { + peeled.push(SqlPayloadAtom::Eq(name.clone(), v.clone())); + false + } + _ => true, + }, + FilterExpr::Expr(SqlExpr::InList { + expr, + list, + negated: false, + }) => match &**expr { + SqlExpr::Column { name, .. } if is_indexed(name) => { + let mut lits = Vec::with_capacity(list.len()); + let all_lit = list.iter().all(|e| { + if let SqlExpr::Literal(v) = e { + lits.push(v.clone()); + true + } else { + false + } + }); + if all_lit { + peeled.push(SqlPayloadAtom::In(name.clone(), lits)); + false + } else { + true + } + } + _ => true, + }, + FilterExpr::Expr(SqlExpr::Between { + expr, + low, + high, + negated: false, + }) => match (&**expr, &**low, &**high) { + ( + SqlExpr::Column { name, .. }, + SqlExpr::Literal(lo), + SqlExpr::Literal(hi), + ) if is_indexed(name) => { + peeled.push(SqlPayloadAtom::Range { + field: name.clone(), + low: Some(lo.clone()), + low_inclusive: true, + high: Some(hi.clone()), + high_inclusive: true, + }); + false + } + _ => true, + }, + FilterExpr::Expr(SqlExpr::BinaryOp { left, op, right }) + if matches!( + op, + BinaryOp::Lt | BinaryOp::Le | BinaryOp::Gt | BinaryOp::Ge + ) => + { + match (&**left, &**right) { + (SqlExpr::Column { name, .. }, SqlExpr::Literal(v)) + if is_indexed(name) => + { + let inclusive = matches!(op, BinaryOp::Le | BinaryOp::Ge); + let upper = matches!(op, BinaryOp::Lt | BinaryOp::Le); + peeled.push(SqlPayloadAtom::Range { + field: name.clone(), + low: if upper { None } else { Some(v.clone()) }, + low_inclusive: !upper && inclusive, + high: if upper { Some(v.clone()) } else { None }, + high_inclusive: upper && inclusive, + }); + false + } + _ => true, + } + } + _ => true, + }); + *payload_filters = peeled; + } + } + } + Ok(()) +} + +#[cfg(test)] +mod tests { + use super::super::fixtures::plan_select_sql; + use super::apply_vector_payload; + use crate::catalog::{SqlCatalog, SqlCatalogError}; + use crate::types::CollectionInfo; + + struct ChangingCatalog; + + impl SqlCatalog for ChangingCatalog { + fn get_collection( + &self, + _: nodedb_types::DatabaseId, + _: &str, + ) -> Result, SqlCatalogError> { + Err(SqlCatalogError::RetryableSchemaChanged { + descriptor: "collection embeddings".into(), + }) + } + } + + #[test] + fn vector_payload_lookup_preserves_catalog_error() { + let mut plan = plan_select_sql( + "SELECT id FROM embeddings ORDER BY vector_distance(embedding, [1.0, 0.0]) LIMIT 5", + ); + let error = apply_vector_payload(&mut plan, &ChangingCatalog, None, None).unwrap_err(); + assert_eq!( + error, + crate::SqlError::from(SqlCatalogError::RetryableSchemaChanged { + descriptor: "collection embeddings".into(), + }) + ); + } +} diff --git a/nodedb-sql/src/planner/select/entry/query.rs b/nodedb-sql/src/planner/select/entry/query.rs new file mode 100644 index 000000000..ae00f8bc0 --- /dev/null +++ b/nodedb-sql/src/planner/select/entry/query.rs @@ -0,0 +1,232 @@ +// SPDX-License-Identifier: Apache-2.0 + +//! Top-level query entry: CTE handling and UNION dispatch. ORDER BY and +//! search-trigger detection live in `order_by.rs`; LIMIT / OFFSET application +//! lives in `limit.rs`. + +use sqlparser::ast::{Query, SetExpr}; + +use crate::error::{Result, SqlError}; +use crate::functions::registry::FunctionRegistry; +use crate::planner::select::cte_catalog::CteCatalog; +use crate::reserved::check_ast_identifier; +use crate::resolver::derived::{infer_subquery_relation, rename_output_columns}; +use crate::temporal::TemporalScope; +use crate::types::{CtePlan, SqlCatalog, SqlPlan}; + +/// Plan a SELECT query that produces the statement's result rows. +/// +/// Only this SELECT list may hold a Control-Plane-computed item (a sequence +/// accessor over a relation): the Control Plane evaluates the statement's +/// output rows and nothing deeper. +pub fn plan_statement_query( + query: &Query, + catalog: &dyn SqlCatalog, + functions: &FunctionRegistry, + temporal: TemporalScope, +) -> Result { + plan_query_at(query, catalog, functions, temporal, true) +} + +/// Plan a nested SELECT query: a subquery, CTE body, UNION branch, derived +/// table, INSERT source, or MERGE source. Its SELECT list refuses a +/// sequence accessor over a relation. +pub fn plan_query( + query: &Query, + catalog: &dyn SqlCatalog, + functions: &FunctionRegistry, + temporal: TemporalScope, +) -> Result { + plan_query_at(query, catalog, functions, temporal, false) +} + +/// Plan a SELECT query. `statement_output` says whether its rows are the +/// statement's result; a WITH clause passes it to the outer query only. +fn plan_query_at( + query: &Query, + catalog: &dyn SqlCatalog, + functions: &FunctionRegistry, + temporal: TemporalScope, + statement_output: bool, +) -> Result { + // Handle CTEs (WITH clause). + if let Some(with) = &query.with + && with.recursive + { + return crate::planner::cte::plan_recursive_cte(query, catalog, functions, temporal); + } + // Non-recursive CTEs: plan each CTE subquery and the outer query. + if let Some(with) = &query.with + && !with.cte_tables.is_empty() + { + let inner_query = Query { + with: None, + body: query.body.clone(), + order_by: query.order_by.clone(), + limit_clause: query.limit_clause.clone(), + fetch: query.fetch.clone(), + locks: query.locks.clone(), + for_clause: query.for_clause.clone(), + settings: query.settings.clone(), + format_clause: query.format_clause.clone(), + pipe_operators: query.pipe_operators.clone(), + }; + + // Plan each CTE subquery and infer the relation it exposes. + let mut definitions = Vec::new(); + let mut relations = Vec::new(); + for cte in &with.cte_tables { + let name = check_ast_identifier(&cte.alias.name)?; + let declared: Vec = cte + .alias + .columns + .iter() + .map(|column| check_ast_identifier(&column.name)) + .collect::>()?; + let cte_plan = plan_query(&cte.query, catalog, functions, temporal)?; + let info = infer_subquery_relation( + catalog, + &name, + &cte.query, + Some(&cte_plan), + functions, + temporal, + )?; + definitions.push((name.clone(), cte_plan)); + relations.push((name, rename_output_columns(info, &declared))); + } + + // Build CTE-aware catalog so the outer query can reference CTE names. + let cte_catalog = CteCatalog { + inner: catalog, + relations, + }; + let outer = plan_query_at( + &inner_query, + &cte_catalog, + functions, + temporal, + statement_output, + )?; + + return Ok(SqlPlan::Cte(CtePlan { + definitions, + outer: Box::new(outer), + })); + } + + // Handle UNION. + match &*query.body { + SetExpr::Select(select) => super::search::plan_select_query( + query, + select, + catalog, + functions, + temporal, + statement_output, + ), + SetExpr::SetOperation { + op, + left, + right, + set_quantifier, + } => crate::planner::union::plan_set_operation( + op, + left, + right, + set_quantifier, + catalog, + functions, + temporal, + ), + _ => Err(SqlError::Unsupported { + detail: format!("query body type: {}", query.body), + }), + } +} + +#[cfg(test)] +mod tests { + use super::super::fixtures::plan_select_sql; + use crate::types::*; + #[test] + fn aggregate_subquery_join_filters_input_before_aggregation() { + let plan = plan_select_sql( + "SELECT AVG(price) FROM products WHERE category IN (SELECT DISTINCT category FROM products WHERE qty > 100)", + ); + + let SqlPlan::Aggregate { input, .. } = plan else { + panic!("expected aggregate plan"); + }; + + let SqlPlan::Join { + left, + join_type, + on, + .. + } = *input + else { + panic!("expected semi-join below aggregate"); + }; + + assert_eq!(join_type, JoinType::Semi); + assert_eq!(on, vec![("category".into(), "category".into())]); + assert!(matches!(*left, SqlPlan::Scan { .. })); + } + + #[test] + fn scalar_subquery_defers_projection_until_after_join_filter() { + let plan = plan_select_sql( + "SELECT user_id FROM orders WHERE amount > (SELECT AVG(amount) FROM orders)", + ); + + let SqlPlan::Join { + left, + projection, + filters, + .. + } = plan + else { + panic!("expected join plan"); + }; + + let SqlPlan::Scan { + projection: scan_projection, + .. + } = *left + else { + panic!("expected scan on join left"); + }; + + assert!(scan_projection.is_empty(), "scan projected too early"); + assert_eq!(projection.len(), 1); + match &projection[0] { + Projection::Column(name) => assert_eq!(name, "user_id"), + other => panic!("expected user_id projection, got {other:?}"), + } + assert!( + !filters.is_empty(), + "scalar comparison should stay post-join" + ); + } + + #[test] + fn chained_join_preserves_qualified_on_keys() { + let plan = plan_select_sql( + "SELECT d.name, t.tag, p.theme \ + FROM docs d \ + LEFT JOIN tags t ON d.id = t.doc_id \ + INNER JOIN user_prefs p ON d.id = p.key", + ); + + let SqlPlan::Join { left, on, .. } = plan else { + panic!("expected outer join plan"); + }; + assert_eq!(on, vec![("d.id".into(), "p.key".into())]); + + let SqlPlan::Join { on: inner_on, .. } = *left else { + panic!("expected nested left join"); + }; + assert_eq!(inner_on, vec![("d.id".into(), "t.doc_id".into())]); + } +} diff --git a/nodedb-sql/src/planner/select/entry/search.rs b/nodedb-sql/src/planner/select/entry/search.rs new file mode 100644 index 000000000..d3ba9d08b --- /dev/null +++ b/nodedb-sql/src/planner/select/entry/search.rs @@ -0,0 +1,797 @@ +// SPDX-License-Identifier: Apache-2.0 + +//! SELECT-body search triggers and post-processing order. + +use super::payload::apply_vector_payload; +use super::pk_prefilter::apply_vector_pk_prefilter; +use crate::error::Result; +use crate::functions::registry::FunctionRegistry; +use crate::planner::aggregate_cp_wrap::wrap_aggregate_cp_items; +use crate::planner::select::limit::apply_limit; +use crate::planner::select::order_by::{apply_order_by, try_hybrid_from_projection}; +use crate::planner::select::query_tail::QueryTail; +use crate::planner::select::select_stmt::{has_aggregation, plan_select}; +use crate::temporal::TemporalScope; +use crate::types::{Projection, SqlCatalog, SqlPlan}; +use sqlparser::ast::{Query, Select}; + +pub(super) fn plan_select_query( + query: &Query, + select: &Select, + catalog: &dyn SqlCatalog, + functions: &FunctionRegistry, + temporal: TemporalScope, + statement_output: bool, +) -> Result { + // ORDER BY / LIMIT belong to the query, not to its SELECT body, + // but the scan planner needs them to pick an access path that can + // honour them — so they travel down with the SELECT. + let tail = QueryTail { + order_by: query.order_by.as_ref(), + limit_clause: &query.limit_clause, + fetch: query.fetch.as_ref(), + }; + let planned = plan_select( + select, + catalog, + functions, + temporal, + &tail, + statement_output, + )?; + let scope = planned.scope; + let mut plan = planned.plan; + // Snapshot the projection before ORDER BY transforms the plan, + // in case `apply_order_by` converts a Scan into VectorSearch. + let pre_order_by_projection: Option> = match &plan { + SqlPlan::Scan { projection, .. } => Some(projection.clone()), + _ => None, + }; + let pre_order_by_collection: Option = match &plan { + SqlPlan::Scan { collection, .. } => Some(collection.clone()), + _ => None, + }; + if let Some(order_by) = &query.order_by { + plan = apply_order_by(&plan, order_by, functions, &select.projection, &scope)?; + } + // Fall back to a SELECT-projection scan for hybrid-search and + // text-search triggers. The `SELECT id, rrf_score(...) AS score + // FROM c WHERE ... LIMIT N` shape has no ORDER BY, so + // `apply_order_by` cannot fire. The same applies to + // `SELECT id, bm25_score(field, term) FROM c ORDER BY id` where + // ORDER BY does not contain a search trigger. + // + // Also fires when the plan is already `TextSearch` (set by the + // WHERE `text_match(...)` path) and the SELECT list additionally + // contains `bm25_score(...)` — in that case each call joins the + // plan as a score column the executor injects. + // + // `apply_order_by` may have wrapped a search plan in a + // post-processing tail to carry the sort, so the upgrade inspects + // the body and is re-wrapped in place — otherwise the score column + // the SELECT list asked for would never be attached. + let upgrade = { + let leaf = match &plan { + SqlPlan::Subquery { input, .. } => input.as_ref(), + other => other, + }; + match scope.single_table() { + Some(table) if matches!(leaf, SqlPlan::Scan { .. } | SqlPlan::TextSearch(_)) => { + try_hybrid_from_projection(leaf, &select.projection, functions, table)? + } + _ => None, + } + }; + if let Some(upgraded_leaf) = upgrade { + plan = match plan { + SqlPlan::Subquery { + filters, + projection, + window_functions, + sort_keys, + offset, + distinct, + limit, + .. + } => SqlPlan::Subquery { + input: Box::new(upgraded_leaf), + filters, + projection, + window_functions, + sort_keys, + offset, + distinct, + limit, + }, + _ => upgraded_leaf, + }; + } + apply_vector_pk_prefilter(&mut plan, catalog)?; + apply_vector_payload( + &mut plan, + catalog, + pre_order_by_projection.as_deref(), + pre_order_by_collection.as_deref(), + )?; + let plan = apply_limit(plan, &tail)?; + // ORDER BY and LIMIT sit on the aggregate by now, so the wrap + // only restates the output columns around it. + wrap_aggregate_cp_items( + plan, + &select.projection, + has_aggregation(select, functions), + functions, + &scope, + ) +} + +#[cfg(test)] +mod tests { + use super::super::fixtures::{plan_select_sql, try_plan_select_sql}; + use crate::error::SqlError; + use crate::types::*; + use nodedb_types::text_search::{QueryMode, TextColumnFault, TextSearchParams}; + #[test] + fn order_by_vector_distance_with_array_join_fuses_into_vector_search() { + let plan = plan_select_sql( + "SELECT v.id FROM embeddings v \ + JOIN ARRAY_SLICE('genome', '{chrom: [1, 1], pos: [0, 50000]}') AS s \ + ON v.id = s.qual \ + ORDER BY vector_distance(v.embedding, [1.0, 0.0, 0.0]) \ + LIMIT 10", + ); + + let SqlPlan::VectorSearch { + collection, + top_k, + array_prefilter, + .. + } = plan + else { + panic!("expected fused VectorSearch plan"); + }; + assert_eq!(collection, "embeddings"); + assert_eq!(top_k, 10); + let prefilter = array_prefilter.expect("array_prefilter must be set on fused plan"); + assert_eq!(prefilter.array_name, "genome"); + assert_eq!(prefilter.slice.dim_ranges.len(), 2); + } + + #[test] + fn vector_distance_two_args_produces_default_ann_options() { + let plan = plan_select_sql( + "SELECT id FROM embeddings ORDER BY vector_distance(embedding, [1.0, 0.0]) LIMIT 5", + ); + let SqlPlan::VectorSearch { ann_options, .. } = plan else { + panic!("expected VectorSearch plan"); + }; + assert_eq!(ann_options, VectorAnnOptions::default()); + } + + #[test] + fn order_by_sparse_score_desc_routes_to_sparse_search() { + // `ORDER BY sparse_score(field, '{dim: weight, ...}') DESC LIMIT k` must + // route to `SqlPlan::SparseSearch` exactly as `vector_distance(...)` routes + // to `SqlPlan::VectorSearch`. The query literal is parsed into sorted + // `(dimension, weight)` entries and `top_k` tracks the LIMIT. + let plan = plan_select_sql( + "SELECT id FROM embeddings \ + ORDER BY sparse_score(terms, '{3: 1.0, 7: 0.5}') DESC LIMIT 5", + ); + let SqlPlan::SparseSearch { + collection, + field, + query_entries, + top_k, + .. + } = plan + else { + panic!("expected SparseSearch plan"); + }; + assert_eq!(collection, "embeddings"); + assert_eq!(field, "terms"); + assert_eq!(top_k, 5); + assert_eq!(query_entries, vec![(3, 1.0), (7, 0.5)]); + } + + #[test] + fn vector_distance_named_args_parses_ann_options() { + let plan = plan_select_sql( + "SELECT id FROM embeddings ORDER BY vector_distance(embedding, [1.0, 0.0], quantization => 'rabitq', oversample => 3) LIMIT 5", + ); + let SqlPlan::VectorSearch { + ann_options, + ef_search, + top_k, + .. + } = plan + else { + panic!("expected VectorSearch plan"); + }; + assert_eq!(ann_options.quantization, Some(VectorQuantization::RaBitQ)); + assert_eq!(ann_options.oversample, Some(3)); + // ef_search falls back to top_k * 2 (no ef_search_override supplied). + assert_eq!(ef_search, top_k * 2); + } + + #[test] + fn vector_distance_ef_search_override_applied() { + let plan = plan_select_sql( + "SELECT id FROM embeddings ORDER BY vector_distance(embedding, [1.0], ef_search => 150) LIMIT 5", + ); + let SqlPlan::VectorSearch { ef_search, .. } = plan else { + panic!("expected VectorSearch plan"); + }; + assert_eq!(ef_search, 150); + } + + #[test] + fn arrow_distance_operator_yields_l2_metric() { + // The <-> operator rewrites to vector_distance(...) via the preprocessor. + // Use the function form here since sqlparser handles bracket-array syntax. + let plan = plan_select_sql( + "SELECT id FROM embeddings ORDER BY vector_distance(embedding, [1.0, 0.0, 0.0]) LIMIT 5", + ); + let SqlPlan::VectorSearch { metric, .. } = plan else { + panic!("expected VectorSearch plan"); + }; + assert_eq!(metric, DistanceMetric::L2); + } + + #[test] + fn a_lone_column_argument_matches_no_vector_distance_signature() { + let err = try_plan_select_sql( + "SELECT id FROM embeddings ORDER BY vector_distance(embedding) LIMIT 5", + ) + .unwrap_err(); + assert_eq!( + err, + SqlError::UndefinedFunction { + name: "vector_distance".into() + } + ); + // The one-argument query-vector form still plans a search. + let plan = plan_select_sql( + "SELECT id FROM embeddings ORDER BY vector_distance(ARRAY[1.0, 0.0]) LIMIT 5", + ); + assert!(matches!(plan, SqlPlan::VectorSearch { .. })); + } + + #[test] + fn cosine_distance_operator_yields_cosine_metric() { + // The <=> operator rewrites to vector_cosine_distance(...). + let plan = plan_select_sql( + "SELECT id FROM embeddings ORDER BY vector_cosine_distance(embedding, [1.0, 0.0, 0.0]) LIMIT 5", + ); + let SqlPlan::VectorSearch { metric, .. } = plan else { + panic!("expected VectorSearch plan"); + }; + assert_eq!(metric, DistanceMetric::Cosine); + } + + #[test] + fn neg_inner_product_operator_yields_inner_product_metric() { + // The <#> operator rewrites to vector_neg_inner_product(...). + let plan = plan_select_sql( + "SELECT id FROM embeddings ORDER BY vector_neg_inner_product(embedding, [1.0, 0.0, 0.0]) LIMIT 5", + ); + let SqlPlan::VectorSearch { metric, .. } = plan else { + panic!("expected VectorSearch plan"); + }; + assert_eq!(metric, DistanceMetric::InnerProduct); + } + + #[test] + fn vector_distance_function_yields_l2_metric() { + let plan = plan_select_sql( + "SELECT id FROM embeddings ORDER BY vector_distance(embedding, [1.0, 0.0]) LIMIT 5", + ); + let SqlPlan::VectorSearch { metric, .. } = plan else { + panic!("expected VectorSearch plan"); + }; + assert_eq!(metric, DistanceMetric::L2); + } + + #[test] + fn vector_cosine_distance_function_yields_cosine_metric() { + let plan = plan_select_sql( + "SELECT id FROM embeddings ORDER BY vector_cosine_distance(embedding, [1.0, 0.0]) LIMIT 5", + ); + let SqlPlan::VectorSearch { metric, .. } = plan else { + panic!("expected VectorSearch plan"); + }; + assert_eq!(metric, DistanceMetric::Cosine); + } + + #[test] + fn vector_neg_inner_product_function_yields_inner_product_metric() { + let plan = plan_select_sql( + "SELECT id FROM embeddings ORDER BY vector_neg_inner_product(embedding, [1.0, 0.0]) LIMIT 5", + ); + let SqlPlan::VectorSearch { metric, .. } = plan else { + panic!("expected VectorSearch plan"); + }; + assert_eq!(metric, DistanceMetric::InnerProduct); + } + + /// The text plan of `plan`, under at most one post-processing tail. + fn text_plan(plan: &SqlPlan) -> &TextSearchPlan { + match plan { + SqlPlan::TextSearch(search) => search, + SqlPlan::Subquery { input, .. } => match input.as_ref() { + SqlPlan::TextSearch(search) => search, + other => panic!("expected TextSearch under the tail, got {other:?}"), + }, + other => panic!("expected TextSearch, got {other:?}"), + } + } + + /// The field and top-k of a `Match` shape. + fn match_shape(search: &TextSearchPlan) -> (Option<&str>, Option) { + match &search.shape { + TextSearchShape::Match { field, top_k, .. } => (field.as_deref(), *top_k), + TextSearchShape::ScoreScan => panic!("expected a Match shape"), + } + } + + #[test] + fn where_text_match_scopes_to_the_column_and_returns_every_match() { + let plan = plan_select_sql("SELECT id FROM docs WHERE text_match(title, 'rust')"); + let search = text_plan(&plan); + assert_eq!(match_shape(search), (Some("title"), None)); + assert!(search.scores.is_empty()); + assert!(search.filters.is_empty()); + } + + #[test] + fn where_text_match_star_reads_the_whole_document() { + let plan = plan_select_sql("SELECT id FROM docs WHERE text_match(*, 'rust')"); + assert_eq!(match_shape(text_plan(&plan)), (None, None)); + let plan = plan_select_sql("SELECT id FROM docs WHERE text_match(docs.*, 'rust')"); + assert_eq!(match_shape(text_plan(&plan)), (None, None)); + } + + #[test] + fn where_text_match_limit_is_the_top_k() { + let plan = plan_select_sql("SELECT id FROM docs WHERE text_match(body, 'rust') LIMIT 3"); + assert_eq!(match_shape(text_plan(&plan)), (Some("body"), Some(3))); + } + + #[test] + fn a_score_beside_a_match_keeps_the_match_shape() { + let plan = plan_select_sql( + "SELECT id, bm25_score(title, 'x') AS s FROM docs WHERE text_match(body, 'y')", + ); + let search = text_plan(&plan); + assert_eq!(match_shape(search), (Some("body"), None)); + assert_eq!( + search.scores, + vec![TextScoreColumn { + field: Some("title".into()), + query: "x".into(), + mode: QueryMode::Or, + fuzzy: false, + alias: "s".into(), + }] + ); + } + + #[test] + fn every_sibling_predicate_restricts_the_match() { + let plan = plan_select_sql( + "SELECT id FROM docs WHERE text_match(body, 'w') AND tag = 'a' AND n > 1 LIMIT 3", + ); + let search = text_plan(&plan); + assert_eq!(match_shape(search), (Some("body"), Some(3))); + assert_eq!(search.filters.len(), 2); + } + + #[test] + fn a_score_without_a_match_scans_the_filtered_rows_in_order() { + let plan = plan_select_sql( + "SELECT id, bm25_score(body, 'x') FROM docs WHERE tag = 'a' ORDER BY id LIMIT 2", + ); + let SqlPlan::Subquery { + sort_keys, limit, .. + } = &plan + else { + panic!("expected a post-processing tail, got {plan:?}"); + }; + assert_eq!(*limit, Some(2)); + assert_eq!(sort_keys.len(), 1); + let search = text_plan(&plan); + assert!(matches!(search.shape, TextSearchShape::ScoreScan)); + assert_eq!(search.filters.len(), 1); + assert_eq!(search.scores.len(), 1); + assert_eq!(search.scores[0].field.as_deref(), Some("body")); + } + + #[test] + fn order_by_score_sorts_by_the_score_column() { + let plan = + plan_select_sql("SELECT id FROM docs ORDER BY bm25_score(body, 'x') DESC LIMIT 5"); + let SqlPlan::Subquery { + sort_keys, limit, .. + } = &plan + else { + panic!("expected a post-processing tail, got {plan:?}"); + }; + assert_eq!(*limit, Some(5)); + let search = text_plan(&plan); + assert!(matches!(search.shape, TextSearchShape::ScoreScan)); + let alias = &search.scores[0].alias; + assert!(matches!( + &sort_keys[0].expr, + SqlExpr::Column { table: None, name } if name == alias + )); + assert!(!sort_keys[0].ascending); + // A null score takes the default NULL placement of a DESC key: first. + assert!(sort_keys[0].nulls_first); + } + + #[test] + fn order_by_score_honors_an_explicit_nulls_clause() { + let plan = plan_select_sql( + "SELECT id FROM docs ORDER BY bm25_score(body, 'x') DESC NULLS LAST LIMIT 5", + ); + let SqlPlan::Subquery { sort_keys, .. } = &plan else { + panic!("expected a post-processing tail, got {plan:?}"); + }; + assert!(!sort_keys[0].ascending); + assert!(!sort_keys[0].nulls_first); + + let plan = plan_select_sql("SELECT id FROM docs ORDER BY bm25_score(body, 'x') LIMIT 5"); + let SqlPlan::Subquery { sort_keys, .. } = &plan else { + panic!("expected a post-processing tail, got {plan:?}"); + }; + assert!(sort_keys[0].ascending); + assert!(!sort_keys[0].nulls_first); + } + + #[test] + fn order_by_score_alias_reuses_the_select_alias() { + let plan = plan_select_sql( + "SELECT id, bm25_score(body, 'x') AS score FROM docs \ + WHERE text_match(body, 'x') ORDER BY score DESC LIMIT 20", + ); + let search = text_plan(&plan); + assert_eq!(match_shape(search), (Some("body"), None)); + assert_eq!(search.scores.len(), 1); + assert_eq!(search.scores[0].alias, "score"); + } + + #[test] + fn order_by_qualifier_of_the_table_names_its_column() { + let plan = plan_select_sql("SELECT d.id FROM docs d ORDER BY bm25_score(d.body, 'x')"); + assert_eq!(text_plan(&plan).scores[0].field.as_deref(), Some("body")); + } + + #[test] + fn order_by_vector_qualifier_of_the_table_names_its_column() { + let plan = plan_select_sql( + "SELECT e.id FROM embeddings e ORDER BY vector_distance(e.embedding, [1.0, 0.0]) LIMIT 5", + ); + let SqlPlan::VectorSearch { field, .. } = plan else { + panic!("expected VectorSearch plan"); + }; + assert_eq!(field, "embedding"); + } + + #[test] + fn order_by_foreign_qualifier_is_an_unknown_table() { + let err = try_plan_select_sql("SELECT id FROM docs ORDER BY bm25_score(zz.body, 'x')") + .unwrap_err(); + assert_eq!(err, SqlError::UnknownTable { name: "zz".into() }); + } + + #[test] + fn a_literal_is_not_a_text_column() { + let err = + try_plan_select_sql("SELECT id FROM docs WHERE text_match('lit', 'x')").unwrap_err(); + assert!(matches!( + err, + SqlError::TextColumn { + fault: TextColumnFault::NotAColumn, + .. + } + )); + } + + #[test] + fn strict_columns_are_checked_for_text() { + let plan = plan_select_sql("SELECT id FROM articles WHERE text_match(title, 'x')"); + assert_eq!(match_shape(text_plan(&plan)), (Some("title"), None)); + + let err = try_plan_select_sql("SELECT id FROM articles WHERE text_match(views, 'x')") + .unwrap_err(); + assert!(matches!( + err, + SqlError::TextColumn { + fault: TextColumnFault::NotText { .. }, + .. + } + )); + + let err = try_plan_select_sql("SELECT id FROM articles WHERE text_match(ghost, 'x')") + .unwrap_err(); + assert!(matches!( + err, + SqlError::TextColumn { + ref collection, + fault: TextColumnFault::Undeclared, + .. + } if collection == "articles" + )); + } + + #[test] + fn a_text_match_with_one_argument_is_an_arity_error() { + let err = try_plan_select_sql("SELECT id FROM docs WHERE text_match(body)").unwrap_err(); + assert!(matches!(err, SqlError::Arity { .. })); + } + + #[test] + fn distinct_over_a_score_scan_is_refused() { + let err = + try_plan_select_sql("SELECT DISTINCT id, bm25_score(body, 'x') FROM docs").unwrap_err(); + assert!(matches!(err, SqlError::Unsupported { .. })); + } + + #[test] + fn aggregation_over_a_where_match_is_refused() { + let err = try_plan_select_sql("SELECT count(*) FROM docs WHERE text_match(body, 'x')") + .unwrap_err(); + assert!(matches!(err, SqlError::Unsupported { .. })); + } + + #[test] + fn hybrid_carries_both_columns_and_the_filters() { + let plan = plan_select_sql( + "SELECT id, rrf_score(vector_distance(emb, [1.0, 0.0]), bm25_score(title, 'rust')) \ + AS s FROM docs WHERE tag = 'a' LIMIT 5", + ); + let SqlPlan::HybridSearch(hybrid) = plan else { + panic!("expected HybridSearch, got {plan:?}"); + }; + assert_eq!(hybrid.vector_field, "emb"); + assert_eq!(hybrid.text_field.as_deref(), Some("title")); + assert_eq!(hybrid.query_text, "rust"); + assert_eq!(hybrid.filters.len(), 1); + assert_eq!(hybrid.top_k, 5); + } + + /// The mode and fuzzy flag of a `Match` shape. + fn match_options(search: &TextSearchPlan) -> (QueryMode, bool) { + match &search.shape { + TextSearchShape::Match { query, mode, .. } => (*mode, query.is_fuzzy()), + TextSearchShape::ScoreScan => panic!("expected a Match shape"), + } + } + + /// The `Unsupported` detail of a planning error. + fn unsupported_detail(sql: &str) -> String { + match try_plan_select_sql(sql).unwrap_err() { + SqlError::Unsupported { detail } => detail, + other => panic!("expected Unsupported, got {other:?}"), + } + } + + #[test] + fn text_match_without_options_runs_the_trait_default() { + let defaults = TextSearchParams::default(); + let plan = plan_select_sql("SELECT id FROM docs WHERE text_match(body, 'rust db')"); + assert_eq!( + match_options(text_plan(&plan)), + (defaults.mode, defaults.fuzzy) + ); + let plan = plan_select_sql("SELECT id FROM docs ORDER BY bm25_score(body, 'x') LIMIT 5"); + let score = &text_plan(&plan).scores[0]; + assert_eq!((score.mode, score.fuzzy), (defaults.mode, defaults.fuzzy)); + } + + #[test] + fn text_match_mode_option_reaches_the_plan() { + let plan = + plan_select_sql("SELECT id FROM docs WHERE text_match(body, 'rust db', mode => 'and')"); + assert_eq!(match_options(text_plan(&plan)), (QueryMode::And, false)); + let plan = + plan_select_sql("SELECT id FROM docs WHERE text_match(body, 'rust db', mode => 'or')"); + assert_eq!(match_options(text_plan(&plan)), (QueryMode::Or, false)); + } + + #[test] + fn text_match_fuzzy_option_reaches_the_plan() { + let plan = + plan_select_sql("SELECT id FROM docs WHERE text_match(body, 'databse', fuzzy => true)"); + assert_eq!(match_options(text_plan(&plan)), (QueryMode::Or, true)); + let plan = plan_select_sql( + "SELECT id FROM docs WHERE text_match(body, 'databse', fuzzy => false, mode => 'and')", + ); + assert_eq!(match_options(text_plan(&plan)), (QueryMode::And, false)); + } + + #[test] + fn bm25_score_options_reach_the_score_column() { + let plan = plan_select_sql( + "SELECT id FROM docs \ + ORDER BY bm25_score(body, 'x y', mode => 'and', fuzzy => true) DESC LIMIT 5", + ); + let score = &text_plan(&plan).scores[0]; + assert_eq!((score.mode, score.fuzzy), (QueryMode::And, true)); + let plan = plan_select_sql( + "SELECT id, bm25_score(title, 'x', fuzzy => true) AS s FROM docs \ + WHERE text_match(body, 'y', mode => 'and')", + ); + let search = text_plan(&plan); + assert_eq!(match_options(search), (QueryMode::And, false)); + assert_eq!( + (search.scores[0].mode, search.scores[0].fuzzy), + (QueryMode::Or, true) + ); + } + + #[test] + fn hybrid_text_leg_takes_the_bm25_options() { + let plan = plan_select_sql( + "SELECT id, rrf_score(vector_distance(emb, [1.0, 0.0]), \ + bm25_score(title, 'rust', mode => 'and', fuzzy => true)) AS s FROM docs LIMIT 5", + ); + let SqlPlan::HybridSearch(hybrid) = plan else { + panic!("expected HybridSearch, got {plan:?}"); + }; + assert_eq!((hybrid.mode, hybrid.fuzzy), (QueryMode::And, true)); + } + + #[test] + fn an_unknown_text_option_is_refused() { + let detail = + unsupported_detail("SELECT id FROM docs WHERE text_match(body, 'x', boost => 2)"); + assert!( + detail.contains("unknown text-search option 'boost'"), + "{detail}" + ); + let detail = unsupported_detail( + "SELECT id FROM docs ORDER BY bm25_score(body, 'x', slop => 1) LIMIT 5", + ); + assert!( + detail.contains("unknown text-search option 'slop'"), + "{detail}" + ); + } + + #[test] + fn an_equals_text_option_is_refused() { + let detail = + unsupported_detail("SELECT id FROM docs WHERE text_match(body, 'x', mode = 'or')"); + assert!(detail.contains("use '=>'"), "{detail}"); + } + + #[test] + fn a_positional_third_text_argument_is_refused() { + let detail = unsupported_detail("SELECT id FROM docs WHERE text_match(body, 'x', 'fuzzy')"); + assert!(detail.contains("third positional argument"), "{detail}"); + let detail = + unsupported_detail("SELECT id FROM docs WHERE text_match(body, 'x', { fuzzy: true })"); + assert!(detail.contains("third positional argument"), "{detail}"); + } + + #[test] + fn options_on_a_phrase_query_are_refused() { + let detail = unsupported_detail( + "SELECT id FROM docs WHERE text_match(body, '\"quick fox\"', fuzzy => true)", + ); + assert!(detail.contains("phrase"), "{detail}"); + let plan = plan_select_sql("SELECT id FROM docs WHERE text_match(body, '\"quick fox\"')"); + assert!(matches!( + &text_plan(&plan).shape, + TextSearchShape::Match { + query: crate::fts_types::FtsQuery::Phrase(_), + .. + } + )); + } + + #[test] + fn hybrid_without_a_vector_leg_is_an_error() { + let err = try_plan_select_sql( + "SELECT id, rrf_score(1, bm25_score(title, 'rust')) AS s FROM docs LIMIT 5", + ) + .unwrap_err(); + assert!(matches!(err, SqlError::InvalidFunction { .. })); + } + + const RRF: &str = "rrf_score(vector_distance(emb, [1.0, 0.0]), bm25_score(title, 'rust'))"; + + /// The hybrid plan of `plan`, under at most one post-processing tail. + fn hybrid_plan(plan: &SqlPlan) -> &HybridSearchPlan { + match plan { + SqlPlan::HybridSearch(hybrid) => hybrid, + SqlPlan::Subquery { input, .. } => match input.as_ref() { + SqlPlan::HybridSearch(hybrid) => hybrid, + other => panic!("expected HybridSearch under the tail, got {other:?}"), + }, + other => panic!("expected HybridSearch, got {other:?}"), + } + } + + /// The `InvalidFunction` detail of a planning error. + fn invalid_function_detail(sql: &str) -> String { + match try_plan_select_sql(sql).unwrap_err() { + SqlError::InvalidFunction { detail } => detail, + other => panic!("expected InvalidFunction for {sql}, got {other:?}"), + } + } + + #[test] + fn a_projected_hybrid_sorted_by_another_column_ranks_limit_plus_offset_rows() { + let plan = plan_select_sql(&format!( + "SELECT id, {RRF} AS s FROM docs ORDER BY id LIMIT 30 OFFSET 4" + )); + let SqlPlan::Subquery { + sort_keys, + limit, + offset, + .. + } = &plan + else { + panic!("expected a post-processing tail, got {plan:?}"); + }; + assert_eq!((*limit, *offset, sort_keys.len()), (Some(30), 4, 1)); + assert_eq!(hybrid_plan(&plan).top_k, 34); + } + + #[test] + fn a_hybrid_ranked_by_its_score_ranks_limit_plus_offset_rows() { + let plan = plan_select_sql(&format!( + "SELECT id, {RRF} AS s FROM docs ORDER BY s DESC LIMIT 5 OFFSET 2" + )); + assert_eq!(hybrid_plan(&plan).top_k, 7); + let plan = plan_select_sql(&format!( + "SELECT id, {RRF} AS s FROM docs ORDER BY s DESC LIMIT 25" + )); + assert_eq!(hybrid_plan(&plan).top_k, 25); + } + + #[test] + fn a_hybrid_with_no_limit_is_an_error() { + for sql in [ + format!("SELECT id, {RRF} AS s FROM docs"), + format!("SELECT id, {RRF} AS s FROM docs ORDER BY id"), + format!("SELECT id, {RRF} AS s FROM docs ORDER BY s DESC"), + format!("SELECT id FROM docs ORDER BY {RRF} DESC"), + ] { + let detail = invalid_function_detail(&sql); + assert!(detail.contains("LIMIT"), "{sql}: {detail}"); + } + } + + #[test] + fn a_malformed_rrf_constant_is_an_error() { + let legs = "vector_distance(emb, [1.0, 0.0]), bm25_score(title, 'rust')"; + for k in ["'x'", "0", "-5", "title"] { + let sql = format!("SELECT id, rrf_score({legs}, {k}) AS s FROM docs LIMIT 5"); + let detail = invalid_function_detail(&sql); + assert!(detail.contains("k1"), "{sql}: {detail}"); + let sql = format!("SELECT id, rrf_score({legs}, 60, {k}) AS s FROM docs LIMIT 5"); + let detail = invalid_function_detail(&sql); + assert!(detail.contains("k2"), "{sql}: {detail}"); + } + let plan = plan_select_sql(&format!( + "SELECT id, rrf_score({legs}, 30, 10) AS s FROM docs LIMIT 5" + )); + let hybrid = hybrid_plan(&plan); + assert!((hybrid.vector_weight - 0.25).abs() < f32::EPSILON); + } + + #[test] + fn a_text_leg_other_than_bm25_score_is_an_error() { + for leg in ["text_match(title, 'rust')", "upper(title)", "title"] { + let sql = format!( + "SELECT id, rrf_score(vector_distance(emb, [1.0, 0.0]), {leg}) AS s \ + FROM docs LIMIT 5" + ); + let detail = invalid_function_detail(&sql); + assert!(detail.contains("bm25_score"), "{sql}: {detail}"); + } + } +} diff --git a/nodedb-sql/src/planner/select/helpers.rs b/nodedb-sql/src/planner/select/helpers.rs index 54a0df2d4..7de98f5c9 100644 --- a/nodedb-sql/src/planner/select/helpers.rs +++ b/nodedb-sql/src/planner/select/helpers.rs @@ -22,7 +22,7 @@ pub(super) fn source_projection(plan: &SqlPlan) -> Vec { match plan { SqlPlan::Scan { projection, .. } | SqlPlan::Join { projection, .. } - | SqlPlan::TextSearch { projection, .. } => projection.clone(), + | SqlPlan::TextSearch(TextSearchPlan { projection, .. }) => projection.clone(), _ => Vec::new(), } } diff --git a/nodedb-sql/src/planner/select/limit.rs b/nodedb-sql/src/planner/select/limit.rs index a69ccd29c..6d2c68986 100644 --- a/nodedb-sql/src/planner/select/limit.rs +++ b/nodedb-sql/src/planner/select/limit.rs @@ -11,6 +11,10 @@ use super::post_process::post_process; use super::query_tail::QueryTail; use crate::error::{Result, SqlError}; use crate::types::SqlPlan; +use crate::types::{ + ArraySlicePlan, CtePlan, DocumentIndexLookupPlan, HybridSearchPlan, HybridSearchTriplePlan, + TextSearchPlan, TextSearchShape, TimeseriesScanPlan, +}; /// Default `ef_search` multiplier applied when LIMIT is the only signal /// available for sizing the HNSW beam (e.g. on a fused VectorSearch that @@ -32,11 +36,11 @@ pub(in crate::planner::select) fn apply_limit( // The LIMIT belongs to the query reading the CTE, not to the CTE body, so // it lands on the outer plan — without this a derived table like // `FROM (...) s LIMIT n` comes back unbounded. - if let SqlPlan::Cte { definitions, outer } = plan { - return Ok(SqlPlan::Cte { + if let SqlPlan::Cte(CtePlan { definitions, outer }) = plan { + return Ok(SqlPlan::Cte(CtePlan { definitions, outer: Box::new(apply_limit(*outer, tail)?), - }); + })); } // Only these three variants carry an OFFSET of their own. For the rest the @@ -75,11 +79,11 @@ pub(in crate::planner::select) fn apply_limit( // The index-lookup rewrite of a document scan. It carries the same // row bound as the scan it replaced — without this the converter // substitutes its own default and `LIMIT 1` returns 10,000 rows. - SqlPlan::DocumentIndexLookup { + SqlPlan::DocumentIndexLookup(DocumentIndexLookupPlan { ref mut limit, ref mut offset, .. - } => { + }) => { *limit = limit_val; *offset = offset_val; } @@ -93,9 +97,9 @@ pub(in crate::planner::select) fn apply_limit( *limit = limit_val; *offset = offset_val; } - SqlPlan::TimeseriesScan { + SqlPlan::TimeseriesScan(TimeseriesScanPlan { limit: ref mut l, .. - } => { + }) => { if let Some(lv) = limit_val { *l = lv; } @@ -107,9 +111,9 @@ pub(in crate::planner::select) fn apply_limit( *l = lv; } } - SqlPlan::ArraySlice { + SqlPlan::ArraySlice(ArraySlicePlan { limit: ref mut l, .. - } => { + }) => { if let Some(lv) = limit_val { // The slice bound is a `u32` (0 = unlimited); a LIMIT that // does not fit would wrap into a smaller bound. @@ -122,18 +126,34 @@ pub(in crate::planner::select) fn apply_limit( // `LIMIT N` *is* the top-k bound. `ef_search` is deliberately left // alone on the fused-search variants: a wider beam than the final N // costs distance computations, never correctness. - SqlPlan::TextSearch { - top_k: ref mut k, .. + // A text match ranks its hits, so `LIMIT N` is its top-k bound. With + // no LIMIT it returns every match. + SqlPlan::TextSearch(TextSearchPlan { + shape: TextSearchShape::Match { + top_k: ref mut k, .. + }, + .. + }) => { + *k = limit_val; } - | SqlPlan::MultiVectorSearch { - top_k: ref mut k, .. + // A score scan has no ranked cut: its rows are bounded by a tail. + SqlPlan::TextSearch(TextSearchPlan { + shape: TextSearchShape::ScoreScan, + .. + }) => { + if limit_val.is_some() { + return post_process(plan, Vec::new(), limit_val, 0); + } } - | SqlPlan::HybridSearch { + SqlPlan::MultiVectorSearch { top_k: ref mut k, .. } - | SqlPlan::HybridSearchTriple { + | SqlPlan::HybridSearch(HybridSearchPlan { top_k: ref mut k, .. - } => { + }) + | SqlPlan::HybridSearchTriple(HybridSearchTriplePlan { + top_k: ref mut k, .. + }) => { if let Some(lv) = limit_val { *k = lv; } @@ -200,6 +220,7 @@ mod tests { use super::*; use crate::temporal::TemporalScope; + use crate::types::LateralLoopPlan; use crate::types::query::{EngineType, JoinType}; fn minimal_scan() -> SqlPlan { @@ -274,7 +295,7 @@ mod tests { } fn lateral_loop_plan() -> SqlPlan { - SqlPlan::LateralLoop { + SqlPlan::LateralLoop(LateralLoopPlan { outer: Box::new(minimal_scan()), outer_alias: None, inner: Box::new(minimal_scan()), @@ -283,7 +304,7 @@ mod tests { projection: vec![], outer_row_cap: 10, left_join: false, - } + }) } #[test] diff --git a/nodedb-sql/src/planner/select/mod.rs b/nodedb-sql/src/planner/select/mod.rs index a27bff678..3de2a457a 100644 --- a/nodedb-sql/src/planner/select/mod.rs +++ b/nodedb-sql/src/planner/select/mod.rs @@ -17,6 +17,8 @@ mod order_by; mod post_process; mod query_tail; mod select_stmt; +mod text_call; +mod text_options; mod where_search; pub(crate) use cte_catalog::CteCatalog; diff --git a/nodedb-sql/src/planner/select/order_by/apply.rs b/nodedb-sql/src/planner/select/order_by/apply.rs index 6789b353c..f11f7122d 100644 --- a/nodedb-sql/src/planner/select/order_by/apply.rs +++ b/nodedb-sql/src/planner/select/order_by/apply.rs @@ -9,7 +9,7 @@ use sqlparser::ast; use super::aliases::{resolve_order_by_target, select_output_aliases}; -use super::triggers::try_extract_sort_search; +use super::triggers::{SortSearch, try_extract_sort_search}; use crate::error::Result; use crate::functions::registry::FunctionRegistry; use crate::planner::agg_bind::{BindName, bind_aggregate_calls}; @@ -51,10 +51,49 @@ pub(in crate::planner::select) fn apply_order_by( // same call under an alias, and propagate that alias. let first = &exprs[0]; let (resolved_expr, score_alias) = resolve_order_by_target(&first.expr, select_items); - if let Some(search_plan) = - try_extract_sort_search(resolved_expr, plan, functions, score_alias.as_deref())? - { - return Ok(search_plan); + match try_extract_sort_search( + resolved_expr, + plan, + functions, + score_alias.as_deref(), + scope.single_table(), + )? { + Some(SortSearch::Ranked(search_plan)) => return Ok(search_plan), + // The leading key sorts by the score column the plan carries. The + // rest convert as plain sort keys. + // + // A row the index does not hold scores `null`. Its key takes the + // NULL placement of every other ORDER BY key: last ascending, first + // descending, unless the clause names `NULLS FIRST` / `NULLS LAST`. + Some(SortSearch::Scored { + plan: scored, + alias, + }) => { + let sort_scope = scope.with_output_names(select_output_aliases(select_items)); + let mut keys = vec![SortKey { + expr: SqlExpr::Column { + table: None, + name: alias, + }, + ascending: first.options.asc.unwrap_or(true), + nulls_first: first + .options + .nulls_first + .unwrap_or(!first.options.asc.unwrap_or(true)), + }]; + for o in &exprs[1..] { + keys.push(SortKey { + expr: convert_expr(&o.expr, &ColumnScope::Relations(&sort_scope))?, + ascending: o.options.asc.unwrap_or(true), + nulls_first: o + .options + .nulls_first + .unwrap_or(!o.options.asc.unwrap_or(true)), + }); + } + return post_process(scored, keys, None, 0); + } + None => {} } // After GROUP BY, an ORDER BY term may name an aggregate — projected or @@ -177,7 +216,7 @@ pub(in crate::planner::select) fn apply_order_by( // rows in the order it finds them (memtable, then partitions), which // is not the order the client asked for. Dropping the sort here would // silently answer `ORDER BY ts DESC` with ascending rows. - SqlPlan::TimeseriesScan { + SqlPlan::TimeseriesScan(TimeseriesScanPlan { collection, time_range, bucket_interval_ms, @@ -190,7 +229,7 @@ pub(in crate::planner::select) fn apply_order_by( tiered, temporal, .. - } => Ok(SqlPlan::TimeseriesScan { + }) => Ok(SqlPlan::TimeseriesScan(TimeseriesScanPlan { collection: collection.clone(), time_range: *time_range, bucket_interval_ms: *bucket_interval_ms, @@ -203,12 +242,12 @@ pub(in crate::planner::select) fn apply_order_by( sort_keys, tiered: *tiered, temporal: *temporal, - }), + })), // Cte wraps an inner outer plan; push ORDER BY into that outer // so derived-table queries (`SELECT … FROM (…) AS t ORDER BY …`) // honour the sort. inline_cte downstream merges the outer Scan // with the inner subquery plan; the sort_keys ride along. - SqlPlan::Cte { definitions, outer } => Ok(SqlPlan::Cte { + SqlPlan::Cte(CtePlan { definitions, outer }) => Ok(SqlPlan::Cte(CtePlan { definitions: definitions.clone(), outer: Box::new(apply_order_by( outer, @@ -217,7 +256,7 @@ pub(in crate::planner::select) fn apply_order_by( select_items, scope, )?), - }), + })), // The clause is non-empty here (`exprs` was checked above), so these // keys were asked for. A variant with no slot to hold them must not // pass through unchanged — that answers the query in whatever order diff --git a/nodedb-sql/src/planner/select/order_by/graph_score.rs b/nodedb-sql/src/planner/select/order_by/graph_score.rs new file mode 100644 index 000000000..8fdd43a8e --- /dev/null +++ b/nodedb-sql/src/planner/select/order_by/graph_score.rs @@ -0,0 +1,218 @@ +// SPDX-License-Identifier: Apache-2.0 + +//! Arguments of the `graph_score(node_id_col, 'seed', depth => N, label => 'edge')` +//! leg of a three-source `rrf_score(...)`. +//! +//! Two positional arguments: the node-id column (the executor resolves +//! surrogates from the collection) and the seed node id, a string literal. +//! The options are a closed set, each named with `=>`: `depth` (a positive +//! integer, default 1) and `label` (a string literal, default every label). +//! A third positional argument, an unknown or repeated option, `=` in place +//! of `=>`, and a value of the wrong type are typed errors. + +use sqlparser::ast::{self, FunctionArgOperator}; + +use super::super::helpers::extract_string_literal; +use crate::error::{Result, SqlError}; + +/// The options `graph_score` accepts. +const OPTION_NAMES: &str = "depth, label"; + +/// The BFS spec of a `graph_score(...)` leg. +#[derive(Debug, Clone, PartialEq, Eq)] +pub(super) struct GraphScoreArgs { + pub seed_id: String, + pub depth: usize, + pub edge_label: Option, +} + +/// Parse the argument list of a `graph_score(...)` call. +pub(super) fn parse_graph_score_args(func: &ast::Function) -> Result { + let ast::FunctionArguments::List(list) = &func.args else { + return Err(arity_error()); + }; + let mut positional: Vec<&ast::Expr> = Vec::new(); + let mut depth: Option = None; + let mut edge_label: Option = None; + for arg in &list.args { + let (key, value, operator) = match arg { + ast::FunctionArg::Unnamed(ast::FunctionArgExpr::Expr(expr)) => { + if positional.len() >= 2 { + return Err(invalid(format!( + "graph_score() takes two positional arguments (node_id_col, 'seed'); \ + name the options: {OPTION_NAMES} (e.g. depth => 2), got {expr}" + ))); + } + positional.push(expr); + continue; + } + ast::FunctionArg::Unnamed(other) => { + return Err(invalid(format!( + "graph_score() expects a value expression, got {other}" + ))); + } + ast::FunctionArg::Named { + name, + arg, + operator, + } => (name.value.to_ascii_lowercase(), arg, operator), + ast::FunctionArg::ExprNamed { + name: ast::Expr::Identifier(ident), + arg, + operator, + } => (ident.value.to_ascii_lowercase(), arg, operator), + ast::FunctionArg::ExprNamed { name, .. } => { + return Err(invalid(format!( + "graph_score(): option name {name} is not an identifier; \ + name an option: {OPTION_NAMES}" + ))); + } + }; + if *operator != FunctionArgOperator::RightArrow { + return Err(invalid(format!( + "graph_score(): use '=>' not '{operator}' for option '{key}'" + ))); + } + let ast::FunctionArgExpr::Expr(value) = value else { + return Err(invalid(format!( + "graph_score(): option '{key}' expects a value, got {value}" + ))); + }; + match key.as_str() { + "depth" => { + if depth.is_some() { + return Err(repeated(&key)); + } + depth = Some(parse_depth(value)?); + } + "label" => { + if edge_label.is_some() { + return Err(repeated(&key)); + } + edge_label = Some(extract_string_literal(value).map_err(|_| { + invalid(format!( + "graph_score(): option 'label' expects a string literal, got {value}" + )) + })?); + } + _ => { + return Err(invalid(format!( + "graph_score(): unknown option '{key}'; the options are {OPTION_NAMES}" + ))); + } + } + } + let [_, seed] = positional.as_slice() else { + return Err(arity_error()); + }; + let seed_id = extract_string_literal(seed).map_err(|_| { + invalid(format!( + "graph_score(): the seed node id must be a string literal, got {seed}" + )) + })?; + Ok(GraphScoreArgs { + seed_id, + depth: depth.unwrap_or(1), + edge_label, + }) +} + +/// A `depth => N` value: a positive integer literal. +fn parse_depth(value: &ast::Expr) -> Result { + let error = || { + invalid(format!( + "graph_score(): option 'depth' expects a positive integer, got {value}" + )) + }; + let ast::Expr::Value(v) = value else { + return Err(error()); + }; + let ast::Value::Number(n, _) = &v.value else { + return Err(error()); + }; + match n.parse::() { + Ok(depth) if depth > 0 => Ok(depth), + _ => Err(error()), + } +} + +fn arity_error() -> SqlError { + SqlError::Arity { + detail: "graph_score() takes (node_id_col, 'seed', [depth => N, label => 'edge'])".into(), + } +} + +fn repeated(key: &str) -> SqlError { + invalid(format!("graph_score(): option '{key}' is given twice")) +} + +fn invalid(detail: String) -> SqlError { + SqlError::InvalidFunction { detail } +} + +#[cfg(test)] +mod tests { + use sqlparser::dialect::GenericDialect; + use sqlparser::parser::Parser; + + use super::*; + + fn parse(call: &str) -> Result { + let sql = format!("SELECT {call}"); + let statements = Parser::parse_sql(&GenericDialect {}, &sql).expect("parses"); + let ast::Statement::Query(query) = &statements[0] else { + panic!("expected a query"); + }; + let ast::SetExpr::Select(select) = query.body.as_ref() else { + panic!("expected a select"); + }; + let ast::SelectItem::UnnamedExpr(ast::Expr::Function(func)) = &select.projection[0] else { + panic!("expected a function call"); + }; + parse_graph_score_args(func) + } + + #[test] + fn named_options_reach_the_spec() { + assert_eq!( + parse("graph_score(id, 'n1', depth => 2, label => 'hop')").expect("valid"), + GraphScoreArgs { + seed_id: "n1".into(), + depth: 2, + edge_label: Some("hop".into()), + } + ); + assert_eq!( + parse("graph_score(id, 'n1')").expect("valid"), + GraphScoreArgs { + seed_id: "n1".into(), + depth: 1, + edge_label: None, + } + ); + } + + #[test] + fn malformed_arguments_are_typed_errors() { + for call in [ + "graph_score(id, 7)", + "graph_score(id, 'n1', depth => 0)", + "graph_score(id, 'n1', depth => 1.5)", + "graph_score(id, 'n1', depth => 'x')", + "graph_score(id, 'n1', label => 3)", + "graph_score(id, 'n1', hops => 2)", + "graph_score(id, 'n1', depth => 1, depth => 2)", + "graph_score(id, 'n1', 2)", + ] { + let err = parse(call).expect_err(call); + assert!( + matches!(err, SqlError::InvalidFunction { .. }), + "{call}: {err:?}" + ); + } + assert!(matches!( + parse("graph_score(id)").expect_err("one argument"), + SqlError::Arity { .. } + )); + } +} diff --git a/nodedb-sql/src/planner/select/order_by/hybrid.rs b/nodedb-sql/src/planner/select/order_by/hybrid.rs index 3b38641cb..fd8ecf008 100644 --- a/nodedb-sql/src/planner/select/order_by/hybrid.rs +++ b/nodedb-sql/src/planner/select/order_by/hybrid.rs @@ -16,59 +16,96 @@ //! should use for the RRF score column — without it, the executor falls back //! to the fixed internal name `rrf_score`. //! +//! The fused search returns its best `top_k` rows. `top_k` is the query's +//! `LIMIT + OFFSET`: the rows any tail above the search can return. A query +//! with no LIMIT is a typed error, since the fusion ranks a bounded list. +//! //! Validation: //! - Fewer than 2 source args: typed error. //! - Exactly 4 or more than 6 args where arg[2] is numeric: typed error //! (3 sources require arg[2] to be graph_score(...), not a k-constant). //! - 3 sources + 2 k constants: typed error (inconsistent arity). //! - 3 sources + 3 k constants: valid triple-source form. +//! - A leg that is not `vector_distance(...)`, `bm25_score(...)`, or +//! `graph_score(...)` in its position: typed error. +//! - A k-constant that is not a positive finite number: typed error. +use crate::types::{HybridSearchPlan, HybridSearchTriplePlan}; use sqlparser::ast; use super::super::helpers::{ - extract_float, extract_float_array, extract_func_args, extract_string_literal, - source_projection, + extract_float, extract_float_array, extract_func_args, source_projection, }; +use super::super::text_call::{TextCall, resolve_table_column, resolve_text_call}; +use super::aliases::function_call_name; +use super::graph_score::parse_graph_score_args; +use super::text_score::refuse_scan_clauses; use crate::error::{Result, SqlError}; -use crate::types::{Projection, SqlPlan}; +use crate::functions::registry::{FunctionRegistry, SearchTrigger}; +use crate::resolver::columns::ResolvedTable; +use crate::types::{Filter, Projection, SqlPlan}; + +/// The RRF constant of a leg that names none. +const DEFAULT_RRF_K: f64 = 60.0; + +/// What every hybrid plan takes from the plan it replaces. +struct HybridBase<'a> { + table: &'a ResolvedTable, + functions: &'a FunctionRegistry, + top_k: usize, + filters: Vec, + score_alias: Option<&'a str>, + projection: Vec, +} /// Build a `SqlPlan::HybridSearch` or `SqlPlan::HybridSearchTriple` from a -/// `rrf_score(...)` call depending on argument arity. +/// `rrf_score(...)` call depending on argument arity. `None` when `plan` is +/// no `Scan`: no other plan becomes a hybrid search. pub(super) fn plan_hybrid_from_sort( args: &[ast::Expr], - collection: &str, + table: &ResolvedTable, plan: &SqlPlan, score_alias: Option<&str>, + functions: &FunctionRegistry, ) -> Result> { if args.len() < 2 { return Err(no_args_rrf_score_error()); } - - let limit = match plan { - SqlPlan::Scan { limit, .. } => limit.unwrap_or(10), - _ => 10, + let SqlPlan::Scan { + limit, + offset, + filters, + .. + } = plan + else { + return Ok(None); + }; + refuse_scan_clauses(plan, "rrf_score()")?; + let Some(limit) = limit else { + return Err(missing_limit_error()); + }; + let base = HybridBase { + table, + functions, + top_k: limit.saturating_add(*offset), + filters: filters.clone(), + score_alias, + projection: source_projection(plan), }; - let projection = source_projection(plan); // Determine whether args[2] (if present) is a function call (graph source) // or a numeric literal (k-constant for the two-source form). let third_is_graph_score = args.get(2).is_some_and(is_function_call); if third_is_graph_score { - plan_hybrid_triple(args, collection, limit, score_alias, projection) + plan_hybrid_triple(args, base) } else { - plan_hybrid_two_source(args, collection, limit, score_alias, projection) + plan_hybrid_two_source(args, base) } } /// Two-source: `rrf_score(vector_distance(...), bm25_score(...), k1?, k2?)`. -fn plan_hybrid_two_source( - args: &[ast::Expr], - collection: &str, - limit: usize, - score_alias: Option<&str>, - projection: Vec, -) -> Result> { +fn plan_hybrid_two_source(args: &[ast::Expr], base: HybridBase<'_>) -> Result> { // args[2] and args[3] are optional k-constants. If there are more than 4 // args in the two-source form, something is wrong. if args.len() > 4 { @@ -83,40 +120,32 @@ fn plan_hybrid_two_source( }); } - let vector = extract_vector_arg(&args[0])?; - let text = extract_text_arg(&args[1])?; - let k1 = args - .get(2) - .and_then(|e| extract_float(e).ok()) - .unwrap_or(60.0); - let k2 = args - .get(3) - .and_then(|e| extract_float(e).ok()) - .unwrap_or(60.0); + let (vector_field, vector) = extract_vector_arg(&args[0], &base)?; + let text = extract_text_arg(&args[1], &base)?; + let k1 = rrf_k(args, 2, "k1")?; + let k2 = rrf_k(args, 3, "k2")?; let vector_weight = k2 as f32 / (k1 as f32 + k2 as f32); - Ok(Some(SqlPlan::HybridSearch { - collection: collection.into(), + Ok(Some(SqlPlan::HybridSearch(HybridSearchPlan { + collection: base.table.name.clone(), + vector_field, query_vector: vector, - query_text: text, - top_k: limit, - ef_search: limit * 2, + text_field: text.field, + query_text: text.query, + filters: base.filters, + top_k: base.top_k, + ef_search: base.top_k.saturating_mul(2), vector_weight, - fuzzy: true, - score_alias: score_alias.map(|s| s.to_string()), - projection, - })) + mode: text.options.params.mode, + fuzzy: text.options.params.fuzzy, + score_alias: base.score_alias.map(|s| s.to_string()), + projection: base.projection, + }))) } /// Three-source: `rrf_score(vector_distance(...), bm25_score(...), graph_score(...), k1?, k2?, k3?)`. -fn plan_hybrid_triple( - args: &[ast::Expr], - collection: &str, - limit: usize, - score_alias: Option<&str>, - projection: Vec, -) -> Result> { +fn plan_hybrid_triple(args: &[ast::Expr], base: HybridBase<'_>) -> Result> { // After the three source functions, we accept 0 or 3 k-constants. // Anything else (e.g. 1 or 2 k-constants) is an inconsistent arity. let k_count = args.len().saturating_sub(3); @@ -139,110 +168,125 @@ fn plan_hybrid_triple( }); } - let vector = extract_vector_arg(&args[0])?; - let text = extract_text_arg(&args[1])?; - let (graph_seed_id, graph_depth, graph_edge_label) = extract_graph_score_args(&args[2])?; + let (vector_field, vector) = extract_vector_arg(&args[0], &base)?; + let text = extract_text_arg(&args[1], &base)?; + let graph = extract_graph_arg(&args[2], &base)?; - let k1 = args - .get(3) - .and_then(|e| extract_float(e).ok()) - .unwrap_or(60.0); - let k2 = args - .get(4) - .and_then(|e| extract_float(e).ok()) - .unwrap_or(60.0); - let k3 = args - .get(5) - .and_then(|e| extract_float(e).ok()) - .unwrap_or(60.0); + let k1 = rrf_k(args, 3, "k1")?; + let k2 = rrf_k(args, 4, "k2")?; + let k3 = rrf_k(args, 5, "k3")?; - Ok(Some(SqlPlan::HybridSearchTriple { - collection: collection.into(), + Ok(Some(SqlPlan::HybridSearchTriple(HybridSearchTriplePlan { + collection: base.table.name.clone(), + vector_field, query_vector: vector, - query_text: text, - graph_seed_id, - graph_depth, - graph_edge_label, - top_k: limit, - ef_search: limit * 2, - fuzzy: true, + text_field: text.field, + query_text: text.query, + filters: base.filters, + graph_seed_id: graph.seed_id, + graph_depth: graph.depth, + graph_edge_label: graph.edge_label, + top_k: base.top_k, + ef_search: base.top_k.saturating_mul(2), + mode: text.options.params.mode, + fuzzy: text.options.params.fuzzy, rrf_k: (k1, k2, k3), - score_alias: score_alias.map(|s| s.to_string()), - projection, - })) -} - -/// Extract the float-array from a `vector_distance(col, ARRAY[...])` expression. -fn extract_vector_arg(expr: &ast::Expr) -> Result> { - Ok(match expr { - ast::Expr::Function(f) => { - let inner_args = extract_func_args(f)?; - if inner_args.len() >= 2 { - extract_float_array(&inner_args[1]).unwrap_or_default() - } else { - Vec::new() - } - } - _ => Vec::new(), - }) + score_alias: base.score_alias.map(|s| s.to_string()), + projection: base.projection, + }))) } -/// Extract the query string from a `bm25_score(col, 'query')` expression. -fn extract_text_arg(expr: &ast::Expr) -> Result { - Ok(match expr { - ast::Expr::Function(f) => { - let inner_args = extract_func_args(f)?; - if inner_args.len() >= 2 { - extract_string_literal(&inner_args[1]).unwrap_or_default() - } else { - String::new() - } - } - _ => String::new(), - }) +/// The function call of the leg in argument `position` of `rrf_score(...)`, +/// when its search trigger is `trigger`. Any other expression is a typed +/// error naming `shape`, the call the position takes. +fn leg_call<'e>( + expr: &'e ast::Expr, + base: &HybridBase<'_>, + trigger: SearchTrigger, + position: &str, + shape: &str, +) -> Result<(&'e ast::Function, String)> { + let leg_error = || SqlError::InvalidFunction { + detail: format!("rrf_score(): the {position} argument must be {shape}; got {expr}"), + }; + let ast::Expr::Function(f) = expr else { + return Err(leg_error()); + }; + let name = function_call_name(expr).ok_or_else(leg_error)?; + if base.functions.search_trigger(&name) != trigger { + return Err(leg_error()); + } + Ok((f, name)) } -/// Extract `(seed_id, depth, edge_label)` from a -/// `graph_score(node_id_col, seed_id, depth => N, label => 'edge')` expression. -/// -/// Named args (`depth => N`, `label => 'e'`) are represented in sqlparser as -/// `Expr::Named { name, arg }`. Positional arg[0] is the node_id column -/// (ignored — the executor resolves surrogates from the collection), arg[1] -/// is the seed node id string. -fn extract_graph_score_args(expr: &ast::Expr) -> Result<(String, usize, Option)> { - let ast::Expr::Function(f) = expr else { - return Ok((String::new(), 1, None)); +/// The column and query vector of the `vector_distance(column, [...])` leg. +fn extract_vector_arg(expr: &ast::Expr, base: &HybridBase<'_>) -> Result<(String, Vec)> { + const SHAPE: &str = "vector_distance(column, [...])"; + let (f, _) = leg_call(expr, base, SearchTrigger::VectorSearch, "first", SHAPE)?; + let leg_error = || SqlError::InvalidFunction { + detail: format!("rrf_score(): the first argument must be {SHAPE}; got {expr}"), }; let inner_args = extract_func_args(f)?; + let [column, query, ..] = inner_args.as_slice() else { + return Err(leg_error()); + }; + let field = resolve_table_column(column, base.table)?.ok_or_else(leg_error)?; + Ok((field, extract_float_array(query)?)) +} - // arg[1] is the seed node id. - let seed_id = inner_args - .get(1) - .and_then(|e| extract_string_literal(e).ok()) - .unwrap_or_default(); +/// The column, query string, and options of the +/// `bm25_score(column, 'query', ...)` leg. `None` column for +/// `bm25_score(*, 'query')`: the whole-document index. +fn extract_text_arg(expr: &ast::Expr, base: &HybridBase<'_>) -> Result { + let (f, name) = leg_call( + expr, + base, + SearchTrigger::TextSearch, + "second", + "bm25_score(column, 'query')", + )?; + resolve_text_call(&name, f, base.table) +} - // Remaining args may be named: `depth => N` and `label => 'edge'`. - let mut depth: usize = 1; - let mut edge_label: Option = None; +/// The BFS spec of the `graph_score(node_id_col, 'seed', ...)` leg. +fn extract_graph_arg( + expr: &ast::Expr, + base: &HybridBase<'_>, +) -> Result { + let (f, _) = leg_call( + expr, + base, + SearchTrigger::GraphSearch, + "third", + "graph_score(node_id_col, 'seed', ...)", + )?; + parse_graph_score_args(f) +} - for arg in inner_args.iter().skip(2) { - if let ast::Expr::Named { name, expr } = arg { - let key = name.value.to_ascii_lowercase(); - match key.as_str() { - "depth" => { - if let Ok(d) = extract_float(expr) { - depth = d as usize; - } - } - "label" => { - edge_label = extract_string_literal(expr).ok(); - } - _ => {} - } - } +/// The RRF constant `name` at argument `index`, or [`DEFAULT_RRF_K`] when +/// the call names none. It must be a positive finite number. +fn rrf_k(args: &[ast::Expr], index: usize, name: &str) -> Result { + let Some(expr) = args.get(index) else { + return Ok(DEFAULT_RRF_K); + }; + let k = extract_float(expr).map_err(|_| SqlError::InvalidFunction { + detail: format!("rrf_score(): {name} must be a numeric constant; got {expr}"), + })?; + if !k.is_finite() || k <= 0.0 { + return Err(SqlError::InvalidFunction { + detail: format!("rrf_score(): {name} must be a positive number; got {expr}"), + }); } + Ok(k) +} - Ok((seed_id, depth, edge_label)) +/// The typed error for a `rrf_score(...)` query with no LIMIT. +fn missing_limit_error() -> SqlError { + SqlError::InvalidFunction { + detail: "rrf_score() fuses a ranked list of LIMIT rows; add a LIMIT clause \ + (e.g. ... ORDER BY score DESC LIMIT 10)" + .into(), + } } /// Returns true when `expr` is a `Function` call (rather than a numeric literal). diff --git a/nodedb-sql/src/planner/select/order_by/mod.rs b/nodedb-sql/src/planner/select/order_by/mod.rs index b6066d946..aeac2c57f 100644 --- a/nodedb-sql/src/planner/select/order_by/mod.rs +++ b/nodedb-sql/src/planner/select/order_by/mod.rs @@ -12,12 +12,16 @@ //! - `aliases` — alias resolution between ORDER BY and SELECT projection. //! - `triggers` — generic `SearchTrigger` → `SqlPlan` detection. //! - `hybrid` — `rrf_score(...)` → `SqlPlan::HybridSearch` construction. +//! - `graph_score` — the `graph_score(...)` leg of a three-source `rrf_score`. +//! - `text_score` — `bm25_score(...)` columns attached to a text or scan plan. //! - `vector_join` — `vector_distance ⋈ ARRAY_SLICE` fusion target detection. mod aliases; mod apply; +mod graph_score; mod hybrid; mod projection; +mod text_score; mod triggers; mod vector_join; diff --git a/nodedb-sql/src/planner/select/order_by/projection.rs b/nodedb-sql/src/planner/select/order_by/projection.rs index 9b99f3fb3..b5519545d 100644 --- a/nodedb-sql/src/planner/select/order_by/projection.rs +++ b/nodedb-sql/src/planner/select/order_by/projection.rs @@ -1,6 +1,6 @@ // SPDX-License-Identifier: Apache-2.0 -//! SELECT-projection fallback for hybrid-search and text-search trigger detection. +//! SELECT-projection pass for hybrid-search and text-score columns. //! //! When `apply_order_by` left the plan as a `Scan` (no ORDER BY, or an //! ORDER BY that did not match any search trigger), the `rrf_score(...)` or @@ -10,42 +10,41 @@ //! A score call no search plan serves is refused at plan time //! (`planner::search_scope`): it has no per-row value. //! -//! Text-search shape: `SELECT id, bm25_score(field, term) FROM c ORDER BY id`. -//! The plan stays a Scan after ORDER BY (non-search sort key). This pass -//! detects `bm25_score` or `text_match` in the SELECT list and promotes the -//! plan to `SqlPlan::TextSearch` with `score_alias` set. The converter then -//! emits `TextOp::BM25ScoreScan` — a full-collection scan where each row -//! receives the BM25 score for the query term (null if the term does not occur). +//! Every `bm25_score(column, q)` / `text_match(column, q)` in the SELECT list +//! becomes one score column. On a `TextSearch` plan the columns join the plan +//! and its shape stays as it is. A `Scan` becomes a score scan: every row its +//! filters admit, each with its score. A row the scoped index holds but the +//! query does not match scores `0.0`. A row the index does not hold scores +//! `null`. use sqlparser::ast; -use super::super::helpers::{extract_func_args, extract_string_literal}; +use super::super::helpers::extract_func_args; +use super::super::text_call::resolve_text_call; use super::aliases::function_call_name; use super::hybrid::{no_args_rrf_score_error, plan_hybrid_from_sort}; +use super::text_score::{ScanSort, attach_scores}; use crate::error::Result; -use crate::fts_types::FtsQuery; use crate::functions::registry::{FunctionRegistry, SearchTrigger}; use crate::parser::normalize::normalize_ident; -use crate::types::SqlPlan; +use crate::planner::select::post_process::post_process; +use crate::resolver::columns::ResolvedTable; +use crate::types::{SqlPlan, TextScoreColumn}; -/// Try to fire a hybrid-search trigger from the SELECT projection alone. +/// Fire a hybrid search or attach text-score columns from the SELECT list. /// -/// Also handles `bm25_score(field, term)` and `text_match(field, term)` in -/// the SELECT list when no ORDER BY search trigger fired. When detected, the -/// Scan (or existing TextSearch) is promoted to a `SqlPlan::TextSearch` with -/// `score_alias` set so the converter can emit `TextOp::BM25ScoreScan`. +/// `Ok(None)` when the SELECT list holds neither, or `plan` is neither a +/// `Scan` nor a `TextSearch`: no other plan becomes a search here. pub(in crate::planner::select) fn try_hybrid_from_projection( plan: &SqlPlan, select_items: &[ast::SelectItem], functions: &FunctionRegistry, + table: &ResolvedTable, ) -> Result> { - // Try hybrid first (rrf_score). - let collection = match plan { - SqlPlan::Scan { collection, .. } => collection.clone(), - SqlPlan::TextSearch { collection, .. } => collection.clone(), - _ => return Ok(None), - }; - + if !matches!(plan, SqlPlan::Scan { .. } | SqlPlan::TextSearch(_)) { + return Ok(None); + } + let mut scores = Vec::new(); for item in select_items { let (expr, alias) = match item { ast::SelectItem::ExprWithAlias { expr, alias } => (expr, Some(normalize_ident(alias))), @@ -56,118 +55,46 @@ pub(in crate::planner::select) fn try_hybrid_from_projection( continue; }; let name = function_call_name(expr).unwrap_or_default(); - match functions.search_trigger(&name) { SearchTrigger::HybridSearch => { - // Only fire hybrid trigger for Scan plans. - if !matches!(plan, SqlPlan::Scan { .. }) { + // Only a Scan becomes a hybrid search. + let SqlPlan::Scan { sort_keys, .. } = plan else { continue; - } + }; let args = extract_func_args(func)?; if args.is_empty() { return Err(no_args_rrf_score_error()); } - return plan_hybrid_from_sort(&args, &collection, plan, alias.as_deref()); - } - SearchTrigger::TextSearch => { - // bm25_score(field, term) in SELECT — promote to BM25ScoreScan. - let args = extract_func_args(func)?; - if args.len() < 2 { - continue; - } - let query_text = extract_string_literal(&args[1]).unwrap_or_default(); - if query_text.is_empty() { - continue; + let Some(hybrid) = + plan_hybrid_from_sort(&args, table, plan, alias.as_deref(), functions)? + else { + return Ok(None); + }; + // An ORDER BY that named no search trigger still sorts the + // fused rows. + if sort_keys.is_empty() { + return Ok(Some(hybrid)); } - // Use the explicit AS alias when present; otherwise use the - // stringified expression so the injected row field key matches - // the lookup key the pgwire projection layer derives from - // `UnnamedExpr.to_string()`. - let score_alias = alias.clone().unwrap_or_else(|| expr.to_string()); - return Ok(Some(build_text_search_score_scan( - plan, - &collection, - query_text, - score_alias, - ))); + return post_process(hybrid, sort_keys.clone(), None, 0).map(Some); } - SearchTrigger::TextMatch => { - // text_match(field, term) in SELECT — same as bm25_score. - let args = extract_func_args(func)?; - if args.len() < 2 { - continue; - } - let query_text = extract_string_literal(&args[1]).unwrap_or_default(); - if query_text.is_empty() { - continue; - } - let score_alias = alias.clone().unwrap_or_else(|| expr.to_string()); - return Ok(Some(build_text_search_score_scan( - plan, - &collection, - query_text, - score_alias, - ))); + SearchTrigger::TextSearch | SearchTrigger::TextMatch => { + let call = resolve_text_call(&name, func, table)?; + // The explicit AS alias, else the stringified expression, the + // key the pgwire projection layer derives from + // `UnnamedExpr.to_string()`. + scores.push(TextScoreColumn { + field: call.field, + query: call.query, + mode: call.options.params.mode, + fuzzy: call.options.params.fuzzy, + alias: alias.unwrap_or_else(|| expr.to_string()), + }); } _ => {} } } - Ok(None) -} - -/// Build a `SqlPlan::TextSearch` with `score_alias` set for a BM25ScoreScan. -/// -/// Carries forward filters from an existing `Scan` or `TextSearch` plan. -/// When the input is already a `TextSearch` (from a WHERE `text_match(...)`), -/// the existing query and filters are preserved and only the `score_alias` -/// is attached. -fn build_text_search_score_scan( - plan: &SqlPlan, - collection: &str, - query_text: String, - score_alias: String, -) -> SqlPlan { - match plan { - SqlPlan::TextSearch { - query, - top_k, - filters, - projection, - .. - } => SqlPlan::TextSearch { - collection: collection.to_string(), - query: query.clone(), - top_k: *top_k, - filters: filters.clone(), - score_alias: Some(score_alias), - projection: projection.clone(), - }, - SqlPlan::Scan { - filters, - limit, - projection, - .. - } => SqlPlan::TextSearch { - collection: collection.to_string(), - query: FtsQuery::Plain { - text: query_text, - fuzzy: true, - }, - top_k: limit.unwrap_or(10_000), - filters: filters.clone(), - score_alias: Some(score_alias), - projection: projection.clone(), - }, - _ => SqlPlan::TextSearch { - collection: collection.to_string(), - query: FtsQuery::Plain { - text: query_text, - fuzzy: true, - }, - top_k: 10_000, - filters: Vec::new(), - score_alias: Some(score_alias), - projection: Vec::new(), - }, + if scores.is_empty() { + return Ok(None); } + attach_scores(plan, scores, ScanSort::Keep) } diff --git a/nodedb-sql/src/planner/select/order_by/text_score.rs b/nodedb-sql/src/planner/select/order_by/text_score.rs new file mode 100644 index 000000000..d1a8c6269 --- /dev/null +++ b/nodedb-sql/src/planner/select/order_by/text_score.rs @@ -0,0 +1,94 @@ +// SPDX-License-Identifier: Apache-2.0 + +//! Attach `bm25_score(...)` columns to a plan. +//! +//! A `TextSearch` plan takes the columns as they are. A `Scan` becomes a +//! score scan over every row its filters admit. No other plan takes a score +//! column: it has no per-row BM25 value. + +use crate::error::{Result, SqlError}; +use crate::planner::select::post_process::post_process; +use crate::types::{SqlPlan, TextScoreColumn, TextSearchPlan, TextSearchShape}; + +/// Whether the scan's ORDER BY moves onto the score scan. +#[derive(Clone, Copy, PartialEq, Eq)] +pub(super) enum ScanSort { + /// The scan's sort keys come from an ORDER BY the caller does not + /// consume: the score scan keeps them in a post-processing tail. + Keep, + /// The caller consumes the ORDER BY: the scan's sort keys are dropped. + Consumed, +} + +/// `plan` with `scores` attached. `None` when `plan` takes no score column. +pub(super) fn attach_scores( + plan: &SqlPlan, + scores: Vec, + sort: ScanSort, +) -> Result> { + match plan { + SqlPlan::TextSearch(search) => { + let mut search = search.clone(); + for column in scores { + search.add_score(column); + } + Ok(Some(SqlPlan::TextSearch(search))) + } + SqlPlan::Scan { + collection, + filters, + projection, + sort_keys, + .. + } => { + refuse_scan_clauses(plan, "a bm25_score() scan")?; + let mut search = TextSearchPlan { + collection: collection.clone(), + shape: TextSearchShape::ScoreScan, + filters: filters.clone(), + scores: Vec::new(), + projection: projection.clone(), + }; + for column in scores { + search.add_score(column); + } + let search = SqlPlan::TextSearch(search); + if sort == ScanSort::Keep && !sort_keys.is_empty() { + return post_process(search, sort_keys.clone(), None, 0).map(Some); + } + Ok(Some(search)) + } + _ => Ok(None), + } +} + +/// Refuse rewriting a `Scan` into `target` when the scan carries a clause the +/// rewritten plan has no slot for. Rewriting would answer without it. +pub(super) fn refuse_scan_clauses(plan: &SqlPlan, target: &str) -> Result<()> { + let SqlPlan::Scan { + distinct, + window_functions, + temporal, + .. + } = plan + else { + return Ok(()); + }; + let dropped = if *distinct { + Some("DISTINCT") + } else if !window_functions.is_empty() { + Some("a window function") + } else if temporal.is_temporal() { + Some("AS OF") + } else { + None + }; + match dropped { + Some(clause) => Err(SqlError::Unsupported { + detail: format!( + "{clause} with {target} is not supported; wrap the search in a subquery" + ), + }), + None => Ok(()), + } +} diff --git a/nodedb-sql/src/planner/select/order_by/triggers.rs b/nodedb-sql/src/planner/select/order_by/triggers.rs index b023efee9..ebd9a5134 100644 --- a/nodedb-sql/src/planner/select/order_by/triggers.rs +++ b/nodedb-sql/src/planner/select/order_by/triggers.rs @@ -14,11 +14,14 @@ use super::super::helpers::{ extract_column_name, extract_float_array, extract_func_args, extract_string_literal, metric_from_func_name, source_projection, }; +use super::super::text_call::{resolve_table_column, resolve_text_call}; use super::aliases::function_call_name; use super::hybrid::{no_args_rrf_score_error, plan_hybrid_from_sort}; +use super::text_score::{ScanSort, attach_scores}; use super::vector_join::extract_vector_join_target; use crate::error::{Result, SqlError}; use crate::functions::registry::{FunctionRegistry, SearchTrigger}; +use crate::resolver::columns::ResolvedTable; use crate::types::*; /// Default `ef_search` multiplier applied when the user has not supplied @@ -26,20 +29,106 @@ use crate::types::*; /// the standard HNSW heuristic. const DEFAULT_EF_SEARCH_MULTIPLIER: usize = 2; +/// The vector column `function(query)` searches when it names none: the +/// vector-primary field, else the one declared vector column, else the +/// collection-level index (`""`). Several declared vector columns are +/// ambiguous, a typed error naming them. +fn default_vector_field(function: &str, table: Option<&ResolvedTable>) -> Result { + let Some(table) = table else { + return Ok(String::new()); + }; + if let Some(vp) = &table.info.vector_primary { + return Ok(vp.vector_field.clone()); + } + let vector_columns: Vec<&str> = table + .info + .columns + .iter() + .filter(|c| matches!(c.data_type, SqlDataType::Vector(_))) + .map(|c| c.name.as_str()) + .collect(); + match vector_columns.as_slice() { + [] => Ok(String::new()), + [only] => Ok((*only).to_owned()), + several => Err(SqlError::Unsupported { + detail: format!( + "{function}(query) names no column and '{}' has several vector columns \ + ({}); name one: {function}(column, query)", + table.name, + several.join(", ") + ), + }), + } +} + +/// `function(column)`: the call names the searched column but no query +/// vector. The one-argument signature takes a query vector, so no signature +/// matches: PostgreSQL reports this as `undefined_function` (`42883`). +fn no_query_vector_error(function: &str) -> SqlError { + SqlError::UndefinedFunction { + name: function.to_owned(), + } +} + +/// A search plan an ORDER BY trigger produced. +pub(super) enum SortSearch { + /// The plan returns its rows in the order the ORDER BY asked for. + Ranked(SqlPlan), + /// A text plan that carries the score under `alias`. Its rows are + /// sorted by that column. + Scored { plan: SqlPlan, alias: String }, +} + /// Try to detect a search-triggering function call. /// -/// `score_alias` is propagated only into hybrid-search plans — vector and -/// text searches return a fixed-shape response. +/// `score_alias` names the hybrid score column, and the score column of a +/// `bm25_score(...)` sort. `table` is the single relation in scope: the +/// text and hybrid triggers fire only over one. pub(super) fn try_extract_sort_search( expr: &ast::Expr, plan: &SqlPlan, functions: &FunctionRegistry, score_alias: Option<&str>, -) -> Result> { + table: Option<&ResolvedTable>, +) -> Result> { let ast::Expr::Function(func) = expr else { return Ok(None); }; let name = function_call_name(expr).unwrap_or_default(); + match functions.search_trigger(&name) { + // ORDER BY bm25_score(column, q): the plan carries the score column + // and the rows sort by it. + SearchTrigger::TextSearch => { + let Some(table) = table else { + return Ok(None); + }; + let call = resolve_text_call(&name, func, table)?; + let alias = score_alias.map_or_else(|| expr.to_string(), str::to_owned); + let column = TextScoreColumn { + field: call.field, + query: call.query, + mode: call.options.params.mode, + fuzzy: call.options.params.fuzzy, + alias: alias.clone(), + }; + return Ok(attach_scores(plan, vec![column], ScanSort::Consumed)? + .map(|plan| SortSearch::Scored { plan, alias })); + } + SearchTrigger::HybridSearch => { + let (Some(table), SqlPlan::Scan { .. }) = (table, plan) else { + return Ok(None); + }; + let args = extract_func_args(func)?; + if args.is_empty() { + return Err(no_args_rrf_score_error()); + } + return Ok( + plan_hybrid_from_sort(&args, table, plan, score_alias, functions)? + .map(SortSearch::Ranked), + ); + } + _ => {} + } let (collection, array_prefilter) = match plan { SqlPlan::Scan { collection, .. } => (collection.clone(), None), SqlPlan::Join { left, right, .. } => match extract_vector_join_target(left, right) { @@ -56,11 +145,31 @@ pub(super) fn try_extract_sort_search( match functions.search_trigger(&name) { SearchTrigger::VectorSearch => { - if args.len() < 2 { - return Ok(None); - } - let field = extract_column_name(&args[0])?; - let vector = extract_float_array(&args[1])?; + let (field, query_arg) = match args.as_slice() { + [] => return Ok(None), + // A lone column is the searched field with no query vector: + // no signature takes it, as PostgreSQL reports `42883`. + [ast::Expr::Identifier(_) | ast::Expr::CompoundIdentifier(_)] => { + return Err(no_query_vector_error(&name)); + } + // `vector_distance(query)` (the `SEARCH c USING VECTOR(q, k)` + // form) searches the collection's default vector column. + [query] => (default_vector_field(&name, table)?, query), + // Over one relation, its qualifier is redundant: `t.embedding` + // names the `embedding` column. + [column, query, ..] => { + let field = match table { + Some(table) => resolve_table_column(column, table)?.ok_or_else(|| { + SqlError::Unsupported { + detail: format!("expected column name, got: {column}"), + } + })?, + None => extract_column_name(column)?, + }; + (field, query) + } + }; + let vector = extract_float_array(query_arg)?; let ann_options = parse_ann_options(raw_func_args)?; let limit = match plan { SqlPlan::Scan { limit, .. } => limit.unwrap_or(10), @@ -73,7 +182,7 @@ pub(super) fn try_extract_sort_search( .ef_search_override .unwrap_or(limit * DEFAULT_EF_SEARCH_MULTIPLIER); let metric = metric_from_func_name(&name); - Ok(Some(SqlPlan::VectorSearch { + Ok(Some(SortSearch::Ranked(SqlPlan::VectorSearch { collection, field, query_vector: vector, @@ -91,8 +200,9 @@ pub(super) fn try_extract_sort_search( // fields after `apply_order_by` returns. skip_payload_fetch: false, payload_filters: Vec::new(), + pk_prefilter: None, projection: source_projection(plan), - })) + }))) } SearchTrigger::SparseSearch => { if args.len() < 2 { @@ -115,41 +225,13 @@ pub(super) fn try_extract_sort_search( SqlPlan::Join { limit, .. } => limit.unwrap_or(10), _ => 10, }; - Ok(Some(SqlPlan::SparseSearch { + Ok(Some(SortSearch::Ranked(SqlPlan::SparseSearch { collection, field, query_entries, top_k, projection: source_projection(plan), - })) - } - SearchTrigger::TextSearch if args.len() >= 2 => { - let query_text = extract_string_literal(&args[1])?; - let limit = match plan { - SqlPlan::Scan { limit, .. } => limit.unwrap_or(10), - _ => 10, - }; - Ok(Some(SqlPlan::TextSearch { - collection, - query: crate::fts_types::FtsQuery::Plain { - text: query_text, - fuzzy: true, - }, - top_k: limit, - filters: match plan { - SqlPlan::Scan { filters, .. } => filters.clone(), - _ => Vec::new(), - }, - score_alias: score_alias.map(|s| s.to_string()), - projection: source_projection(plan), - })) - } - SearchTrigger::TextSearch => Ok(None), - SearchTrigger::HybridSearch => { - if args.is_empty() { - return Err(no_args_rrf_score_error()); - } - plan_hybrid_from_sort(&args, &collection, plan, score_alias) + }))) } _ => Ok(None), } diff --git a/nodedb-sql/src/planner/select/text_call.rs b/nodedb-sql/src/planner/select/text_call.rs new file mode 100644 index 000000000..b04c4f7d4 --- /dev/null +++ b/nodedb-sql/src/planner/select/text_call.rs @@ -0,0 +1,113 @@ +// SPDX-License-Identifier: Apache-2.0 + +//! Argument resolution for the search functions that name a column: +//! `text_match(column, q, ...)`, `bm25_score(column, q, ...)`, and the vector +//! column of `vector_distance(column, q)` inside `rrf_score(...)`. + +use sqlparser::ast; + +use super::helpers::{extract_column_name, extract_string_literal}; +use super::text_options::{TextOptions, parse_text_options}; +use crate::error::{Result, SqlError}; +use crate::parser::normalize::{normalize_ident, normalize_object_name_checked}; +use crate::resolver::columns::ResolvedTable; + +/// The column, query, and options of `text_match(column, q, ...)` / +/// `bm25_score(column, q, ...)`. +pub(super) struct TextCall { + /// `None` for `*` or `.*`: the whole-document index. + pub field: Option, + pub query: String, + /// The named `mode` / `fuzzy` options. + pub options: TextOptions, +} + +/// Resolve a column reference to `table`: `col` or `
.col`, where the +/// qualifier is the table's name or alias. `None` when `expr` is no column +/// reference. A qualifier naming another relation is `UnknownTable`. +pub(super) fn resolve_table_column( + expr: &ast::Expr, + table: &ResolvedTable, +) -> Result> { + match expr { + ast::Expr::Identifier(ident) => Ok(Some(normalize_ident(ident))), + ast::Expr::CompoundIdentifier(parts) if parts.len() == 2 => { + let qualifier = normalize_ident(&parts[0]); + check_qualifier(&qualifier, table)?; + Ok(Some(normalize_ident(&parts[1]))) + } + ast::Expr::CompoundIdentifier(_) => extract_column_name(expr).map(Some), + _ => Ok(None), + } +} + +/// Accept `qualifier` when it names `table` or its alias. +fn check_qualifier(qualifier: &str, table: &ResolvedTable) -> Result<()> { + if qualifier == table.name || table.alias.as_deref() == Some(qualifier) { + Ok(()) + } else { + Err(SqlError::UnknownTable { + name: qualifier.to_owned(), + }) + } +} + +/// Resolve the arguments of a text-search call against `table`. +/// +/// The first argument is a text column of `table` (bare or qualified by the +/// table's name or alias), or `*` / `
.*` for the whole document. Any +/// other expression is a typed error. Named options follow the query; see +/// [`parse_text_options`]. +pub(super) fn resolve_text_call( + function: &str, + func: &ast::Function, + table: &ResolvedTable, +) -> Result { + let arity = || SqlError::Arity { + detail: format!("{function}() takes (column, 'query', [mode => ..., fuzzy => ...])"), + }; + let ast::FunctionArguments::List(list) = &func.args else { + return Err(arity()); + }; + let options = parse_text_options(function, &list.args)?; + let mut unnamed = list.args.iter().filter_map(|a| match a { + ast::FunctionArg::Unnamed(e) => Some(e), + _ => None, + }); + let first = unnamed.next().ok_or_else(arity)?; + let ast::FunctionArgExpr::Expr(query_expr) = unnamed.next().ok_or_else(arity)? else { + return Err(arity()); + }; + let field = match first { + ast::FunctionArgExpr::Wildcard => None, + ast::FunctionArgExpr::QualifiedWildcard(name) => { + check_qualifier(&normalize_object_name_checked(name)?, table)?; + None + } + ast::FunctionArgExpr::Expr(expr) => match resolve_table_column(expr, table)? { + Some(column) => { + table.check_text_column(function, &column)?; + Some(column) + } + None => return Err(not_a_column(function, table, first)), + }, + ast::FunctionArgExpr::WildcardWithOptions(_) => { + return Err(not_a_column(function, table, first)); + } + }; + Ok(TextCall { + field, + query: extract_string_literal(query_expr)?, + options, + }) +} + +/// The error for a text-search argument that is no column reference. +fn not_a_column(function: &str, table: &ResolvedTable, arg: &ast::FunctionArgExpr) -> SqlError { + SqlError::TextColumn { + function: function.to_owned(), + collection: table.name.clone(), + column: arg.to_string(), + fault: nodedb_types::text_search::TextColumnFault::NotAColumn, + } +} diff --git a/nodedb-sql/src/planner/select/text_options.rs b/nodedb-sql/src/planner/select/text_options.rs new file mode 100644 index 000000000..578c7c8bf --- /dev/null +++ b/nodedb-sql/src/planner/select/text_options.rs @@ -0,0 +1,252 @@ +// SPDX-License-Identifier: Apache-2.0 + +//! Named options of the text-search calls `text_match(column, 'q', ...)`, +//! `search(column, 'q', ...)`, and `bm25_score(column, 'q', ...)`. +//! +//! ```sql +//! SELECT id FROM docs WHERE text_match(body, 'rust db', mode => 'and', fuzzy => true) +//! ``` +//! +//! The keys are a closed set: `mode` (`'or'` or `'and'`) and `fuzzy` (`true` +//! or `false`). Each takes `=>`. An omitted option takes its +//! [`TextSearchParams::default`] value, the default of the native `text_search` +//! API. A third positional argument, an unknown key, a repeated key, and `=` +//! in place of `=>` are typed errors. + +use nodedb_types::text_search::{QueryMode, TextSearchParams}; +use sqlparser::ast::{self, FunctionArgOperator}; + +use crate::error::{Result, SqlError}; + +/// The keys a text-search call accepts. +const OPTION_NAMES: &str = "mode, fuzzy"; + +/// The options of one text-search call. +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub(super) struct TextOptions { + pub params: TextSearchParams, + /// Whether the call names any option. + pub named: bool, +} + +/// Parse the named options of `function`'s argument list. The first two +/// positional arguments (column, query) are skipped. +pub(super) fn parse_text_options(function: &str, args: &[ast::FunctionArg]) -> Result { + let mut params = TextSearchParams::default(); + let mut mode_seen = false; + let mut fuzzy_seen = false; + let mut positional: usize = 0; + for arg in args { + let (key, value, operator) = match arg { + ast::FunctionArg::Unnamed(arg_expr) => { + if positional >= 2 { + return Err(positional_option(function, arg_expr)); + } + positional += 1; + continue; + } + ast::FunctionArg::Named { + name, + arg, + operator, + } => (name.value.to_ascii_lowercase(), arg, operator), + ast::FunctionArg::ExprNamed { + name: ast::Expr::Identifier(ident), + arg, + operator, + } => (ident.value.to_ascii_lowercase(), arg, operator), + ast::FunctionArg::ExprNamed { name, .. } => { + return Err(SqlError::Unsupported { + detail: format!( + "{function}(): option name {name} is not an identifier; \ + name an option: {OPTION_NAMES} (e.g. mode => 'or')" + ), + }); + } + }; + if *operator != FunctionArgOperator::RightArrow { + return Err(SqlError::Unsupported { + detail: format!( + "{function}(): use '=>' not '{operator}' for text-search options \ + (e.g. mode => 'or')" + ), + }); + } + let ast::FunctionArgExpr::Expr(value) = value else { + return Err(SqlError::Unsupported { + detail: format!("{function}(): option '{key}' expects a value, got {value}"), + }); + }; + match key.as_str() { + "mode" => { + if mode_seen { + return Err(duplicate(function, "mode")); + } + mode_seen = true; + params.mode = parse_mode(function, value)?; + } + "fuzzy" => { + if fuzzy_seen { + return Err(duplicate(function, "fuzzy")); + } + fuzzy_seen = true; + params.fuzzy = parse_fuzzy(function, value)?; + } + other => { + return Err(SqlError::Unsupported { + detail: format!( + "{function}(): unknown text-search option '{other}'; \ + valid options: {OPTION_NAMES}" + ), + }); + } + } + } + Ok(TextOptions { + params, + named: mode_seen || fuzzy_seen, + }) +} + +/// The error for a third positional argument. `key = value` names the `=>` +/// form it meant. +fn positional_option(function: &str, arg: &ast::FunctionArgExpr) -> SqlError { + let is_eq_option = matches!( + arg, + ast::FunctionArgExpr::Expr(ast::Expr::BinaryOp { + left, + op: ast::BinaryOperator::Eq, + .. + }) if matches!(**left, ast::Expr::Identifier(_)) + ); + let detail = if is_eq_option { + format!("{function}(): use '=>' not '=' for text-search options (e.g. mode => 'or')") + } else { + format!( + "{function}() takes (column, 'query') and named options {OPTION_NAMES}; \ + a third positional argument is not accepted, got {arg}. \ + Use named options: {function}(column, 'query', mode => 'or', fuzzy => true)" + ) + }; + SqlError::Unsupported { detail } +} + +fn duplicate(function: &str, name: &str) -> SqlError { + SqlError::Unsupported { + detail: format!("{function}(): option '{name}' specified more than once"), + } +} + +/// `'or'` or `'and'`, in any ASCII case. +fn parse_mode(function: &str, value: &ast::Expr) -> Result { + let text = match value { + ast::Expr::Value(v) => match &v.value { + ast::Value::SingleQuotedString(s) => Some(s.as_str()), + _ => None, + }, + _ => None, + }; + text.and_then(QueryMode::parse) + .ok_or_else(|| SqlError::Unsupported { + detail: format!("{function}(): option 'mode' expects 'or' or 'and', got {value}"), + }) +} + +/// `true` or `false`. +fn parse_fuzzy(function: &str, value: &ast::Expr) -> Result { + let flag = match value { + ast::Expr::Value(v) => match &v.value { + ast::Value::Boolean(b) => Some(*b), + _ => None, + }, + _ => None, + }; + flag.ok_or_else(|| SqlError::Unsupported { + detail: format!("{function}(): option 'fuzzy' expects true or false, got {value}"), + }) +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::parser::statement::parse_sql; + + /// The argument list of the first projected call of `SELECT `. + fn call_args(call: &str) -> Vec { + let statements = parse_sql(&format!("SELECT {call}")).expect("parse"); + let ast::Statement::Query(query) = &statements[0] else { + panic!("expected a query"); + }; + let ast::SetExpr::Select(select) = query.body.as_ref() else { + panic!("expected a SELECT"); + }; + let ast::SelectItem::UnnamedExpr(ast::Expr::Function(func)) = &select.projection[0] else { + panic!("expected a function call"); + }; + let ast::FunctionArguments::List(list) = &func.args else { + panic!("expected an argument list"); + }; + list.args.clone() + } + + fn parse(call: &str) -> Result { + parse_text_options("text_match", &call_args(call)) + } + + fn detail(err: SqlError) -> String { + match err { + SqlError::Unsupported { detail } => detail, + other => panic!("expected Unsupported, got {other:?}"), + } + } + + #[test] + fn no_options_take_the_trait_default() { + let opts = parse("text_match(body, 'q')").expect("parse"); + assert_eq!(opts.params, TextSearchParams::default()); + assert!(!opts.named); + } + + #[test] + fn mode_and_fuzzy_parse() { + let opts = parse("text_match(body, 'q', mode => 'AND', fuzzy => true)").expect("parse"); + assert_eq!(opts.params.mode, QueryMode::And); + assert!(opts.params.fuzzy); + assert!(opts.named); + let opts = parse("text_match(body, 'q', fuzzy => false)").expect("parse"); + assert_eq!(opts.params.mode, QueryMode::Or); + assert!(!opts.params.fuzzy); + } + + #[test] + fn an_unknown_key_is_refused() { + let err = parse("text_match(body, 'q', boost => 2)").expect_err("unknown"); + assert!(detail(err).contains("unknown text-search option 'boost'")); + } + + #[test] + fn an_equals_operator_is_refused() { + let err = parse("text_match(body, 'q', mode = 'or')").expect_err("="); + assert!(detail(err).contains("use '=>'")); + } + + #[test] + fn a_positional_third_argument_is_refused() { + let err = parse("text_match(body, 'q', '{\"fuzzy\":true}')").expect_err("positional"); + assert!(detail(err).contains("third positional argument")); + } + + #[test] + fn a_repeated_key_is_refused() { + let err = parse("text_match(body, 'q', mode => 'or', mode => 'and')").expect_err("dup"); + assert!(detail(err).contains("more than once")); + } + + #[test] + fn a_bad_value_is_refused() { + let err = parse("text_match(body, 'q', mode => 'xor')").expect_err("mode"); + assert!(detail(err).contains("'or' or 'and'")); + let err = parse("text_match(body, 'q', fuzzy => 'yes')").expect_err("fuzzy"); + assert!(detail(err).contains("true or false")); + } +} diff --git a/nodedb-sql/src/planner/select/where_search.rs b/nodedb-sql/src/planner/select/where_search.rs index e1be4403e..e978a068f 100644 --- a/nodedb-sql/src/planner/select/where_search.rs +++ b/nodedb-sql/src/planner/select/where_search.rs @@ -17,8 +17,9 @@ use sqlparser::ast; use super::entry_ann::parse_ann_options; use super::helpers::{ convert_where_to_filters, extract_column_name, extract_float, extract_float_array, - extract_func_args, extract_string_literal, metric_from_func_name, + extract_func_args, metric_from_func_name, }; +use super::text_call::{TextCall, resolve_text_call}; use crate::error::{Result, SqlError}; use crate::functions::registry::{FunctionRegistry, SearchTrigger}; use crate::parser::normalize::normalize_ident; @@ -42,25 +43,25 @@ pub(super) fn try_extract_where_search( functions: &FunctionRegistry, projection: &[Projection], ) -> Result> { - try_extract_with_extra_filters(expr, table, functions, None, projection) + try_extract_with_extra_filters(expr, table, functions, &[], projection) } -/// Internal entry that threads an optional sibling-AND filter through to the -/// concrete trigger handler. The public entry calls this with `None`; the -/// AND recursion branch calls it with the *other* side of the AND so a -/// vector / spatial / text trigger can carry the sibling predicate as a -/// scan filter instead of silently dropping it. +/// Internal entry that threads the sibling-AND predicates through to the +/// concrete trigger handler. The public entry calls this with none; each AND +/// recursion adds the *other* side of that AND, so a vector / spatial / text +/// trigger nested in `a AND (t AND b)` carries both `a` and `b` as scan +/// filters instead of silently dropping one. fn try_extract_with_extra_filters( expr: &ast::Expr, table: &crate::resolver::columns::ResolvedTable, functions: &FunctionRegistry, - extra_filter: Option<&ast::Expr>, + extra_filters: &[&ast::Expr], projection: &[Projection], ) -> Result> { match expr { ast::Expr::Function(func) => { let name = function_name(func); - dispatch_trigger(&name, func, table, functions, extra_filter, projection) + dispatch_trigger(&name, func, table, functions, extra_filters, projection) } // AND: recurse on each side, carrying the other side as a scan filter. ast::Expr::BinaryOp { @@ -68,32 +69,69 @@ fn try_extract_with_extra_filters( op: ast::BinaryOperator::And, right, } => { - // Try left as the trigger, with right as the carried filter. - if let Some(plan) = try_extract_with_extra_filters( - left, - table, - functions, - Some(right.as_ref()), - projection, - )? { + // Try left as the trigger, with right as a carried filter. + let mut with_right = extra_filters.to_vec(); + with_right.push(right.as_ref()); + if let Some(plan) = + try_extract_with_extra_filters(left, table, functions, &with_right, projection)? + { return Ok(Some(plan)); } - // Try right as the trigger, with left as the carried filter. - if let Some(plan) = try_extract_with_extra_filters( - right, - table, - functions, - Some(left.as_ref()), - projection, - )? { + // Try right as the trigger, with left as a carried filter. + let mut with_left = extra_filters.to_vec(); + with_left.push(left.as_ref()); + if let Some(plan) = + try_extract_with_extra_filters(right, table, functions, &with_left, projection)? + { return Ok(Some(plan)); } Ok(None) } + // A parenthesised conjunction is the same conjunction. + ast::Expr::Nested(inner) => { + try_extract_with_extra_filters(inner, table, functions, extra_filters, projection) + } _ => Ok(None), } } +/// The SELECT clauses a WHERE-derived search plan has no slot for. +pub(super) struct SearchBodyClauses<'a> { + pub temporal: &'a crate::temporal::TemporalScope, + pub has_subqueries: bool, + pub aggregates: bool, + pub distinct: bool, + pub windows: bool, +} + +/// Refuse a WHERE-derived search plan whose SELECT carries a clause the plan +/// cannot hold. Returning the plan would answer the query without it. +pub(super) fn refuse_dropped_clauses(plan: &SqlPlan, clauses: SearchBodyClauses<'_>) -> Result<()> { + let dropped = if clauses.temporal.is_temporal() { + Some("AS OF") + } else if clauses.has_subqueries { + Some("a WHERE subquery") + } else if clauses.aggregates { + Some("aggregation") + } else if clauses.distinct { + Some("DISTINCT") + } else if clauses.windows { + Some("a window function") + } else { + None + }; + match dropped { + Some(clause) => Err(SqlError::Unsupported { + detail: format!( + "{clause} over a WHERE-clause {} is not supported; \ + wrap the search in a subquery", + plan.variant_name() + ), + }), + None => Ok(()), + } +} + fn function_name(func: &ast::Function) -> String { func.name .0 @@ -111,7 +149,7 @@ fn dispatch_trigger( func: &ast::Function, table: &crate::resolver::columns::ResolvedTable, functions: &FunctionRegistry, - extra_filter: Option<&ast::Expr>, + extra_filters: &[&ast::Expr], projection: &[Projection], ) -> Result> { // Exhaustive match on `SearchTrigger`: when a new trigger is added, @@ -119,18 +157,20 @@ fn dispatch_trigger( // for it. This is the structural fix for the original bug class — // silent fall-through on unhandled triggers. match functions.search_trigger(name) { - SearchTrigger::TextMatch => plan_text_from_where(func, table, extra_filter, projection), + SearchTrigger::TextMatch => { + plan_text_from_where(name, func, table, extra_filters, projection) + } SearchTrigger::SpatialDWithin | SearchTrigger::SpatialContains | SearchTrigger::SpatialIntersects | SearchTrigger::SpatialWithin => { - plan_spatial_from_where(name, func, table, extra_filter, projection) + plan_spatial_from_where(name, func, table, extra_filters, projection) } SearchTrigger::VectorSearch => { - plan_vector_from_where(name, func, table, extra_filter, projection) + plan_vector_from_where(name, func, table, extra_filters, projection) } SearchTrigger::MultiVectorSearch => { - plan_multi_vector_from_where(func, table, extra_filter, projection) + plan_multi_vector_from_where(func, table, extra_filters, projection) } // The remaining triggers either have no WHERE-clause shape advertised // anywhere in the docs (`HybridSearch`, `TextSearch`, the array TVFs, @@ -160,32 +200,37 @@ fn dispatch_trigger( } } +/// The conjunction of every sibling-AND predicate, as scan filters. fn extra_filter_to_filters( - extra: Option<&ast::Expr>, + extra: &[&ast::Expr], table: &crate::resolver::columns::ResolvedTable, ) -> Result> { - match extra { - Some(e) => { - let scope = crate::resolver::columns::TableScope::single(table.clone())?; - convert_where_to_filters(e, &scope) - } - None => Ok(Vec::new()), + if extra.is_empty() { + return Ok(Vec::new()); + } + let scope = crate::resolver::columns::TableScope::single(table.clone())?; + let mut filters = Vec::new(); + for e in extra { + filters.extend(convert_where_to_filters(e, &scope)?); } + Ok(filters) } fn plan_text_from_where( + name: &str, func: &ast::Function, table: &crate::resolver::columns::ResolvedTable, - extra_filter: Option<&ast::Expr>, + extra_filters: &[&ast::Expr], projection: &[Projection], ) -> Result> { use crate::fts_types::FtsQuery; - let args = extract_func_args(func)?; - if args.len() < 2 { - return Ok(None); - } - let query_text = extract_string_literal(&args[1])?; + let TextCall { + field, + query: query_text, + options, + } = resolve_text_call(name, func, table)?; + let fuzzy = options.params.fuzzy; // Detect a phrase query: query_text surrounded by double-quotes. // SQL form: `text_match(body, '"quick brown fox"')`. @@ -200,31 +245,45 @@ fn plan_text_from_where( } else { FtsQuery::Plain { text: inner.to_string(), - fuzzy: true, + fuzzy, } } } else { FtsQuery::Plain { text: query_text, - fuzzy: true, + fuzzy, } }; + // A phrase matches its exact terms in order: neither option applies. + if options.named && matches!(fts_query, FtsQuery::Phrase(_)) { + return Err(SqlError::Unsupported { + detail: format!( + "{name}(): a phrase query matches its exact terms in order; \ + the mode and fuzzy options do not apply to it" + ), + }); + } - Ok(Some(SqlPlan::TextSearch { + // `top_k: None` returns every match; `apply_limit` sets a LIMIT's bound. + Ok(Some(SqlPlan::TextSearch(TextSearchPlan { collection: table.name.clone(), - query: fts_query, - top_k: 1000, - filters: extra_filter_to_filters(extra_filter, table)?, - score_alias: None, + shape: TextSearchShape::Match { + field, + query: fts_query, + mode: options.params.mode, + top_k: None, + }, + filters: extra_filter_to_filters(extra_filters, table)?, + scores: Vec::new(), projection: projection.to_vec(), - })) + }))) } fn plan_vector_from_where( name: &str, func: &ast::Function, table: &crate::resolver::columns::ResolvedTable, - extra_filter: Option<&ast::Expr>, + extra_filters: &[&ast::Expr], projection: &[Projection], ) -> Result> { let args = extract_func_args(func)?; @@ -250,7 +309,7 @@ fn plan_vector_from_where( top_k: DEFAULT_TOP_K, ef_search, metric: metric_from_func_name(name), - filters: extra_filter_to_filters(extra_filter, table)?, + filters: extra_filter_to_filters(extra_filters, table)?, array_prefilter: None, ann_options, // Vector-primary skip-payload-fetch and payload-filter peeling are @@ -260,6 +319,8 @@ fn plan_vector_from_where( // and need no special handling here. skip_payload_fetch: false, payload_filters: Vec::new(), + // Key conjuncts move here in the same post-pass. + pk_prefilter: None, projection: projection.to_vec(), })) } @@ -267,7 +328,7 @@ fn plan_vector_from_where( fn plan_multi_vector_from_where( func: &ast::Function, table: &crate::resolver::columns::ResolvedTable, - extra_filter: Option<&ast::Expr>, + extra_filters: &[&ast::Expr], projection: &[Projection], ) -> Result> { let args = extract_func_args(func)?; @@ -279,7 +340,7 @@ fn plan_multi_vector_from_where( // once the executor accepts filters on MultiVectorSearch; for now we keep // parity with the existing variant fields and raise on a sibling filter // so the user gets a clear "not yet supported here" instead of silent drop. - if extra_filter.is_some() { + if !extra_filters.is_empty() { return Err(SqlError::Unsupported { detail: "AND-combined predicates with multi_vector_distance(...) in WHERE are not supported; \ @@ -302,7 +363,7 @@ fn plan_spatial_from_where( name: &str, func: &ast::Function, table: &crate::resolver::columns::ResolvedTable, - extra_filter: Option<&ast::Expr>, + extra_filters: &[&ast::Expr], projection: &[Projection], ) -> Result> { let predicate = match name { @@ -346,7 +407,7 @@ fn plan_spatial_from_where( predicate, query_geometry: geometry, distance_meters: distance, - attribute_filters: extra_filter_to_filters(extra_filter, table)?, + attribute_filters: extra_filter_to_filters(extra_filters, table)?, limit: 1000, projection: projection.to_vec(), })) diff --git a/nodedb-sql/src/resolver/mod.rs b/nodedb-sql/src/resolver/mod.rs index 10ca28677..5ccdc5b82 100644 --- a/nodedb-sql/src/resolver/mod.rs +++ b/nodedb-sql/src/resolver/mod.rs @@ -5,5 +5,6 @@ pub mod columns; pub mod derived; pub mod expr; pub mod scope; +pub mod text_column; pub use scope::ColumnScope; diff --git a/nodedb-sql/src/resolver/text_column.rs b/nodedb-sql/src/resolver/text_column.rs new file mode 100644 index 000000000..ff817d36c --- /dev/null +++ b/nodedb-sql/src/resolver/text_column.rs @@ -0,0 +1,116 @@ +// SPDX-License-Identifier: Apache-2.0 + +//! Whether a column can hold the text a full-text search reads. + +use nodedb_types::text_search::TextColumnFault; + +use super::columns::ResolvedTable; +use crate::error::{Result, SqlError}; +use crate::types::SqlDataType; + +impl ResolvedTable { + /// Accept `column` as the field of `function(column, q)`: a declared + /// string column, or any name on an open-schema collection. + pub fn check_text_column(&self, function: &str, column: &str) -> Result<()> { + let fault = match self.info.columns.iter().find(|c| c.name == column) { + Some(c) if c.data_type == SqlDataType::String => return Ok(()), + Some(c) => TextColumnFault::NotText { + data_type: c + .raw_type + .clone() + .unwrap_or_else(|| format!("{:?}", c.data_type)), + }, + None if self.info.open_schema => return Ok(()), + None => TextColumnFault::Undeclared, + }; + Err(SqlError::TextColumn { + function: function.to_owned(), + collection: self.name.clone(), + column: column.to_owned(), + fault, + }) + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::types::{CollectionInfo, ColumnInfo, EngineType}; + + fn column(name: &str, data_type: SqlDataType, raw: &str) -> ColumnInfo { + ColumnInfo { + name: name.into(), + data_type, + nullable: true, + is_primary_key: false, + default: None, + raw_type: Some(raw.into()), + int_width: None, + float_width: None, + } + } + + fn strict() -> ResolvedTable { + let info = CollectionInfo { + name: "articles".into(), + engine: EngineType::DocumentStrict, + columns: vec![ + column("title", SqlDataType::String, "TEXT"), + column("views", SqlDataType::Int64, "INT"), + ], + primary_key: None, + has_auto_tier: false, + indexes: Vec::new(), + bitemporal: false, + primary: nodedb_types::PrimaryEngine::Document, + vector_primary: None, + partition_strategy: nodedb_types::PartitionStrategy::CollectionHomed, + open_schema: false, + }; + ResolvedTable { + name: "articles".into(), + alias: None, + info, + } + } + + #[test] + fn a_text_column_is_accepted() { + assert_eq!(strict().check_text_column("text_match", "title"), Ok(())); + } + + #[test] + fn an_int_column_is_not_text() { + let err = strict().check_text_column("bm25_score", "views"); + assert_eq!( + err, + Err(SqlError::TextColumn { + function: "bm25_score".into(), + collection: "articles".into(), + column: "views".into(), + fault: TextColumnFault::NotText { + data_type: "INT".into() + }, + }) + ); + } + + #[test] + fn an_undeclared_column_names_the_collection() { + let Err(SqlError::TextColumn { + collection, fault, .. + }) = strict().check_text_column("text_match", "ghost") + else { + panic!("expected a TextColumn error"); + }; + assert_eq!(collection, "articles"); + assert_eq!(fault, TextColumnFault::Undeclared); + } + + #[test] + fn an_open_schema_collection_accepts_any_name() { + let mut table = strict(); + table.info.open_schema = true; + assert_eq!(table.check_text_column("text_match", "ghost"), Ok(())); + } +} diff --git a/nodedb-sql/src/types/plan/variants/text.rs b/nodedb-sql/src/types/plan/variants/text.rs new file mode 100644 index 000000000..b96ccc212 --- /dev/null +++ b/nodedb-sql/src/types/plan/variants/text.rs @@ -0,0 +1,63 @@ +// SPDX-License-Identifier: Apache-2.0 + +//! Full-text search plan payload. + +use nodedb_types::text_search::QueryMode; + +use crate::fts_types::FtsQuery; +use crate::types::filter::Filter; +use crate::types::query::Projection; + +/// One `bm25_score(column, query)` the SELECT list or ORDER BY reads. +#[derive(Debug, Clone, PartialEq)] +pub struct TextScoreColumn { + /// `None` for `bm25_score(*, q)`: the whole-document index. + pub field: Option, + pub query: String, + /// `mode => 'and' | 'or'` of the call. + pub mode: QueryMode, + /// `fuzzy => true | false` of the call. + pub fuzzy: bool, + /// Output column the score lands in. + pub alias: String, +} + +/// What a text search returns. +#[derive(Debug, Clone)] +pub enum TextSearchShape { + /// `WHERE text_match(column, q)`: matching rows, best first. + /// `top_k: None` returns every match. + Match { + /// `None` for `text_match(*, q)`: the whole-document index. + field: Option, + /// The query. A `Plain` query carries the call's `fuzzy` option. + query: FtsQuery, + /// `mode => 'and' | 'or'` of the call. + mode: QueryMode, + top_k: Option, + }, + /// `bm25_score(...)` with no `text_match`: every row the filters admit. + ScoreScan, +} + +/// Payload of [`SqlPlan::TextSearch`](crate::types::SqlPlan::TextSearch). +#[derive(Debug, Clone)] +pub struct TextSearchPlan { + pub collection: String, + pub shape: TextSearchShape, + /// Residual WHERE predicates. They restrict candidates before ranking. + pub filters: Vec, + /// Score columns. Empty when no `bm25_score` is read. + pub scores: Vec, + /// Resolved SELECT target list, for output-schema derivation. + pub projection: Vec, +} + +impl TextSearchPlan { + /// Add `column` unless a score column of the same alias exists. + pub fn add_score(&mut self, column: TextScoreColumn) { + if !self.scores.iter().any(|s| s.alias == column.alias) { + self.scores.push(column); + } + } +} diff --git a/nodedb-sql/src/visitor/plan_visitor/args.rs b/nodedb-sql/src/visitor/plan_visitor/args.rs index 8fb1846fa..87043c344 100644 --- a/nodedb-sql/src/visitor/plan_visitor/args.rs +++ b/nodedb-sql/src/visitor/plan_visitor/args.rs @@ -10,8 +10,8 @@ use crate::temporal::TemporalScope; use crate::types::SqlPlan; use crate::types::filter::Filter; use crate::types::plan::{ - ArrayPrefilter, MergePlanClause, VectorAnnOptions, VectorPrimaryInsertIntent, VectorPrimaryRow, - WriteRoute, + ArrayPrefilter, MergePlanClause, TextScoreColumn, TextSearchShape, VectorAnnOptions, + VectorPrimaryInsertIntent, VectorPrimaryRow, WriteRoute, }; use crate::types::query::{ AggregateExpr, EngineType, JoinType, Projection, SortKey, SpatialPredicate, WindowSpec, @@ -153,16 +153,34 @@ pub struct VectorSearchVisitArgs<'a> { pub ann_options: &'a VectorAnnOptions, pub skip_payload_fetch: bool, pub payload_filters: &'a [SqlPayloadAtom], + /// Primary keys the search ranks among. `None`: every row. + pub pk_prefilter: Option<&'a [SqlValue]>, +} + +/// Parameters for [`super::trait_def::PlanVisitor::text_search`]. +pub struct TextSearchVisitArgs<'a> { + pub collection: &'a str, + pub shape: &'a TextSearchShape, + /// Residual WHERE predicates, applied before ranking. + pub filters: &'a [Filter], + /// Score columns the SELECT list or ORDER BY reads. + pub scores: &'a [TextScoreColumn], } /// Parameters for [`super::trait_def::PlanVisitor::hybrid_search`]. pub struct HybridSearchVisitArgs<'a> { pub collection: &'a str, + pub vector_field: &'a str, pub query_vector: &'a [f32], + /// `None` for `bm25_score(*, q)`: the whole-document index. + pub text_field: Option<&'a str>, pub query_text: &'a str, + /// Residual WHERE predicates, applied to both legs before fusion. + pub filters: &'a [Filter], pub top_k: usize, pub ef_search: usize, pub vector_weight: f32, + pub mode: nodedb_types::text_search::QueryMode, pub fuzzy: bool, pub score_alias: Option<&'a str>, } @@ -170,13 +188,19 @@ pub struct HybridSearchVisitArgs<'a> { /// Parameters for [`super::trait_def::PlanVisitor::hybrid_search_triple`]. pub struct HybridSearchTripleVisitArgs<'a> { pub collection: &'a str, + pub vector_field: &'a str, pub query_vector: &'a [f32], + /// `None` for `bm25_score(*, q)`: the whole-document index. + pub text_field: Option<&'a str>, pub query_text: &'a str, + /// Residual WHERE predicates, applied to every leg before fusion. + pub filters: &'a [Filter], pub graph_seed_id: &'a str, pub graph_depth: usize, pub graph_edge_label: Option<&'a str>, pub top_k: usize, pub ef_search: usize, + pub mode: nodedb_types::text_search::QueryMode, pub fuzzy: bool, pub rrf_k: (f64, f64, f64), pub score_alias: Option<&'a str>, diff --git a/nodedb-sql/src/visitor/plan_visitor/trait_def.rs b/nodedb-sql/src/visitor/plan_visitor/trait_def.rs index 420e7a292..f7e441673 100644 --- a/nodedb-sql/src/visitor/plan_visitor/trait_def.rs +++ b/nodedb-sql/src/visitor/plan_visitor/trait_def.rs @@ -10,10 +10,10 @@ use super::args::{ HybridSearchTripleVisitArgs, HybridSearchVisitArgs, InsertVisitArgs, JoinVisitArgs, LateralLoopVisitArgs, LateralTopKVisitArgs, MergeVisitArgs, RecursiveScanVisitArgs, RecursiveValueVisitArgs, ScanVisitArgs, SpatialScanVisitArgs, SubqueryVisitArgs, - TimeseriesScanVisitArgs, UpdateFromVisitArgs, UpsertVisitArgs, VectorPrimaryDeleteVisitArgs, - VectorPrimaryInsertVisitArgs, VectorPrimaryUpdateVisitArgs, VectorSearchVisitArgs, + TextSearchVisitArgs, TimeseriesScanVisitArgs, UpdateFromVisitArgs, UpsertVisitArgs, + VectorPrimaryDeleteVisitArgs, VectorPrimaryInsertVisitArgs, VectorPrimaryUpdateVisitArgs, + VectorSearchVisitArgs, }; -use crate::fts_types::FtsQuery; use crate::temporal::TemporalScope; use crate::types::SqlPlan; use crate::types::filter::Filter; @@ -172,14 +172,7 @@ pub trait PlanVisitor { ) -> Result; /// Handle [`SqlPlan::TextSearch`]. - fn text_search( - &mut self, - collection: &str, - query: &FtsQuery, - top_k: usize, - filters: &[Filter], - score_alias: Option<&str>, - ) -> Result; + fn text_search(&mut self, args: TextSearchVisitArgs<'_>) -> Result; /// Handle [`SqlPlan::HybridSearch`]. fn hybrid_search( diff --git a/nodedb-types/src/text_search.rs b/nodedb-types/src/text_search.rs index 237af0be6..f429d4d24 100644 --- a/nodedb-types/src/text_search.rs +++ b/nodedb-types/src/text_search.rs @@ -11,8 +11,18 @@ use serde::{Deserialize, Serialize}; /// Boolean query mode for full-text search. -#[derive(Debug, Clone, Copy, PartialEq, Eq, Default, Serialize, Deserialize)] -#[non_exhaustive] +#[derive( + Debug, + Clone, + Copy, + PartialEq, + Eq, + Default, + Serialize, + Deserialize, + zerompk::ToMessagePack, + zerompk::FromMessagePack, +)] pub enum QueryMode { /// Any query term can match (union). Most permissive — best recall. #[default] @@ -21,6 +31,28 @@ pub enum QueryMode { And, } +impl QueryMode { + /// The SQL spelling of the mode: `'or'` or `'and'`, as + /// `text_match(col, 'q', mode => 'or')` takes it. + pub fn as_str(self) -> &'static str { + match self { + Self::Or => "or", + Self::And => "and", + } + } + + /// Parse the SQL spelling, ignoring ASCII case. `None` names no mode. + pub fn parse(text: &str) -> Option { + if text.eq_ignore_ascii_case("or") { + Some(Self::Or) + } else if text.eq_ignore_ascii_case("and") { + Some(Self::And) + } else { + None + } + } +} + /// BM25 ranking parameters. /// /// Controls how term frequency and document length affect scoring. @@ -51,7 +83,8 @@ impl Default for Bm25Params { /// on document characteristics (length, vocabulary), not on individual queries. /// /// Pass [`TextSearchParams::default()`] for standard OR-mode non-fuzzy search. -#[derive(Debug, Clone, Serialize, Deserialize)] +/// SQL `text_match` / `bm25_score` with no options run the same defaults. +#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)] pub struct TextSearchParams { /// Boolean query mode: `Or` (any term) or `And` (all terms). /// Default: `Or`. @@ -70,3 +103,71 @@ impl Default for TextSearchParams { } } } + +/// Why a column cannot serve a full-text search. +/// +/// Carried by the planner error, the Data-Plane error code, and the server +/// error, so each surface renders the same SQLSTATE through [`Self::sqlstate`]. +#[derive(Debug, Clone, PartialEq, Eq, thiserror::Error, Serialize, Deserialize)] +pub enum TextColumnFault { + /// The strict schema declares no column of that name. + #[error("is not a column of the collection")] + Undeclared, + /// The column is declared with a non-text type. + #[error("has type {data_type}, not text")] + NotText { data_type: String }, + /// The argument is an expression or literal, not a column reference. + #[error("is not a column reference; name a text column or use *")] + NotAColumn, + /// No document of the collection holds text under that field, while + /// other fields hold text. + #[error("holds no indexed text in any document")] + NotIndexed, +} + +impl TextColumnFault { + /// SQLSTATE of the fault: `42703` (undefined_column) for a column that + /// does not exist as text, `42804` (datatype_mismatch) for an argument + /// that is not a text column. + pub fn sqlstate(&self) -> &'static str { + match self { + Self::Undeclared | Self::NotIndexed => crate::error::sqlstate::UNDEFINED_COLUMN, + Self::NotText { .. } | Self::NotAColumn => crate::error::sqlstate::DATATYPE_MISMATCH, + } + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn query_mode_sql_spelling_round_trips() { + for mode in [QueryMode::Or, QueryMode::And] { + assert_eq!(QueryMode::parse(mode.as_str()), Some(mode)); + } + assert_eq!(QueryMode::parse("AND"), Some(QueryMode::And)); + assert_eq!(QueryMode::parse("xor"), None); + } + + #[test] + fn query_mode_round_trips_through_msgpack() { + for mode in [QueryMode::Or, QueryMode::And] { + let bytes = zerompk::to_msgpack_vec(&mode).expect("encode"); + let back: QueryMode = zerompk::from_msgpack(&bytes).expect("decode"); + assert_eq!(back, mode); + } + } + + #[test] + fn faults_map_to_their_sqlstate() { + assert_eq!(TextColumnFault::Undeclared.sqlstate(), "42703"); + assert_eq!(TextColumnFault::NotIndexed.sqlstate(), "42703"); + assert_eq!(TextColumnFault::NotAColumn.sqlstate(), "42804"); + let not_text = TextColumnFault::NotText { + data_type: "INT".into(), + }; + assert_eq!(not_text.sqlstate(), "42804"); + assert_eq!(not_text.to_string(), "has type INT, not text"); + } +} diff --git a/nodedb/src/control/clone/resolver/rewrite.rs b/nodedb/src/control/clone/resolver/rewrite.rs index 80eaff0c0..fd102c785 100644 --- a/nodedb/src/control/clone/resolver/rewrite.rs +++ b/nodedb/src/control/clone/resolver/rewrite.rs @@ -383,30 +383,45 @@ mod tests { "text_search", PhysicalPlan::Text(TextOp::Search { collection: QualifiedCollection::from_stored(COLL.to_string()), + field: None, query: "hello".to_string(), top_k: 4, + mode: nodedb_types::text_search::QueryMode::And, fuzzy: false, prefilter: None, + filters: Vec::new(), rls_filters: Vec::new(), + scores: Vec::new(), }), ), ( "text_bm25_score_scan", PhysicalPlan::Text(TextOp::BM25ScoreScan { collection: QualifiedCollection::from_stored(COLL.to_string()), - query: "hello".to_string(), - score_alias: "score".to_string(), - fuzzy: false, + filters: Vec::new(), + rls_filters: Vec::new(), + scores: vec![nodedb_physical::physical_plan::TextScoreSpec { + field: None, + query: "hello".to_string(), + mode: nodedb_types::text_search::QueryMode::And, + fuzzy: false, + alias: "score".to_string(), + }], + bound: None, }), ), ( "text_hybrid_search", PhysicalPlan::Text(TextOp::HybridSearch { collection: QualifiedCollection::from_stored(COLL.to_string()), + vector_field: String::new(), query_vector: vec![0.0, 1.0], + text_field: None, query_text: "hello".to_string(), + filters: Vec::new(), top_k: 4, ef_search: 16, + mode: nodedb_types::text_search::QueryMode::And, fuzzy: false, vector_weight: 0.5, filter_bitmap: None, @@ -418,13 +433,17 @@ mod tests { "text_hybrid_search_triple", PhysicalPlan::Text(TextOp::HybridSearchTriple { collection: QualifiedCollection::from_stored(COLL.to_string()), + vector_field: String::new(), query_vector: vec![0.0, 1.0], + text_field: None, query_text: "hello".to_string(), + filters: Vec::new(), graph_seed_id: "n1".to_string(), graph_depth: 1, graph_edge_label: None, top_k: 4, ef_search: 16, + mode: nodedb_types::text_search::QueryMode::And, fuzzy: false, rrf_k: (60.0, 60.0, 60.0), filter_bitmap: None, diff --git a/nodedb/src/control/planner/redaction_refusal/mod.rs b/nodedb/src/control/planner/redaction_refusal/mod.rs index 15894a546..fea93149d 100644 --- a/nodedb/src/control/planner/redaction_refusal/mod.rs +++ b/nodedb/src/control/planner/redaction_refusal/mod.rs @@ -5,6 +5,7 @@ mod graph; mod lookup; mod plan; mod streaming_mv; +mod text; pub use plan::{ refuse_unredactable_graph_collection, refuse_unredactable_graph_match, diff --git a/nodedb/src/control/planner/redaction_refusal/text.rs b/nodedb/src/control/planner/redaction_refusal/text.rs new file mode 100644 index 000000000..0bbf38f0b --- /dev/null +++ b/nodedb/src/control/planner/redaction_refusal/text.rs @@ -0,0 +1,243 @@ +// SPDX-License-Identifier: BUSL-1.1 + +//! Refusal of full-text reads over a redacted column. +//! +//! A text match and a BM25 score are computed in the Data Plane over the +//! stored text, which is never redacted. `WHERE text_match(ssn, '123')` +//! selects rows by the stored `ssn`, and `bm25_score(ssn, '123')` returns a +//! number derived from it, so masking the result rows protects nothing. The +//! whole-document index (`text_match(*, q)`) reads every text field, so any +//! rule on the collection covers it. + +use nodedb_physical::physical_plan::{TextOp, TextScoreSpec}; +use nodedb_types::QualifiedCollection; + +use super::lookup::RefusalCtx; + +/// Refuse a text op that reads a redacted column. Exhaustive over [`TextOp`] +/// so a new text operation forces a decision here. +pub(super) fn refuse_text_op(op: &TextOp, ctx: &RefusalCtx<'_>) -> crate::Result<()> { + match op { + TextOp::Search { + collection, + field, + scores, + .. + } + | TextOp::PhraseSearch { + collection, + field, + scores, + .. + } => { + refuse_text_field(ctx, collection, field.as_deref(), "text_match")?; + refuse_scores(ctx, collection, scores) + } + TextOp::BM25ScoreScan { + collection, scores, .. + } => refuse_scores(ctx, collection, scores), + TextOp::HybridSearch { + collection, + text_field, + .. + } + | TextOp::HybridSearchTriple { + collection, + text_field, + .. + } => refuse_text_field(ctx, collection, text_field.as_deref(), "the text leg"), + // Index writes and analyzer configuration return no column values. + TextOp::FtsIndexDoc { .. } | TextOp::FtsDeleteDoc { .. } | TextOp::SetTextConfig { .. } => { + Ok(()) + } + } +} + +fn refuse_scores( + ctx: &RefusalCtx<'_>, + collection: &QualifiedCollection, + scores: &[TextScoreSpec], +) -> crate::Result<()> { + scores.iter().try_for_each(|spec| { + refuse_text_field(ctx, collection, spec.field.as_deref(), "bm25_score") + }) +} + +/// Refuse when `field` is redacted, or when `field` is `None` (the +/// whole-document index) and any rule covers the collection. +fn refuse_text_field( + ctx: &RefusalCtx<'_>, + collection: &QualifiedCollection, + field: Option<&str>, + reader: &str, +) -> crate::Result<()> { + let collection = collection.as_str(); + if collection.is_empty() { + return Ok(()); + } + match field { + Some(field) if ctx.field_is_redacted(collection, field) => Err(crate::Error::PlanError { + detail: format!( + "column '{field}' on '{collection}' is redacted for this role: {reader} over a \ + redacted column is not permitted — the match and score are computed over the \ + stored text, so masking the result would not protect it" + ), + }), + Some(_) => Ok(()), + None if ctx.collection_is_redacted(collection) => Err(crate::Error::PlanError { + detail: format!( + "'{collection}' has redacted columns for this role: {reader} over the whole \ + document is not permitted — it reads every text column, including the \ + redacted ones" + ), + }), + None => Ok(()), + } +} + +#[cfg(test)] +mod tests { + use nodedb_physical::physical_plan::{TextOp, TextScoreSpec}; + use nodedb_types::{DatabaseId, QualifiedCollection}; + + use crate::bridge::envelope::PhysicalPlan; + use crate::control::security::auth_context::AuthContext; + use crate::control::security::redaction::{ + RedactionMode, RedactionPolicy, RedactionRule, RedactionStore, + }; + use crate::types::TenantId; + + const TENANT: u64 = 1; + + fn store_with_rule(collection: &str, role: &str, field: &str) -> RedactionStore { + let store = RedactionStore::new(); + store.create_policy(RedactionPolicy { + name: format!("{collection}_{role}_{field}"), + tenant_id: TENANT, + collection: collection.into(), + display_collection: collection.into(), + for_role: role.into(), + rules: vec![RedactionRule { + field: field.into(), + mode: RedactionMode::Mask("***".into()), + }], + }); + store + } + + fn auth_with_role(role: &str) -> AuthContext { + use crate::control::security::identity::{ + AuthMethod, AuthenticatedIdentity, DatabaseSet, Role, + }; + let identity = AuthenticatedIdentity::new_regular( + 42, + "alice", + TenantId::new(TENANT), + AuthMethod::ScramSha256, + vec![Role::ReadWrite], + None, + DatabaseSet::Some(smallvec::smallvec![DatabaseId::DEFAULT]), + ); + let mut auth = AuthContext::from_identity(&identity, "s_test".into()); + auth.roles = vec![role.to_string()]; + auth + } + + fn check(plan: &PhysicalPlan, store: &RedactionStore) -> crate::Result<()> { + super::super::refuse_unredactable_plan( + plan, + TenantId::new(TENANT), + DatabaseId::DEFAULT, + &auth_with_role("support"), + store, + ) + } + + fn users() -> QualifiedCollection { + QualifiedCollection::new(DatabaseId::DEFAULT, "users") + } + + fn search(field: Option<&str>, scores: Vec) -> PhysicalPlan { + PhysicalPlan::Text(TextOp::Search { + collection: users(), + field: field.map(str::to_string), + query: "123".into(), + top_k: 10, + mode: nodedb_types::text_search::QueryMode::And, + fuzzy: false, + prefilter: None, + filters: Vec::new(), + rls_filters: Vec::new(), + scores, + }) + } + + fn score(field: Option<&str>) -> TextScoreSpec { + TextScoreSpec { + field: field.map(str::to_string), + query: "123".into(), + mode: nodedb_types::text_search::QueryMode::And, + fuzzy: false, + alias: "s".into(), + } + } + + fn assert_refused(result: crate::Result<()>) { + assert!( + matches!(result, Err(crate::Error::PlanError { .. })), + "expected a PlanError refusal, got {result:?}" + ); + } + + /// `WHERE text_match(ssn, …)` selects rows by the stored `ssn`. + #[test] + fn text_match_on_a_redacted_column_is_refused() { + let store = store_with_rule("users", "support", "ssn"); + assert_refused(check(&search(Some("ssn"), Vec::new()), &store)); + assert!(check(&search(Some("bio"), Vec::new()), &store).is_ok()); + } + + /// `bm25_score(ssn, …)` derives a number from the stored `ssn`, in a + /// search and in a score scan. + #[test] + fn bm25_score_on_a_redacted_column_is_refused() { + let store = store_with_rule("users", "support", "ssn"); + assert_refused(check( + &search(Some("bio"), vec![score(Some("ssn"))]), + &store, + )); + let scan = PhysicalPlan::Text(TextOp::BM25ScoreScan { + collection: users(), + filters: Vec::new(), + rls_filters: Vec::new(), + scores: vec![score(Some("ssn"))], + bound: None, + }); + assert_refused(check(&scan, &store)); + } + + /// The whole-document index reads every text column. + #[test] + fn whole_document_match_is_refused_under_any_rule() { + let store = store_with_rule("users", "support", "ssn"); + assert_refused(check(&search(None, Vec::new()), &store)); + let phrase = PhysicalPlan::Text(TextOp::PhraseSearch { + collection: users(), + field: None, + terms: vec!["a".into(), "b".into()], + top_k: 10, + prefilter: None, + filters: Vec::new(), + rls_filters: Vec::new(), + scores: Vec::new(), + }); + assert_refused(check(&phrase, &store)); + } + + /// A role the policy does not name reads the column. + #[test] + fn a_role_without_the_rule_is_not_refused() { + let store = store_with_rule("users", "analyst", "ssn"); + assert!(check(&search(Some("ssn"), vec![score(None)]), &store).is_ok()); + } +} diff --git a/nodedb/src/control/planner/rls_injection/permission_tree/text.rs b/nodedb/src/control/planner/rls_injection/permission_tree/text.rs index af065159d..20d39cab4 100644 --- a/nodedb/src/control/planner/rls_injection/permission_tree/text.rs +++ b/nodedb/src/control/planner/rls_injection/permission_tree/text.rs @@ -11,13 +11,23 @@ use super::context::{PermCtx, PermTreeLevel}; pub(super) fn apply_text(ctx: &PermCtx<'_>, op: &mut TextOp) -> crate::Result<()> { match op { // Filter: the subtree lands in the post-score / post-fusion slot the - // handler applies before the ranked hits are returned. The result may - // hold fewer than `top_k` rows, which is the intended effect. + // handler applies to every row before returning it. A ranked result + // may hold fewer than `top_k` rows, which is the intended effect. TextOp::Search { collection, rls_filters, .. } + | TextOp::BM25ScoreScan { + collection, + rls_filters, + .. + } + | TextOp::PhraseSearch { + collection, + rls_filters, + .. + } | TextOp::HybridSearch { collection, rls_filters, @@ -29,16 +39,6 @@ pub(super) fn apply_text(ctx: &PermCtx<'_>, op: &mut TextOp) -> crate::Result<() .. } => ctx.filter_into(collection, PermTreeLevel::Read, rls_filters), - // Refuse: the score scan emits every document in the collection with a - // score column appended, and the phrase search emits every positional - // hit — neither carries a filter slot for the subtree to occupy. - TextOp::BM25ScoreScan { collection, .. } | TextOp::PhraseSearch { collection, .. } => ctx - .refuse_if_tree( - collection, - "the search returns matched document rows through a response shape that carries \ - no subtree filter", - ), - // Filter (write level, blanket): indexing a document names the row it // indexes, so there is no predicate to narrow. TextOp::FtsIndexDoc { collection, .. } => ctx.authorize(collection, PermTreeLevel::Write), @@ -58,25 +58,32 @@ mod tests { use nodedb_physical::physical_plan::TextOp; use super::super::plan::test_support::{ - apply, assert_refused, cache_with_tree, injected_resources, readable, sorted, + apply, cache_with_tree, injected_resources, readable, sorted, }; use crate::bridge::envelope::PhysicalPlan; - /// A BM25 score scan returns every row of the collection with no slot for - /// the subtree, so it is refused rather than silently over-returning. + fn articles() -> nodedb_types::QualifiedCollection { + nodedb_types::QualifiedCollection::new(nodedb_types::DatabaseId::DEFAULT, "articles") + } + + /// A BM25 score scan carries the slot, so the subtree filters every row. #[test] - fn bm25_score_scan_is_refused_under_a_tree() { + fn bm25_score_scan_receives_the_subtree_filter() { let cache = cache_with_tree("articles"); let mut plan = PhysicalPlan::Text(TextOp::BM25ScoreScan { - collection: nodedb_types::QualifiedCollection::new( - nodedb_types::DatabaseId::DEFAULT, - "articles", - ), - query: "rust".into(), - score_alias: "score".into(), - fuzzy: false, + collection: articles(), + filters: Vec::new(), + rls_filters: Vec::new(), + scores: Vec::new(), + bound: None, }); - assert_refused(apply(&mut plan, &cache), "articles"); + assert!(apply(&mut plan, &cache).is_ok()); + match &plan { + PhysicalPlan::Text(TextOp::BM25ScoreScan { rls_filters, .. }) => { + assert_eq!(sorted(injected_resources(rls_filters)), readable()); + } + other => panic!("plan shape changed: {other:?}"), + } } /// A BM25 search does carry the slot, so the subtree is injected. @@ -84,15 +91,16 @@ mod tests { fn search_receives_the_subtree_filter() { let cache = cache_with_tree("articles"); let mut plan = PhysicalPlan::Text(TextOp::Search { - collection: nodedb_types::QualifiedCollection::new( - nodedb_types::DatabaseId::DEFAULT, - "articles", - ), + collection: articles(), + field: None, query: "rust".into(), top_k: 10, + mode: nodedb_types::text_search::QueryMode::And, fuzzy: false, prefilter: None, + filters: Vec::new(), rls_filters: Vec::new(), + scores: Vec::new(), }); assert!(apply(&mut plan, &cache).is_ok()); match &plan { diff --git a/nodedb/src/control/planner/rls_injection/text.rs b/nodedb/src/control/planner/rls_injection/text.rs index 2ae88ae98..0b084defc 100644 --- a/nodedb/src/control/planner/rls_injection/text.rs +++ b/nodedb/src/control/planner/rls_injection/text.rs @@ -10,14 +10,24 @@ use super::context::RlsCtx; /// between injecting, refusing, and no-op. pub(super) fn inject_text(ctx: &RlsCtx<'_>, op: &mut TextOp) -> crate::Result<()> { match op { - // Inject: the policy lands in the post-score / post-fusion slot the - // handler applies before the ranked hits are returned. The result can - // hold fewer than `top_k` rows, which is the intended effect. + // Inject: the policy lands in the RLS slot. The handler folds it into + // the eligible rows before ranking, so `top_k` counts only rows the + // policy admits. TextOp::Search { collection, rls_filters, .. } + | TextOp::BM25ScoreScan { + collection, + rls_filters, + .. + } + | TextOp::PhraseSearch { + collection, + rls_filters, + .. + } | TextOp::HybridSearch { collection, rls_filters, @@ -29,16 +39,6 @@ pub(super) fn inject_text(ctx: &RlsCtx<'_>, op: &mut TextOp) -> crate::Result<() .. } => ctx.set_post_filters(collection, rls_filters), - // Refuse: the score scan emits every document in the collection with a - // score column appended, and the phrase search emits every positional - // hit — neither carries a filter slot for the policy to occupy. - TextOp::BM25ScoreScan { collection, .. } | TextOp::PhraseSearch { collection, .. } => ctx - .refuse_if_policy( - collection, - "the search returns matched document rows through a response shape that carries \ - no row filter", - ), - // Refuse: an index write carries the extracted text and a surrogate, // not the row body the policy names. Indexing a row the policy hides // makes it reachable by search, which is the disclosure the policy @@ -61,8 +61,8 @@ mod tests { use nodedb_physical::physical_plan::TextOp; use super::super::plan::test_support::{ - assert_refused, assert_write_refused, inject, inject_without_policy, - store_with_read_policy, store_with_write_policy, + assert_write_refused, inject, inject_without_policy, store_with_read_policy, + store_with_write_policy, }; use crate::bridge::envelope::PhysicalPlan; @@ -73,7 +73,7 @@ mod tests { collection, ), surrogate: nodedb_types::Surrogate::new(1), - text: "hello".into(), + fields: vec![("body".into(), "hello".into())], provenance: None, }) } @@ -96,21 +96,59 @@ mod tests { assert_eq!(plan, before); } - /// A BM25 score scan returns every row of the collection with no slot for - /// the policy, so it is refused rather than silently over-returning. + fn articles() -> nodedb_types::QualifiedCollection { + nodedb_types::QualifiedCollection::new(nodedb_types::DatabaseId::DEFAULT, "articles") + } + + /// The RLS slot of a text read. + fn rls_slot(plan: &PhysicalPlan) -> &[u8] { + match plan { + PhysicalPlan::Text( + TextOp::Search { rls_filters, .. } + | TextOp::BM25ScoreScan { rls_filters, .. } + | TextOp::PhraseSearch { rls_filters, .. }, + ) => rls_filters, + other => panic!("plan shape changed: {other:?}"), + } + } + + /// A BM25 score scan applies the policy to every row it emits. #[test] - fn bm25_score_scan_is_refused_under_a_read_policy() { + fn bm25_score_scan_receives_the_policy_filter() { let store = store_with_read_policy("articles"); let mut plan = PhysicalPlan::Text(TextOp::BM25ScoreScan { - collection: nodedb_types::QualifiedCollection::new( - nodedb_types::DatabaseId::DEFAULT, - "articles", - ), - query: "rust".into(), - score_alias: "score".into(), - fuzzy: false, + collection: articles(), + filters: Vec::new(), + rls_filters: Vec::new(), + scores: Vec::new(), + bound: None, }); - assert_refused(inject(&mut plan, &store), "articles"); + assert!(inject(&mut plan, &store).is_ok()); + assert!( + !rls_slot(&plan).is_empty(), + "policy filter must be injected" + ); + } + + /// A phrase search applies the policy to its ranked hits. + #[test] + fn phrase_search_receives_the_policy_filter() { + let store = store_with_read_policy("articles"); + let mut plan = PhysicalPlan::Text(TextOp::PhraseSearch { + collection: articles(), + field: None, + terms: vec!["rust".into(), "lang".into()], + top_k: 10, + prefilter: None, + filters: Vec::new(), + rls_filters: Vec::new(), + scores: Vec::new(), + }); + assert!(inject(&mut plan, &store).is_ok()); + assert!( + !rls_slot(&plan).is_empty(), + "policy filter must be injected" + ); } /// A BM25 search does carry the slot, so the policy is injected. @@ -118,22 +156,56 @@ mod tests { fn search_receives_the_policy_filter() { let store = store_with_read_policy("articles"); let mut plan = PhysicalPlan::Text(TextOp::Search { - collection: nodedb_types::QualifiedCollection::new( - nodedb_types::DatabaseId::DEFAULT, - "articles", - ), + collection: articles(), + field: Some("body".into()), query: "rust".into(), top_k: 10, + mode: nodedb_types::text_search::QueryMode::And, fuzzy: false, prefilter: None, + filters: Vec::new(), rls_filters: Vec::new(), + scores: Vec::new(), }); assert!(inject(&mut plan, &store).is_ok()); - match &plan { - PhysicalPlan::Text(TextOp::Search { rls_filters, .. }) => { - assert!(!rls_filters.is_empty(), "policy filter must be injected") - } - other => panic!("plan shape changed: {other:?}"), - } + assert!( + !rls_slot(&plan).is_empty(), + "policy filter must be injected" + ); + } + + /// A slot that already holds the statement's predicates keeps them: the + /// policy joins them. + #[test] + fn the_policy_joins_existing_post_filters() { + let store = store_with_read_policy("articles"); + let own = vec![nodedb_query::scan_filter::ScanFilter { + field: "tag".into(), + op: nodedb_query::scan_filter::FilterOp::Eq, + value: nodedb_types::Value::String("a".into()), + clauses: Vec::new(), + expr: None, + }]; + let own_bytes = zerompk::to_msgpack_vec(&own).expect("encode own filters"); + let mut plan = PhysicalPlan::Text(TextOp::Search { + collection: articles(), + field: None, + query: "rust".into(), + top_k: 10, + mode: nodedb_types::text_search::QueryMode::And, + fuzzy: false, + prefilter: None, + filters: Vec::new(), + rls_filters: own_bytes, + scores: Vec::new(), + }); + assert!(inject(&mut plan, &store).is_ok()); + let merged: Vec = + zerompk::from_msgpack(rls_slot(&plan)).expect("decode merged filters"); + assert!( + merged.len() > 1, + "policy must join, not replace: {merged:?}" + ); + assert!(merged.iter().any(|f| f.field == "tag")); } } diff --git a/nodedb/src/control/planner/sql_plan_convert/filter/mod.rs b/nodedb/src/control/planner/sql_plan_convert/filter/mod.rs index 903fd27a1..03f35c362 100644 --- a/nodedb/src/control/planner/sql_plan_convert/filter/mod.rs +++ b/nodedb/src/control/planner/sql_plan_convert/filter/mod.rs @@ -12,6 +12,7 @@ mod expr_lower; mod serialize; pub(super) use expr_lower::{expr_filter, expr_filter_qualified}; +pub(crate) use serialize::serialize_filters; pub(super) use serialize::{ - encode_scan_filters, filter_to_scan_filters, serialize_filters, serialize_join_post_filters, + encode_scan_filters, filter_to_scan_filters, serialize_join_post_filters, }; diff --git a/nodedb/src/control/planner/sql_plan_convert/scan/mod.rs b/nodedb/src/control/planner/sql_plan_convert/scan/mod.rs index 63a933f22..1a395f43b 100644 --- a/nodedb/src/control/planner/sql_plan_convert/scan/mod.rs +++ b/nodedb/src/control/planner/sql_plan_convert/scan/mod.rs @@ -11,6 +11,8 @@ mod join_cost; mod recursive; mod search; mod spatial; +mod text_score_bound; +mod text_search; mod timeseries; pub(in crate::control::planner::sql_plan_convert) use core::{ @@ -22,9 +24,13 @@ pub(in crate::control::planner::sql_plan_convert) use recursive::{ }; pub(in crate::control::planner::sql_plan_convert) use search::{ convert_hybrid_search, convert_hybrid_search_triple, convert_sparse_search, - convert_text_search, convert_vector_search, + convert_vector_search, }; pub(in crate::control::planner::sql_plan_convert) use spatial::convert_spatial_scan; +pub(in crate::control::planner::sql_plan_convert) use text_score_bound::{ + TextBodyTail, bound_text_body, +}; +pub(in crate::control::planner::sql_plan_convert) use text_search::convert_text_search; pub(in crate::control::planner::sql_plan_convert) use timeseries::{ convert_timeseries_ingest, convert_timeseries_scan, }; diff --git a/nodedb/src/control/planner/sql_plan_convert/scan/search.rs b/nodedb/src/control/planner/sql_plan_convert/scan/search.rs index 7df0e480f..5c5fdffa4 100644 --- a/nodedb/src/control/planner/sql_plan_convert/scan/search.rs +++ b/nodedb/src/control/planner/sql_plan_convert/scan/search.rs @@ -31,6 +31,10 @@ pub(in crate::control::planner::sql_plan_convert) fn convert_vector_search( None => None, }; let ann_options = p.ann_options.to_runtime(); + let filter_bitmap = match p.pk_prefilter { + Some(keys) => Some(pk_prefilter_bitmap(p.ctx, collection_key, keys)?), + None => None, + }; let payload_filters: Vec = p .payload_filters .iter() @@ -46,7 +50,7 @@ pub(in crate::control::planner::sql_plan_convert) fn convert_vector_search( top_k: *p.top_k, ef_search: *p.ef_search, metric: *p.metric, - filter_bitmap: None, + filter_bitmap, field_name: p.field.to_string(), rls_filters: filter_bytes, inline_prefilter_plan, @@ -59,6 +63,24 @@ pub(in crate::control::planner::sql_plan_convert) fn convert_vector_search( }]) } +/// The surrogates `keys` are bound to in `collection`, as the candidate +/// bitmap of a vector search. A key bound to no row names no candidate, so +/// keys that name no row yield an empty bitmap: the search returns nothing. +fn pk_prefilter_bitmap( + ctx: &super::super::convert::ConvertContext, + collection: nodedb_types::CollectionKey<'_>, + keys: &[nodedb_sql::types::SqlValue], +) -> crate::Result { + let mut bitmap = nodedb_types::SurrogateBitmap::new(); + for key in keys { + let pk_bytes = super::super::value::sql_value_to_string(key).into_bytes(); + if let Some(surrogate) = ctx.surrogate_for_existing_pk(collection, &pk_bytes)? { + bitmap.insert(surrogate); + } + } + Ok(bitmap) +} + pub(in crate::control::planner::sql_plan_convert) fn convert_sparse_search( p: SparseSearchParams<'_>, ) -> crate::Result> { @@ -172,111 +194,20 @@ fn build_array_prefilter_plan( )) } -pub(in crate::control::planner::sql_plan_convert) fn convert_text_search( - collection: &str, - query: &nodedb_sql::fts_types::FtsQuery, - top_k: &usize, - score_alias: Option<&str>, - tenant_id: TenantId, - database_id: crate::types::DatabaseId, -) -> crate::Result> { - use nodedb_sql::fts_types::FtsQuery; - - let collection_key = nodedb_types::CollectionKey::from_bare(database_id, collection); - let qualified_collection = nodedb_types::QualifiedCollection::new(database_id, collection); - let vshard = collection_key.vshard(); - - // Phrase queries emit a dedicated PhraseSearch op rather than going - // through the BM25 plain-string path. Score alias is not meaningful - // for phrase search (no per-row score injection), so it is ignored. - if let FtsQuery::Phrase(terms) = query { - let analyzed_terms: Vec = - terms.iter().flat_map(|t| nodedb_fts::analyze(t)).collect(); - if analyzed_terms.is_empty() { - // No searchable terms after analysis — return empty result via - // a standard search that will match nothing. - return Ok(vec![PhysicalTask { - tenant_id, - vshard_id: vshard, - database_id, - plan: PhysicalPlan::Text(TextOp::Search { - collection: qualified_collection.clone(), - query: String::new(), - top_k: *top_k, - fuzzy: false, - prefilter: None, - rls_filters: Vec::new(), - }), - post_set_op: PostSetOp::None, - txn_id: None, - }]); - } - return Ok(vec![PhysicalTask { - tenant_id, - vshard_id: vshard, - database_id, - plan: PhysicalPlan::Text(TextOp::PhraseSearch { - collection: qualified_collection.clone(), - terms: analyzed_terms, - top_k: *top_k, - prefilter: None, - }), - post_set_op: PostSetOp::None, - txn_id: None, - }]); - } - - let query_str = query - .to_plain_string() - .ok_or_else(|| crate::Error::BadRequest { - detail: "unsupported FTS query form; use plain terms, AND/OR combinations, \ - or phrase queries with text_match(field, '\"phrase here\"')" - .into(), - })?; - let fuzzy = query.is_fuzzy(); - - // When a score alias is present the caller wants a full-collection scan - // with BM25 scores injected per row (all rows appear, non-matching rows - // receive a null score). Emit `BM25ScoreScan` for that shape; emit the - // hit-only `Search` for the WHERE `text_match(...)` shape. - let op = if let Some(alias) = score_alias { - TextOp::BM25ScoreScan { - collection: qualified_collection.clone(), - query: query_str, - score_alias: alias.to_string(), - fuzzy, - } - } else { - TextOp::Search { - collection: qualified_collection.clone(), - query: query_str, - top_k: *top_k, - fuzzy, - prefilter: None, - rls_filters: Vec::new(), - } - }; - - Ok(vec![PhysicalTask { - tenant_id, - vshard_id: vshard, - database_id, - plan: PhysicalPlan::Text(op), - post_set_op: PostSetOp::None, - txn_id: None, - }]) -} - pub(in crate::control::planner::sql_plan_convert) fn convert_hybrid_search( p: HybridSearchParams<'_>, ) -> crate::Result> { let HybridSearchParams { collection, + vector_field, query_vector, + text_field, query_text, + filters, top_k, ef_search, vector_weight, + mode, fuzzy, score_alias, tenant_id, @@ -291,10 +222,14 @@ pub(in crate::control::planner::sql_plan_convert) fn convert_hybrid_search( database_id, plan: PhysicalPlan::Text(TextOp::HybridSearch { collection: qualified_collection, + vector_field: vector_field.to_string(), query_vector: query_vector.to_vec(), + text_field: text_field.map(str::to_owned), query_text: query_text.to_string(), + filters: serialize_filters(filters)?, top_k: *top_k, ef_search: *ef_search, + mode, fuzzy: *fuzzy, vector_weight: *vector_weight, filter_bitmap: None, @@ -311,13 +246,17 @@ pub(in crate::control::planner::sql_plan_convert) fn convert_hybrid_search_tripl ) -> crate::Result> { let HybridSearchTripleParams { collection, + vector_field, query_vector, + text_field, query_text, + filters, graph_seed_id, graph_depth, graph_edge_label, top_k, ef_search, + mode, fuzzy, rrf_k, score_alias, @@ -333,13 +272,17 @@ pub(in crate::control::planner::sql_plan_convert) fn convert_hybrid_search_tripl database_id, plan: PhysicalPlan::Text(TextOp::HybridSearchTriple { collection: qualified_collection, + vector_field: vector_field.to_string(), query_vector: query_vector.to_vec(), + text_field: text_field.map(str::to_owned), query_text: query_text.to_string(), + filters: serialize_filters(filters)?, graph_seed_id: graph_seed_id.to_string(), graph_depth: *graph_depth, graph_edge_label: graph_edge_label.clone(), top_k: *top_k, ef_search: *ef_search, + mode, fuzzy: *fuzzy, rrf_k: *rrf_k, filter_bitmap: None, diff --git a/nodedb/src/control/planner/sql_plan_convert/scan/text_score_bound.rs b/nodedb/src/control/planner/sql_plan_convert/scan/text_score_bound.rs new file mode 100644 index 000000000..bc77a455e --- /dev/null +++ b/nodedb/src/control/planner/sql_plan_convert/scan/text_score_bound.rs @@ -0,0 +1,355 @@ +// SPDX-License-Identifier: BUSL-1.1 + +//! LIMIT pushdown into a text body under a relational tail. +//! +//! The tail above a text body applies the query's ORDER BY and LIMIT. When +//! the tail only cuts rows (no filter, DISTINCT, or window) the body can +//! return fewer rows and the tail's answer is unchanged. +//! +//! A `bm25_score(...)` scan with no `text_match` emits every admitted row: +//! +//! - no ORDER BY: any `limit + offset` admitted rows. +//! - ORDER BY one score column: the first `limit + offset` rows in that +//! order, kept in a bounded top-k by the Data Plane. +//! +//! A `text_match` search ranks its hits by BM25 score descending, then by +//! surrogate ascending. ORDER BY one score column descending, whose spec +//! (field, query, mode, fuzzy) equals the search's, sorts by the same score. +//! The tail's sort is stable, so its first `limit + offset` rows are the +//! search's first `limit + offset` hits, and `top_k` takes that bound. A +//! sharded search keeps the answer too: each shard's first rows hold every +//! row of the gathered first rows. A phrase search ranks by match position, +//! not by BM25 score, so it takes no bound. +//! +//! Any other ORDER BY leaves the body unbounded. The tail still applies its +//! own ORDER BY and LIMIT over the rows returned. + +use nodedb_physical::physical_plan::{ + QueryOp, ScoreScanBound, ScoreScanOrder, TextOp, TextScoreSpec, +}; +use nodedb_sql::types::{SortKey, SqlExpr}; +use nodedb_types::text_search::QueryMode; + +use crate::bridge::envelope::PhysicalPlan; + +/// The tail clauses a pushdown must respect. +pub(in crate::control::planner::sql_plan_convert) struct TextBodyTail<'a> { + pub sort_keys: &'a [SortKey], + pub limit: Option, + pub offset: usize, + /// Whether the tail filters, deduplicates, or computes windows: each + /// reads rows past the cut, so no bound is pushed. + pub reads_past_cut: bool, +} + +/// Bound the text body `plan` is (or gathers) by the tail's LIMIT. +pub(in crate::control::planner::sql_plan_convert) fn bound_text_body( + plan: &mut PhysicalPlan, + tail: &TextBodyTail<'_>, +) { + if tail.reads_past_cut { + return; + } + let Some(limit) = tail.limit else { + return; + }; + let rows = limit.saturating_add(tail.offset); + if let PhysicalPlan::Query(QueryOp::Exchange(exchange)) = plan { + // Each shard returns its own first rows: their union holds the + // first rows of the whole collection. + bound_text_body(&mut exchange.child, tail); + return; + } + let PhysicalPlan::Text(op) = plan else { + return; + }; + if let TextOp::BM25ScoreScan { scores, bound, .. } = op { + if let Some(order) = score_scan_order(tail.sort_keys, scores) { + *bound = Some(ScoreScanBound { rows, order }); + } + return; + } + if let TextOp::Search { + field, + query, + top_k, + mode, + fuzzy, + scores, + .. + } = op + { + let search = SearchSpec { + field: field.as_deref(), + query, + mode: *mode, + fuzzy: *fuzzy, + }; + if sorts_by_search_rank(tail.sort_keys, scores, &search) { + *top_k = (*top_k).min(rows); + } + } +} + +/// The query a `TextOp::Search` ranks by. +struct SearchSpec<'a> { + field: Option<&'a str>, + query: &'a str, + mode: QueryMode, + fuzzy: bool, +} + +/// Whether `sort_keys` order rows the way the search ranks them: one key, +/// descending, naming a score column whose spec equals the search's. +fn sorts_by_search_rank( + sort_keys: &[SortKey], + scores: &[TextScoreSpec], + search: &SearchSpec<'_>, +) -> bool { + let [key] = sort_keys else { + return false; + }; + if key.ascending { + return false; + } + let SqlExpr::Column { name, .. } = &key.expr else { + return false; + }; + // The last column of an alias is the value a row carries under it. + scores + .iter() + .rev() + .find(|s| &s.alias == name) + .is_some_and(|s| { + s.field.as_deref() == search.field + && s.query == search.query + && s.mode == search.mode + && s.fuzzy == search.fuzzy + }) +} + +/// The bounded order of a score scan under `sort_keys`. `None` when the +/// keys leave the scan unbounded. `Some(None)` bounds it with no order. +fn score_scan_order( + sort_keys: &[SortKey], + scores: &[TextScoreSpec], +) -> Option> { + match sort_keys { + [] => Some(None), + [key] => match &key.expr { + SqlExpr::Column { name, .. } if scores.iter().any(|s| &s.alias == name) => { + Some(Some(ScoreScanOrder { + alias: name.clone(), + ascending: key.ascending, + nulls_first: key.nulls_first, + })) + } + _ => None, + }, + _ => None, + } +} + +#[cfg(test)] +mod tests { + use nodedb_types::{DatabaseId, QualifiedCollection}; + + use super::*; + + fn scan() -> PhysicalPlan { + PhysicalPlan::Text(TextOp::BM25ScoreScan { + collection: QualifiedCollection::new(DatabaseId::DEFAULT, "c"), + filters: Vec::new(), + rls_filters: Vec::new(), + scores: vec![TextScoreSpec { + field: None, + query: "q".into(), + mode: nodedb_types::text_search::QueryMode::And, + fuzzy: false, + alias: "s".into(), + }], + bound: None, + }) + } + + fn key(name: &str) -> SortKey { + SortKey { + expr: SqlExpr::Column { + table: None, + name: name.into(), + }, + ascending: false, + nulls_first: false, + } + } + + fn bound_of(plan: &PhysicalPlan) -> Option { + match plan { + PhysicalPlan::Text(TextOp::BM25ScoreScan { bound, .. }) => bound.clone(), + other => panic!("plan shape changed: {other:?}"), + } + } + + fn tail(sort_keys: &[SortKey], reads_past_cut: bool) -> TextBodyTail<'_> { + TextBodyTail { + sort_keys, + limit: Some(10), + offset: 5, + reads_past_cut, + } + } + + #[test] + fn an_unordered_limit_bounds_the_scan() { + let mut plan = scan(); + bound_text_body(&mut plan, &tail(&[], false)); + assert_eq!( + bound_of(&plan), + Some(ScoreScanBound { + rows: 15, + order: None + }) + ); + } + + #[test] + fn a_score_order_is_kept_in_the_bound() { + let mut plan = scan(); + bound_text_body(&mut plan, &tail(&[key("s")], false)); + let bound = bound_of(&plan).expect("bounded"); + assert_eq!(bound.rows, 15); + assert_eq!( + bound.order, + Some(ScoreScanOrder { + alias: "s".into(), + ascending: false, + nulls_first: false, + }) + ); + } + + #[test] + fn a_non_score_order_or_a_filtering_tail_leaves_the_scan_unbounded() { + let mut plan = scan(); + bound_text_body(&mut plan, &tail(&[key("id")], false)); + assert_eq!(bound_of(&plan), None); + bound_text_body(&mut plan, &tail(&[key("s"), key("id")], false)); + assert_eq!(bound_of(&plan), None); + bound_text_body(&mut plan, &tail(&[], true)); + assert_eq!(bound_of(&plan), None); + } + + fn spec(query: &str, mode: QueryMode, fuzzy: bool) -> TextScoreSpec { + TextScoreSpec { + field: Some("body".into()), + query: query.into(), + mode, + fuzzy, + alias: "score".into(), + } + } + + fn search(score: TextScoreSpec) -> PhysicalPlan { + PhysicalPlan::Text(TextOp::Search { + collection: QualifiedCollection::new(DatabaseId::DEFAULT, "c"), + field: Some("body".into()), + query: "rust db".into(), + top_k: usize::MAX, + mode: QueryMode::And, + fuzzy: false, + prefilter: None, + filters: Vec::new(), + rls_filters: Vec::new(), + scores: vec![score], + }) + } + + fn top_k_of(plan: &PhysicalPlan) -> usize { + match plan { + PhysicalPlan::Text(TextOp::Search { top_k, .. }) => *top_k, + other => panic!("plan shape changed: {other:?}"), + } + } + + #[test] + fn a_descending_sort_by_the_search_score_bounds_the_search() { + let mut plan = search(spec("rust db", QueryMode::And, false)); + bound_text_body(&mut plan, &tail(&[key("score")], false)); + assert_eq!(top_k_of(&plan), 15); + } + + #[test] + fn a_sort_the_search_does_not_rank_by_leaves_it_unbounded() { + let ascending = SortKey { + ascending: true, + ..key("score") + }; + let cases: Vec<(PhysicalPlan, Vec, bool)> = vec![ + // Ascending order. + ( + search(spec("rust db", QueryMode::And, false)), + vec![ascending], + false, + ), + // A score of another query, mode, or fuzzy setting. + ( + search(spec("rust", QueryMode::And, false)), + vec![key("score")], + false, + ), + ( + search(spec("rust db", QueryMode::Or, false)), + vec![key("score")], + false, + ), + ( + search(spec("rust db", QueryMode::And, true)), + vec![key("score")], + false, + ), + // A non-score key, a second key, no key, a filtering tail. + ( + search(spec("rust db", QueryMode::And, false)), + vec![key("id")], + false, + ), + ( + search(spec("rust db", QueryMode::And, false)), + vec![key("score"), key("id")], + false, + ), + ( + search(spec("rust db", QueryMode::And, false)), + vec![], + false, + ), + ( + search(spec("rust db", QueryMode::And, false)), + vec![key("score")], + true, + ), + ]; + for (mut plan, keys, reads_past_cut) in cases { + bound_text_body(&mut plan, &tail(&keys, reads_past_cut)); + assert_eq!(top_k_of(&plan), usize::MAX, "keys {keys:?}"); + } + } + + #[test] + fn a_sharded_search_bounds_each_shard() { + let child = search(spec("rust db", QueryMode::And, false)); + let mut plan = PhysicalPlan::Query(QueryOp::Exchange( + nodedb_physical::physical_plan::ExchangeOp { + child: Box::new(child), + mode: nodedb_physical::physical_plan::ExchangeMode::Gather { + as_aggregate: false, + }, + }, + )); + bound_text_body(&mut plan, &tail(&[key("score")], false)); + let PhysicalPlan::Query(QueryOp::Exchange(exchange)) = &plan else { + panic!("plan shape changed: {plan:?}"); + }; + assert_eq!(top_k_of(&exchange.child), 15); + } +} diff --git a/nodedb/src/control/planner/sql_plan_convert/scan/text_search.rs b/nodedb/src/control/planner/sql_plan_convert/scan/text_search.rs new file mode 100644 index 000000000..016940949 --- /dev/null +++ b/nodedb/src/control/planner/sql_plan_convert/scan/text_search.rs @@ -0,0 +1,96 @@ +// SPDX-License-Identifier: BUSL-1.1 + +//! `SqlPlan::TextSearch` → `TextOp` conversion. + +use nodedb_physical::physical_plan::{TextOp, TextScoreSpec}; +use nodedb_physical::physical_task::{PhysicalTask, PostSetOp}; +use nodedb_sql::fts_types::FtsQuery; +use nodedb_sql::types::TextSearchShape; + +use super::super::filter::serialize_filters; +use super::super::scan_params::TextSearchConvertParams; +use crate::bridge::envelope::PhysicalPlan; + +/// Lower a text search. A `Match` becomes `Search` (or `PhraseSearch` for a +/// quoted phrase); a `ScoreScan` becomes `BM25ScoreScan`. A `Match` with no +/// top-k bound returns every match (`usize::MAX`). +pub(in crate::control::planner::sql_plan_convert) fn convert_text_search( + p: TextSearchConvertParams<'_>, +) -> crate::Result> { + let collection_key = nodedb_types::CollectionKey::from_bare(p.database_id, p.collection); + let collection = nodedb_types::QualifiedCollection::new(p.database_id, p.collection); + let filters = serialize_filters(p.filters)?; + let scores: Vec = p + .scores + .iter() + .map(|s| TextScoreSpec { + field: s.field.clone(), + query: s.query.clone(), + mode: s.mode, + fuzzy: s.fuzzy, + alias: s.alias.clone(), + }) + .collect(); + + let op = match p.shape { + TextSearchShape::ScoreScan => TextOp::BM25ScoreScan { + collection, + filters, + rls_filters: Vec::new(), + scores, + bound: None, + }, + // The phrase words go as written. The Data Plane analyzes them with + // the collection's analyzer, the analysis its indexed text took. + TextSearchShape::Match { + field, + query: FtsQuery::Phrase(terms), + top_k, + .. + } => TextOp::PhraseSearch { + collection, + field: field.clone(), + terms: terms.clone(), + top_k: top_k.unwrap_or(usize::MAX), + prefilter: None, + filters, + rls_filters: Vec::new(), + scores, + }, + TextSearchShape::Match { + field, + query, + mode, + top_k, + } => { + let query_str = query + .to_plain_string() + .ok_or_else(|| crate::Error::BadRequest { + detail: "unsupported FTS query form; use plain terms, AND/OR combinations, \ + or phrase queries with text_match(field, '\"phrase here\"')" + .into(), + })?; + TextOp::Search { + collection, + field: field.clone(), + query: query_str, + top_k: top_k.unwrap_or(usize::MAX), + mode: *mode, + fuzzy: query.is_fuzzy(), + prefilter: None, + filters, + rls_filters: Vec::new(), + scores, + } + } + }; + + Ok(vec![PhysicalTask { + tenant_id: p.tenant_id, + vshard_id: collection_key.vshard(), + database_id: p.database_id, + plan: PhysicalPlan::Text(op), + post_set_op: PostSetOp::None, + txn_id: None, + }]) +} diff --git a/nodedb/src/control/planner/sql_plan_convert/scan_params.rs b/nodedb/src/control/planner/sql_plan_convert/scan_params.rs index 8c7523817..2a44d2cd1 100644 --- a/nodedb/src/control/planner/sql_plan_convert/scan_params.rs +++ b/nodedb/src/control/planner/sql_plan_convert/scan_params.rs @@ -106,6 +106,9 @@ pub(super) struct VectorSearchParams<'a> { /// Translated to `nodedb_types::PayloadAtom` and emitted as /// `VectorOp::Search::payload_filters`. pub payload_filters: &'a [nodedb_sql::types::SqlPayloadAtom], + /// Primary keys the search ranks among, lowered to + /// `VectorOp::Search::filter_bitmap`. `None`: every row. + pub pk_prefilter: Option<&'a [nodedb_sql::types::SqlValue]>, } /// Parameters for `convert_sparse_search`. @@ -119,14 +122,31 @@ pub(super) struct SparseSearchParams<'a> { pub database_id: crate::types::DatabaseId, } +/// Parameters for `convert_text_search`. +pub(super) struct TextSearchConvertParams<'a> { + pub collection: &'a str, + pub shape: &'a nodedb_sql::types::TextSearchShape, + /// Residual WHERE predicates, applied before ranking. + pub filters: &'a [Filter], + pub scores: &'a [nodedb_sql::types::TextScoreColumn], + pub tenant_id: TenantId, + pub database_id: crate::types::DatabaseId, +} + /// Parameters for `convert_hybrid_search`. pub(super) struct HybridSearchParams<'a> { pub collection: &'a str, + pub vector_field: &'a str, pub query_vector: &'a [f32], + /// `None` reads the whole-document index. + pub text_field: Option<&'a str>, pub query_text: &'a str, + /// Residual WHERE predicates, applied to both legs before fusion. + pub filters: &'a [Filter], pub top_k: &'a usize, pub ef_search: &'a usize, pub vector_weight: &'a f32, + pub mode: nodedb_types::text_search::QueryMode, pub fuzzy: &'a bool, /// SELECT-list alias for the RRF score column. Forwarded to /// `TextOp::HybridSearch.score_alias` so the executor renames the @@ -139,13 +159,19 @@ pub(super) struct HybridSearchParams<'a> { /// Parameters for `convert_hybrid_search_triple`. pub(super) struct HybridSearchTripleParams<'a> { pub collection: &'a str, + pub vector_field: &'a str, pub query_vector: &'a [f32], + /// `None` reads the whole-document index. + pub text_field: Option<&'a str>, pub query_text: &'a str, + /// Residual WHERE predicates, applied to every leg before fusion. + pub filters: &'a [Filter], pub graph_seed_id: &'a str, pub graph_depth: &'a usize, pub graph_edge_label: &'a Option, pub top_k: &'a usize, pub ef_search: &'a usize, + pub mode: nodedb_types::text_search::QueryMode, pub fuzzy: &'a bool, pub rrf_k: &'a (f64, f64, f64), pub score_alias: Option<&'a str>, diff --git a/nodedb/src/control/planner/sql_plan_convert/set_ops.rs b/nodedb/src/control/planner/sql_plan_convert/set_ops.rs index d498969d2..8e04e8056 100644 --- a/nodedb/src/control/planner/sql_plan_convert/set_ops.rs +++ b/nodedb/src/control/planner/sql_plan_convert/set_ops.rs @@ -316,7 +316,17 @@ pub(super) fn convert_subquery( } = args; // The body is ONE relation, already gathered when sharded. - let child = convert_body_to_single_plan(input, tenant_id, ctx)?; + let mut child = convert_body_to_single_plan(input, tenant_id, ctx)?; + // A text body takes the LIMIT when the tail only cuts rows. + super::scan::bound_text_body( + &mut child, + &super::scan::TextBodyTail { + sort_keys, + limit, + offset, + reads_past_cut: !filters.is_empty() || distinct || !window_functions.is_empty(), + }, + ); // A join / lateral body emits ONE merged document per output row whose // columns keep their table prefix (`a.attnum`), which is why the response @@ -699,4 +709,74 @@ mod tests { assert_eq!(tasks.len(), 1); } + + /// The text search under the tail, looked up through a gather. + fn search_top_k(plan: &PhysicalPlan) -> usize { + match plan { + PhysicalPlan::Query(QueryOp::Exchange(exchange)) => search_top_k(&exchange.child), + PhysicalPlan::Text(TextOp::Search { top_k, .. }) => *top_k, + other => panic!("expected a text search body, got {other:?}"), + } + } + + /// `WHERE text_match(body, q) ORDER BY bm25_score(body, q) DESC LIMIT 3 + /// OFFSET 2`: the search under the tail ranks only the 5 rows the tail + /// can return. + #[test] + fn a_text_match_sorted_by_its_own_score_takes_the_tail_limit() { + use nodedb_sql::fts_types::FtsQuery; + use nodedb_sql::types::{TextScoreColumn, TextSearchPlan, TextSearchShape}; + use nodedb_types::text_search::QueryMode; + + let score = |query: &str| TextScoreColumn { + field: Some("body".into()), + query: query.into(), + mode: QueryMode::And, + fuzzy: false, + alias: "score".into(), + }; + let wrapped = |score_query: &str, ascending: bool| SqlPlan::Subquery { + input: Box::new(SqlPlan::TextSearch(TextSearchPlan { + collection: "docs".into(), + shape: TextSearchShape::Match { + field: Some("body".into()), + query: FtsQuery::Plain { + text: "rust db".into(), + fuzzy: false, + }, + mode: QueryMode::And, + top_k: None, + }, + filters: Vec::new(), + scores: vec![score(score_query)], + projection: Vec::new(), + })), + filters: Vec::new(), + projection: Vec::new(), + window_functions: Vec::new(), + sort_keys: vec![SortKey { + expr: SqlExpr::Column { + table: None, + name: "score".into(), + }, + ascending, + nulls_first: ascending, + }], + offset: 2, + distinct: false, + limit: Some(3), + }; + let body_top_k = |plan: &SqlPlan| { + let tasks = convert_one(plan, TenantId::new(1), &bare_ctx()).expect("converts"); + match &tasks[0].plan { + PhysicalPlan::Query(QueryOp::PostProcess { input, .. }) => search_top_k(input), + other => panic!("expected a post-processing tail, got {other:?}"), + } + }; + + assert_eq!(body_top_k(&wrapped("rust db", false)), 5); + // An ascending sort, or a score of another query, keeps every match. + assert_eq!(body_top_k(&wrapped("rust db", true)), usize::MAX); + assert_eq!(body_top_k(&wrapped("rust", false)), usize::MAX); + } } diff --git a/nodedb/src/control/planner/sql_plan_convert/visitor/arms_scan_search.rs b/nodedb/src/control/planner/sql_plan_convert/visitor/arms_scan_search.rs index 05db0ec62..e0c397185 100644 --- a/nodedb/src/control/planner/sql_plan_convert/visitor/arms_scan_search.rs +++ b/nodedb/src/control/planner/sql_plan_convert/visitor/arms_scan_search.rs @@ -71,6 +71,7 @@ macro_rules! impl_scan_search_arms_for_convert_visitor { ann_options, skip_payload_fetch, payload_filters, + pk_prefilter, } = args; super::super::scan::convert_vector_search( super::super::scan_params::VectorSearchParams { @@ -87,6 +88,7 @@ macro_rules! impl_scan_search_arms_for_convert_visitor { ctx: self.ctx, skip_payload_fetch, payload_filters, + pk_prefilter, }, ) } @@ -112,19 +114,23 @@ macro_rules! impl_scan_search_arms_for_convert_visitor { fn text_search( &mut self, - collection: &str, - query: &nodedb_sql::fts_types::FtsQuery, - top_k: usize, - _filters: &[nodedb_sql::types::filter::Filter], - score_alias: Option<&str>, + args: nodedb_sql::TextSearchVisitArgs<'_>, ) -> crate::Result> { - super::super::scan::convert_text_search( + let nodedb_sql::TextSearchVisitArgs { collection, - query, - &top_k, - score_alias, - self.tenant_id, - self.ctx.database_id, + shape, + filters, + scores, + } = args; + super::super::scan::convert_text_search( + super::super::scan_params::TextSearchConvertParams { + collection, + shape, + filters, + scores, + tenant_id: self.tenant_id, + database_id: self.ctx.database_id, + }, ) } @@ -134,22 +140,30 @@ macro_rules! impl_scan_search_arms_for_convert_visitor { ) -> crate::Result> { let nodedb_sql::HybridSearchVisitArgs { collection, + vector_field, query_vector, + text_field, query_text, + filters, top_k, ef_search, vector_weight, + mode, fuzzy, score_alias, } = args; super::super::scan::convert_hybrid_search( super::super::scan_params::HybridSearchParams { collection, + vector_field, query_vector, + text_field, query_text, + filters, top_k: &top_k, ef_search: &ef_search, vector_weight: &vector_weight, + mode, fuzzy: &fuzzy, score_alias, tenant_id: self.tenant_id, @@ -164,13 +178,17 @@ macro_rules! impl_scan_search_arms_for_convert_visitor { ) -> crate::Result> { let nodedb_sql::HybridSearchTripleVisitArgs { collection, + vector_field, query_vector, + text_field, query_text, + filters, graph_seed_id, graph_depth, graph_edge_label, top_k, ef_search, + mode, fuzzy, rrf_k, score_alias, @@ -179,13 +197,17 @@ macro_rules! impl_scan_search_arms_for_convert_visitor { super::super::scan::convert_hybrid_search_triple( super::super::scan_params::HybridSearchTripleParams { collection, + vector_field, query_vector, + text_field, query_text, + filters, graph_seed_id, graph_depth: &graph_depth, graph_edge_label: &graph_edge_label_owned, top_k: &top_k, ef_search: &ef_search, + mode, fuzzy: &fuzzy, rrf_k: &rrf_k, score_alias, diff --git a/nodedb/src/control/server/native/dispatch/plan_builder/text.rs b/nodedb/src/control/server/native/dispatch/plan_builder/text.rs index 3f95c7255..0e9dd7ee1 100644 --- a/nodedb/src/control/server/native/dispatch/plan_builder/text.rs +++ b/nodedb/src/control/server/native/dispatch/plan_builder/text.rs @@ -4,11 +4,23 @@ use nodedb_types::QualifiedCollection; use nodedb_types::protocol::TextFields; +use nodedb_types::text_search::TextSearchParams; use crate::bridge::envelope::PhysicalPlan; use crate::control::server::native::dispatch::DispatchCtx; use nodedb_physical::physical_plan::TextOp; +/// The query options a request names. An absent field takes its +/// [`TextSearchParams::default`] value, the default of SQL `text_match` and +/// of every `text_search` client. +fn search_params(fields: &TextFields) -> TextSearchParams { + let defaults = TextSearchParams::default(); + TextSearchParams { + mode: fields.text_mode.unwrap_or(defaults.mode), + fuzzy: fields.fuzzy.unwrap_or(defaults.fuzzy), + } +} + pub(crate) async fn build_search( ctx: &DispatchCtx<'_>, fields: &TextFields, @@ -21,15 +33,20 @@ pub(crate) async fn build_search( detail: "missing 'query_text'".to_string(), })?; let top_k = fields.top_k.unwrap_or(10) as usize; - let fuzzy = fields.fuzzy.unwrap_or(false); + let params = search_params(fields); Ok(PhysicalPlan::Text(TextOp::Search { collection: QualifiedCollection::new(ctx.database_id(), collection), + // An absent or empty field reads the whole-document index. + field: fields.field.clone().filter(|f| !f.is_empty()), query: query_text.to_string(), top_k, - fuzzy, + mode: params.mode, + fuzzy: params.fuzzy, prefilter: None, + filters: Vec::new(), rls_filters: Vec::new(), + scores: Vec::new(), })) } @@ -53,15 +70,22 @@ pub(crate) async fn build_hybrid_search( let top_k = fields.top_k.unwrap_or(10) as usize; let vector_weight = fields.vector_weight.unwrap_or(0.5) as f32; let ef_search = fields.ef_search.unwrap_or(0) as usize; - let fuzzy = fields.fuzzy.unwrap_or(false); + let params = search_params(fields); Ok(PhysicalPlan::Text(TextOp::HybridSearch { collection: QualifiedCollection::new(ctx.database_id(), collection), + // `field` names the vector column; absent keys the default index. + vector_field: fields.field.clone().unwrap_or_default(), query_vector: query_vector.clone(), + // The native frame names no text column: the text leg reads the + // whole-document index. + text_field: None, query_text: query_text.clone(), + filters: Vec::new(), top_k, ef_search, - fuzzy, + mode: params.mode, + fuzzy: params.fuzzy, vector_weight, filter_bitmap: None, rls_filters: Vec::new(), @@ -70,3 +94,34 @@ pub(crate) async fn build_hybrid_search( score_alias: None, })) } + +#[cfg(test)] +mod tests { + use nodedb_types::text_search::QueryMode; + + use super::*; + + #[test] + fn absent_options_take_the_trait_default() { + assert_eq!( + search_params(&TextFields::default()), + TextSearchParams::default() + ); + } + + #[test] + fn named_options_reach_the_params() { + let fields = TextFields { + text_mode: Some(QueryMode::And), + fuzzy: Some(true), + ..Default::default() + }; + assert_eq!( + search_params(&fields), + TextSearchParams { + mode: QueryMode::And, + fuzzy: true, + } + ); + } +} diff --git a/nodedb/src/control/server/pgwire/types/mod.rs b/nodedb/src/control/server/pgwire/types/mod.rs index b7aa1ef45..7d62452c8 100644 --- a/nodedb/src/control/server/pgwire/types/mod.rs +++ b/nodedb/src/control/server/pgwire/types/mod.rs @@ -10,16 +10,13 @@ pub mod field; pub mod numeric_sqlstate; pub mod parse; pub mod privilege; +pub mod wire_type; pub use error_map::{ dml_fold_error_to_pg, error_to_pg, error_to_pg_in_context, error_to_sqlstate, notice_warning, response_status_to_sqlstate, shape_error_to_pg, sqlstate_error, }; -pub use field::{ - bool_field, bytea_field, float4_array_field, float4_field, float8_array_field, float8_field, - int2_field, int4_field, int8_field, json_field, jsonb_field, text_field, timestamp_field, - timestamptz_field, type_name_to_pgwire, varchar_field, -}; +pub use field::{text_field, type_name_to_pgwire}; pub use parse::parse_role; pub use privilege::{ require_cluster_admin, require_database_owner, require_database_owner_or_higher, diff --git a/nodedb/src/control/server/response_translate/dispatch.rs b/nodedb/src/control/server/response_translate/dispatch.rs index 5b93f924d..c8695ec7b 100644 --- a/nodedb/src/control/server/response_translate/dispatch.rs +++ b/nodedb/src/control/server/response_translate/dispatch.rs @@ -65,7 +65,11 @@ pub fn translate_search_response( &[], *top_k, ), - PhysicalPlan::Text(TextOp::Search { collection, .. }) => translate_text_search_payload( + PhysicalPlan::Text( + TextOp::Search { collection, .. } + | TextOp::PhraseSearch { collection, .. } + | TextOp::BM25ScoreScan { collection, .. }, + ) => translate_text_search_payload( payload, state, database_id, diff --git a/nodedb/src/control/server/shared/ddl/neutral/column_default.rs b/nodedb/src/control/server/shared/ddl/neutral/column_default.rs index 179145088..2431ca28a 100644 --- a/nodedb/src/control/server/shared/ddl/neutral/column_default.rs +++ b/nodedb/src/control/server/shared/ddl/neutral/column_default.rs @@ -119,7 +119,9 @@ fn clause_error(clause: &str, owner: &str, error: &SqlError) -> DdlError { SqlError::TypeMismatch { .. } => sqlstate::DATATYPE_MISMATCH, SqlError::IntegerOutOfRange { .. } | SqlError::FloatOutOfRange { .. } - | SqlError::ConstantOverflow { .. } => sqlstate::NUMERIC_VALUE_OUT_OF_RANGE, + | SqlError::DecimalOutOfRange { .. } + | SqlError::ConstantOverflow { .. } + | SqlError::NumericLiteralOutOfRange { .. } => sqlstate::NUMERIC_VALUE_OUT_OF_RANGE, SqlError::DivisionByZero => sqlstate::DIVISION_BY_ZERO, SqlError::DataException { .. } => sqlstate::DATA_EXCEPTION, SqlError::InvalidLimitValue { .. } => sqlstate::INVALID_LIMIT_VALUE, @@ -128,6 +130,7 @@ fn clause_error(clause: &str, owner: &str, error: &SqlError) -> DdlError { } SqlError::UnknownColumn { .. } => sqlstate::UNDEFINED_COLUMN, SqlError::AmbiguousColumn { .. } => sqlstate::AMBIGUOUS_COLUMN, + SqlError::TextColumn { fault, .. } => fault.sqlstate(), SqlError::UndefinedObject { .. } => sqlstate::UNDEFINED_OBJECT, SqlError::ObjectNotInPrerequisiteState { .. } => sqlstate::OBJECT_NOT_IN_PREREQUISITE_STATE, SqlError::SequencePerRowUnsupported { .. } diff --git a/nodedb/src/control/server/shared/write_admission/predicate/txn_buffering/classify.rs b/nodedb/src/control/server/shared/write_admission/predicate/txn_buffering/classify.rs index 1aa846e8a..e46c854a0 100644 --- a/nodedb/src/control/server/shared/write_admission/predicate/txn_buffering/classify.rs +++ b/nodedb/src/control/server/shared/write_admission/predicate/txn_buffering/classify.rs @@ -598,6 +598,9 @@ mod tests { conflict_policy: None, timeseries: None, vector_primary: None, + vector_fields: Vec::new(), + declared_columns: Vec::new(), + declared_key: None, }), PhysicalPlan::Document(DocumentOp::IndexLookup { collection: QualifiedCollection::new(DatabaseId::DEFAULT, "c"), @@ -1080,7 +1083,7 @@ mod tests { }), PhysicalPlan::Graph(GraphOp::Hop { start_nodes: Vec::new(), - edge_label: None, + edge_labels: Vec::new(), direction: Direction::Out, depth: 0, options: GraphTraversalOptions::default(), @@ -1090,23 +1093,25 @@ mod tests { }), PhysicalPlan::Graph(GraphOp::Neighbors { node_id: "n".into(), - edge_label: None, + edge_labels: Vec::new(), direction: Direction::Out, rls_filters: Vec::new(), collection: None, }), PhysicalPlan::Graph(GraphOp::NeighborsMulti { node_ids: Vec::new(), - edge_label: None, + edge_labels: Vec::new(), direction: Direction::Out, max_results: 0, rls_filters: Vec::new(), collection: None, + edge_predicate: Vec::new(), + with_properties: false, }), PhysicalPlan::Graph(GraphOp::Path { src: "a".into(), dst: "b".into(), - edge_label: None, + edge_labels: Vec::new(), max_depth: 0, options: GraphTraversalOptions::default(), rls_filters: Vec::new(), @@ -1115,7 +1120,7 @@ mod tests { }), PhysicalPlan::Graph(GraphOp::Subgraph { start_nodes: Vec::new(), - edge_label: None, + edge_labels: Vec::new(), depth: 0, options: GraphTraversalOptions::default(), rls_filters: Vec::new(), @@ -1373,7 +1378,7 @@ mod tests { source_key: Vec::new(), dest_key: Vec::new(), field: "f".into(), - amount: 0.0, + amount: nodedb_physical::physical_plan::TransferAmount::Float(0.0), debit_surrogate: Surrogate::new(1), credit_surrogate: Surrogate::new(2), rls_write_check: nodedb_types::RlsWriteCheck::NoPolicyApplies, @@ -1550,30 +1555,43 @@ mod tests { let plans = vec![ PhysicalPlan::Text(TextOp::Search { collection: QualifiedCollection::new(DatabaseId::DEFAULT, "c"), + field: None, query: "q".into(), top_k: 0, + mode: nodedb_types::text_search::QueryMode::And, fuzzy: false, prefilter: None, + filters: Vec::new(), rls_filters: Vec::new(), + scores: Vec::new(), }), PhysicalPlan::Text(TextOp::BM25ScoreScan { collection: QualifiedCollection::new(DatabaseId::DEFAULT, "c"), - query: "q".into(), - score_alias: "s".into(), - fuzzy: false, + filters: Vec::new(), + rls_filters: Vec::new(), + scores: Vec::new(), + bound: None, }), PhysicalPlan::Text(TextOp::PhraseSearch { collection: QualifiedCollection::new(DatabaseId::DEFAULT, "c"), + field: None, terms: Vec::new(), top_k: 0, prefilter: None, + filters: Vec::new(), + rls_filters: Vec::new(), + scores: Vec::new(), }), PhysicalPlan::Text(TextOp::HybridSearch { collection: QualifiedCollection::new(DatabaseId::DEFAULT, "c"), + vector_field: String::new(), query_vector: Vec::new(), + text_field: None, query_text: "q".into(), + filters: Vec::new(), top_k: 0, ef_search: 0, + mode: nodedb_types::text_search::QueryMode::And, fuzzy: false, vector_weight: 0.0, filter_bitmap: None, @@ -1583,7 +1601,7 @@ mod tests { PhysicalPlan::Text(TextOp::FtsIndexDoc { collection: QualifiedCollection::new(DatabaseId::DEFAULT, "c"), surrogate: Surrogate::new(1), - text: "t".into(), + fields: vec![("body".into(), "t".into())], provenance: None, }), PhysicalPlan::Text(TextOp::FtsDeleteDoc { @@ -1593,13 +1611,17 @@ mod tests { }), PhysicalPlan::Text(TextOp::HybridSearchTriple { collection: QualifiedCollection::new(DatabaseId::DEFAULT, "c"), + vector_field: String::new(), query_vector: Vec::new(), + text_field: None, query_text: "q".into(), + filters: Vec::new(), graph_seed_id: "n".into(), graph_depth: 0, graph_edge_label: None, top_k: 0, ef_search: 0, + mode: nodedb_types::text_search::QueryMode::And, fuzzy: false, rrf_k: (0.0, 0.0, 0.0), filter_bitmap: None, @@ -1946,12 +1968,11 @@ mod tests { }), PhysicalPlan::Meta(MetaOp::MarkSavepoint { txn_id: TxnId::new(1), + savepoint: 1, }), PhysicalPlan::Meta(MetaOp::RollbackToSavepoint { txn_id: TxnId::new(1), - value_marker: 0, - graph_marker: 0, - array_marker: 0, + savepoint: 1, }), PhysicalPlan::Meta(MetaOp::RecordCalvinWriteVersions { tenant_id: tenant(), diff --git a/nodedb/src/data/executor/dispatch/text.rs b/nodedb/src/data/executor/dispatch/text.rs index 49c166278..07d4cacad 100644 --- a/nodedb/src/data/executor/dispatch/text.rs +++ b/nodedb/src/data/executor/dispatch/text.rs @@ -8,6 +8,7 @@ use nodedb_physical::physical_plan::TextOp; use crate::data::executor::core_loop::CoreLoop; use crate::data::executor::handlers::text_search::TextSearchParams; use crate::data::executor::handlers::text_search_hybrid::HybridSearchParams; +use crate::data::executor::handlers::text_search_scan::{PhraseSearchParams, ScoreScanParams}; use crate::data::executor::handlers::text_search_triple::HybridSearchTripleParams; use crate::data::executor::task::ExecutionTask; @@ -17,58 +18,84 @@ impl CoreLoop { match op { TextOp::Search { collection, + field, query, top_k, + mode, fuzzy, prefilter, + filters, rls_filters, + scores, } => self.execute_text_search( task, TextSearchParams { tid, collection: collection.as_str(), + field: field.as_deref(), query, top_k: *top_k, + mode: *mode, fuzzy: *fuzzy, prefilter: prefilter.as_ref(), + filters, rls_filters, + scores, }, ), TextOp::BM25ScoreScan { collection, - query, - score_alias, - fuzzy, + filters, + rls_filters, + scores, + bound, } => self.execute_bm25_score_scan( task, - tid, - collection.as_str(), - query, - score_alias, - *fuzzy, + ScoreScanParams { + tid, + collection: collection.as_str(), + filters, + rls_filters, + scores, + bound: bound.as_ref(), + }, ), TextOp::PhraseSearch { collection, + field, terms, top_k, prefilter, + filters, + rls_filters, + scores, } => self.execute_phrase_search( task, - tid, - collection.as_str(), - terms, - *top_k, - prefilter.as_ref(), + PhraseSearchParams { + tid, + collection: collection.as_str(), + field: field.as_deref(), + terms, + top_k: *top_k, + prefilter: prefilter.as_ref(), + filters, + rls_filters, + scores, + }, ), TextOp::HybridSearch { collection, + vector_field, query_vector, + text_field, query_text, + filters, top_k, ef_search, + mode, fuzzy, vector_weight, filter_bitmap, @@ -79,10 +106,14 @@ impl CoreLoop { HybridSearchParams { tid, collection: collection.as_str(), + vector_field, query_vector, + text_field: text_field.as_deref(), query_text, + filters, top_k: *top_k, ef_search: *ef_search, + mode: *mode, fuzzy: *fuzzy, vector_weight: *vector_weight, filter_bitmap: filter_bitmap.as_ref(), @@ -94,14 +125,14 @@ impl CoreLoop { TextOp::FtsIndexDoc { collection, surrogate, - text, + fields, provenance, } => self.execute_fts_index_doc( task, tid, collection.as_str(), *surrogate, - text, + fields, provenance.as_ref(), ), @@ -119,13 +150,17 @@ impl CoreLoop { TextOp::HybridSearchTriple { collection, + vector_field, query_vector, + text_field, query_text, + filters, graph_seed_id, graph_depth, graph_edge_label, top_k, ef_search, + mode, fuzzy, rrf_k, filter_bitmap, @@ -136,13 +171,17 @@ impl CoreLoop { HybridSearchTripleParams { tid, collection: collection.as_str(), + vector_field, query_vector, + text_field: text_field.as_deref(), query_text, + filters, graph_seed_id, graph_depth: *graph_depth, graph_edge_label: graph_edge_label.as_deref(), top_k: *top_k, ef_search: *ef_search, + mode: *mode, fuzzy: *fuzzy, rrf_k: *rrf_k, filter_bitmap: filter_bitmap.as_ref(), diff --git a/nodedb/src/data/executor/handlers/hybrid_key.rs b/nodedb/src/data/executor/handlers/hybrid_key.rs index 83766f6cc..64c760f39 100644 --- a/nodedb/src/data/executor/handlers/hybrid_key.rs +++ b/nodedb/src/data/executor/handlers/hybrid_key.rs @@ -42,6 +42,29 @@ impl HybridFusionKey { } } +/// Keep the hits of one ranked leg whose row `admitted` holds, and renumber +/// their ranks from 0 in the order kept. A headless hit has no row, so it is +/// never admitted. `None` admits every hit. +/// +/// Each leg is restricted before fusion, so the fused top-k counts only +/// admitted rows and needs no check after fusion. +pub(in crate::data::executor) fn retain_admitted( + ranked: &mut Vec>, + admitted: Option<&nodedb_types::SurrogateBitmap>, +) { + let Some(admitted) = admitted else { + return; + }; + ranked.retain(|r| { + r.document_id + .storage_key() + .is_some_and(|key| admitted.contains(key.surrogate())) + }); + for (rank, r) in ranked.iter_mut().enumerate() { + r.rank = rank; + } +} + impl fmt::Display for HybridFusionKey { fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { match self { diff --git a/nodedb/src/data/executor/handlers/hybrid_overlay.rs b/nodedb/src/data/executor/handlers/hybrid_overlay.rs index 2515f9656..8f0b62279 100644 --- a/nodedb/src/data/executor/handlers/hybrid_overlay.rs +++ b/nodedb/src/data/executor/handlers/hybrid_overlay.rs @@ -1,37 +1,29 @@ // SPDX-License-Identifier: BUSL-1.1 -//! Shared read-your-own-writes overlay splice for the two hybrid search -//! handlers (`text_search_hybrid.rs` = vector + text, `text_search_triple.rs` -//! = vector + text + graph). +//! Shared read-your-own-writes support for the two hybrid search handlers +//! (`text_search_hybrid.rs` = vector + text, `text_search_triple.rs` = +//! vector + text + graph). //! -//! The vector and text legs of a hybrid search read only committed state. Inside -//! a transaction the same query must ALSO observe the transaction's own staged -//! document writes (read-your-own-writes), exactly as the single-source -//! vector-only and FTS-only handlers already do. This module folds those staged -//! writes into the vector and text legs by REUSING the single-source overlay -//! merges — [`CoreLoop::merge_vector_overlay_into_search`] and -//! [`CoreLoop::merge_fts_overlay_into_results`] — rather than duplicating the -//! staged-document scoring. It then rebuilds the RRF-ready ranked lists from the -//! merged legs so a same-transaction INSERT / UPDATE / DELETE is reflected in -//! the fused result. -//! -//! A vector or FTS posting is an inline side effect of the document write, not a -//! stageable write of its own, so both merges re-read the transaction's staged -//! DOCUMENT BODIES and re-score them in memory: the vector leg extracts the -//! declared vector field and re-computes distance under the index's metric; the -//! text leg re-tokenizes and BM25-scores against the base corpus stats. Staged -//! tombstones remove the stale committed entry and staged puts over an existing -//! surrogate replace it, mirroring the single-source paths. +//! The text leg ranks the transaction's staged rows in the BM25 search +//! itself, before its top-k cut, through the same staged view the +//! single-source text search reads ([`CoreLoop::hybrid_text_leg`]). The +//! vector leg reads committed state, so the transaction's staged writes are +//! folded into it by the single-source vector overlay merge +//! ([`CoreLoop::merge_vector_overlay_into_search`]): staged tombstones drop +//! the stale committed entry and staged puts over an existing surrogate +//! replace it. //! //! The graph leg's RYOW is a separate concern and is deliberately not touched //! here — the triple handler still reads committed graph state. use nodedb_fts::posting::TextSearchResult; -use nodedb_types::{Surrogate, SurrogateBitmap}; +use nodedb_fts::{FtsSearchParams, IndexScope}; +use nodedb_types::SurrogateBitmap; use super::hybrid_key::HybridFusionKey; use crate::data::executor::core_loop::CoreLoop; -use crate::data::executor::handlers::transaction::overlay::{FtsMergeParams, VectorMergeParams}; +use crate::data::executor::handlers::transaction::overlay::VectorMergeParams; +use crate::data::executor::task::ExecutionTask; use crate::engine::vector::DistanceMetric; use crate::engine::vector::SearchResult; use crate::engine::vector::collection::VectorCollection; @@ -45,29 +37,54 @@ pub(in crate::data::executor) type HybridRankedLegs = ( ); /// Scope for one hybrid overlay splice: the active transaction, its -/// `(database, tenant, collection)` target, the vector and text query inputs, -/// the per-leg over-fetch bound, and any surrogate prefilter applied to the -/// vector leg. Bundled to keep the entry point to a single parameter. +/// `(database, tenant, collection)` target, the vector query input, the +/// vector leg's over-fetch bound, and any surrogate prefilter applied to it. +/// Bundled to keep the entry point to a single parameter. pub(in crate::data::executor) struct HybridOverlayParams<'a> { pub txn_id: TxnId, pub database_id: DatabaseId, pub tid: TenantId, pub collection: &'a str, + /// Vector column of the vector leg. Empty re-scores every declared + /// vector field. + pub vector_field: &'a str, pub query_vector: &'a [f32], - pub query_text: &'a str, pub fetch_k: usize, + /// Rows both legs may return: the plan's prefilter, the residual + /// filters, and RLS, combined. pub filter_bitmap: Option<&'a SurrogateBitmap>, } impl CoreLoop { + /// The text leg of a hybrid search: BM25 over `index`, with the issuing + /// transaction's staged rows ranked in before the cut. Empty when the + /// collection holds no text. + pub(in crate::data::executor) fn hybrid_text_leg( + &self, + task: &ExecutionTask, + tid: u64, + index: Option>, + params: FtsSearchParams<'_>, + ) -> crate::Result> { + let Some(index) = index else { + return Ok(Vec::new()); + }; + let staged = self.text_staged_view(task, tid, index)?; + self.inverted.search_staged( + task.request.database_id.as_u64(), + TenantId::new(tid), + index, + params, + staged.as_ref(), + ) + } + /// Build the vector and text RRF-ranked lists for a hybrid search, folding - /// the transaction's staged document writes into both legs. + /// the transaction's staged document writes into the vector leg. /// - /// The base committed legs are `vector_results` (raw HNSW/IVF hits, whose - /// local ids are resolved to surrogates via `vector_collection`) and - /// `text_results` (base BM25 hits). This reuses the exact single-source - /// overlay merges to re-score the staged puts and drop the staged - /// tombstones, then emits `(vector_ranked, text_ranked)` keyed by + /// `vector_results` are the raw HNSW/IVF hits, whose local ids are + /// resolved to surrogates via `vector_collection`. `text_results` already + /// rank the staged rows. Emits `(vector_ranked, text_ranked)` keyed by /// [`HybridFusionKey`], the shared RRF key space the caller fuses on. pub(in crate::data::executor) fn hybrid_ranked_with_overlay( &self, @@ -76,17 +93,10 @@ impl CoreLoop { vector_collection: Option<&VectorCollection>, text_results: &[TextSearchResult], ) -> crate::Result { - // Base committed legs, in the shape each single-source overlay merge - // consumes: vector hits carry the surrogate-resolved id, text scores - // carry the FTS surrogate + score + fuzzy flag. let mut vector_hits: Vec<_> = vector_results .iter() .map(|r| super::vector_search::build_search_hit(vector_collection, r.id, r.distance)) .collect(); - let mut text_scored: Vec<(Surrogate, f32, bool)> = text_results - .iter() - .map(|r| (r.doc_id, r.score, r.fuzzy)) - .collect(); let db = params.database_id.as_u64(); let tid_u64 = params.tid.as_u64(); @@ -97,11 +107,15 @@ impl CoreLoop { // vector, so declaring several fields is safe. Metric comes from the // field's committed index (or its DDL params), falling back to L2 when // neither is registered yet. - let mut fields: Vec = self - .strict_vector_fields(db, tid_u64, params.collection) - .into_iter() - .map(|(field, _dim)| field) - .collect(); + // A named vector column re-scores that field only. + let mut fields: Vec = if params.vector_field.is_empty() { + self.strict_vector_fields(db, tid_u64, params.collection) + .into_iter() + .map(|(field, _dim)| field) + .collect() + } else { + vec![params.vector_field.to_string()] + }; if fields.is_empty() { fields = self.schemaless_vector_field_names(db, tid_u64, params.collection); } @@ -130,22 +144,9 @@ impl CoreLoop { )?; } - // Text leg RYOW: re-score staged docs via the FTS-only overlay merge. - self.merge_fts_overlay_into_results( - FtsMergeParams { - txn_id: params.txn_id, - database_id: params.database_id, - tid: params.tid, - collection: params.collection, - query: params.query_text, - top_k: params.fetch_k, - }, - &mut text_scored, - )?; - - // Rebuild the RRF-ready ranked lists from the merged legs. Both legs - // key on `HybridFusionKey`, matching the committed-only construction - // in the handlers. + // Rebuild the RRF-ready ranked lists. Both legs key on + // `HybridFusionKey`, matching the committed-only construction in the + // handlers. let vector_ranked = vector_hits .iter() .enumerate() @@ -156,13 +157,13 @@ impl CoreLoop { source: "vector", }) .collect(); - let text_ranked = text_scored + let text_ranked = text_results .iter() .enumerate() - .map(|(rank, (surrogate, score, _fuzzy))| RankedResult { - document_id: HybridFusionKey::for_surrogate(*surrogate), + .map(|(rank, r)| RankedResult { + document_id: HybridFusionKey::for_surrogate(r.doc_id), rank, - score: *score, + score: r.score, source: "text", }) .collect(); diff --git a/nodedb/src/data/executor/handlers/text_rows.rs b/nodedb/src/data/executor/handlers/text_rows.rs new file mode 100644 index 000000000..1d5af69bd --- /dev/null +++ b/nodedb/src/data/executor/handlers/text_rows.rs @@ -0,0 +1,280 @@ +// SPDX-License-Identifier: BUSL-1.1 + +//! Row sets of a full-text read: the current version of each base row, and +//! the rows the residual WHERE filters and RLS admit before ranking. +//! +//! Hydration, eligibility, and the score scan read rows through the same +//! functions here, so the three can never disagree about which version of a +//! row is current or how a predicate sees it. + +use std::ops::ControlFlow; + +use nodedb_types::{StorageKey, Surrogate, SurrogateBitmap}; + +use crate::bridge::scan_filter::{ScanFilter, decode_scan_filters}; +use crate::data::executor::core_loop::CoreLoop; +use crate::data::executor::handlers::transaction::overlay::Staged; +use crate::data::executor::row_shape::sparse_row_to_doc; +use crate::data::executor::sparse_body_format::SparseBodyFormatRef; +use crate::data::executor::task::ExecutionTask; +use crate::engine::sparse::btree_versioned::VersionedScanParams; +use crate::engine::sparse::scan_stop::never_stop; +use crate::types::TenantId; + +/// Rows read per page of a bitemporal collection's current versions. +const BITEMPORAL_PAGE_ROWS: usize = 1024; + +/// The row predicates a full-text read applies before ranking: the residual +/// WHERE filters and the RLS filters. +pub(in crate::data::executor) struct TextRowGate { + filters: Vec, + rls: Vec, +} + +impl TextRowGate { + /// Decode the serialized `Vec` of both predicate sets. + pub(in crate::data::executor) fn new( + filters: &[u8], + rls_filters: &[u8], + ) -> crate::Result { + Ok(Self { + filters: decode_scan_filters(filters, "text search filters")?, + rls: decode_scan_filters(rls_filters, "text search RLS filters")?, + }) + } + + /// Whether every row passes: no residual filter and no RLS filter. + pub(in crate::data::executor) fn is_open(&self) -> bool { + self.filters.is_empty() && self.rls.is_empty() + } + + /// Whether the row image passes every residual filter and every RLS + /// filter. A residual filter that fails to evaluate fails the query. An + /// RLS filter that fails to evaluate denies the row: RLS fails closed and + /// never tells an erroring row from a missing one. + pub(in crate::data::executor) fn admits(&self, image: &[u8]) -> crate::Result { + if !ScanFilter::all_match_binary(&self.filters, image)? { + return Ok(false); + } + Ok(self.rls.is_empty() || ScanFilter::all_match_binary(&self.rls, image).unwrap_or(false)) + } +} + +/// The msgpack image filters and RLS read for a row: the normalized body +/// with its identity under `identity_column`, the same image the base +/// document scan matches and returns. +pub(in crate::data::executor) fn text_row_image( + key: &StorageKey, + body: &[u8], + format: SparseBodyFormatRef<'_>, + identity_column: &str, +) -> Vec { + sparse_row_to_doc(key, body, format, identity_column).1 +} + +impl CoreLoop { + /// The current version of one base row: the versioned table's current + /// version for a bitemporal collection, the document table otherwise. + pub(in crate::data::executor) fn text_base_body( + &self, + database_id: u64, + tid: u64, + collection: &str, + key: &StorageKey, + ) -> crate::Result>> { + if self.is_bitemporal(database_id, tid, collection) { + self.sparse + .versioned_get_current(database_id, tid, collection, key) + } else { + self.sparse.get(database_id, tid, collection, key) + } + } + + /// Visit the current version of every base row of `collection`, from the + /// tables [`Self::text_base_body`] reads, until `f` breaks. A bitemporal + /// collection is read a page at a time, so at most one page of rows is in + /// memory. + pub(in crate::data::executor) fn for_each_text_base_row( + &self, + database_id: u64, + tid: u64, + collection: &str, + mut f: F, + ) -> crate::Result<()> + where + F: FnMut(&StorageKey, &[u8]) -> crate::Result>, + { + if !self.is_bitemporal(database_id, tid, collection) { + return self + .sparse + .scan_documents_while(database_id, tid, collection, f); + } + let mut after: Option = None; + loop { + let page = self.sparse.versioned_scan_as_of_after( + VersionedScanParams { + database_id, + tenant: tid, + coll: collection, + sys_cutoff_ms: None, + valid_at_ms: None, + limit: BITEMPORAL_PAGE_ROWS, + }, + after.as_ref(), + &|_, _| true, + &never_stop, + )?; + let full = page.len() >= BITEMPORAL_PAGE_ROWS; + for (key, body) in &page { + if f(key, body)?.is_break() { + return Ok(()); + } + } + match page.into_iter().next_back() { + Some((key, _)) if full => after = Some(key), + _ => return Ok(()), + } + } + } + + /// Surrogates of `collection` whose row passes `gate`, judged before + /// ranking. Staged rows of the issuing transaction are judged on their + /// staged body. A staged tombstone or TRUNCATE removes base rows. `None` + /// when the gate is open. + pub(in crate::data::executor) fn text_eligible_rows( + &self, + task: &ExecutionTask, + tid: u64, + collection: &str, + gate: &TextRowGate, + ) -> crate::Result> { + if gate.is_open() { + return Ok(None); + } + let database_id = task.request.database_id; + let tenant = TenantId::new(tid); + let format = self.sparse_body_format(database_id, tenant, collection); + let identity_column = self.identity_column(database_id.as_u64(), tid, collection); + let admits = |key: &StorageKey, bytes: &[u8]| -> crate::Result { + gate.admits(&text_row_image( + key, + bytes, + format.as_format_ref(), + &identity_column, + )) + }; + let coll_key = (database_id, tenant, collection.to_string()); + let overlay = match task.request.txn_id { + Some(txn_id) => { + // Read-your-own-writes refreshes the lease (see the reaper). + self.touch_overlay(txn_id); + self.txn_overlays.get(&txn_id) + } + None => None, + }; + + let mut eligible = SurrogateBitmap::new(); + if overlay.is_none_or(|o| o.base_visible(&coll_key)) { + self.for_each_text_base_row(database_id.as_u64(), tid, collection, |key, bytes| { + if admits(key, bytes)? { + eligible.insert(key.surrogate()); + } + Ok(ControlFlow::Continue(())) + })?; + } + if let Some(overlay) = overlay { + for (surrogate, staged) in overlay.iter_for_collection(&coll_key) { + let surrogate = Surrogate::new(surrogate); + let keep = match staged { + Staged::Put(body) => admits(&StorageKey::for_surrogate(surrogate), body)?, + Staged::Tombstone => false, + }; + if keep { + eligible.insert(surrogate); + } else { + eligible.remove(surrogate); + } + } + } + Ok(Some(eligible)) + } +} + +/// The rows both the plan's prefilter and the gated rows admit. +pub(in crate::data::executor) fn combine_eligible( + prefilter: Option<&SurrogateBitmap>, + eligible: Option, +) -> Option { + match (prefilter, eligible) { + (Some(prefilter), Some(eligible)) => Some(prefilter.intersect(&eligible)), + (Some(prefilter), None) => Some(prefilter.clone()), + (None, eligible) => eligible, + } +} + +#[cfg(test)] +mod tests { + use super::*; + + fn msgpack(value: serde_json::Value) -> Vec { + nodedb_types::json_to_msgpack_or_empty(&value) + } + + fn id_filter(id: &str) -> Vec { + let filter = ScanFilter { + field: "id".into(), + op: crate::bridge::scan_filter::FilterOp::Eq, + value: nodedb_types::Value::String(id.into()), + clauses: Vec::new(), + expr: None, + }; + zerompk::to_msgpack_vec(&vec![filter]).unwrap() + } + + /// A row stored without an `id` field matches `id = ''`: the + /// image carries the identity the storage key holds. + #[test] + fn row_image_carries_the_key_identity() { + let key = StorageKey::for_surrogate(Surrogate::new(42)); + let body = msgpack(serde_json::json!({ "body": "rust" })); + let image = text_row_image(&key, &body, SparseBodyFormatRef::Document, "id"); + let identity = key.to_identity(); + let gate = TextRowGate::new(&id_filter(identity.as_str()), &[]).unwrap(); + assert!(gate.admits(&image).unwrap()); + let other = TextRowGate::new(&id_filter("nope"), &[]).unwrap(); + assert!(!other.admits(&image).unwrap()); + } + + /// RLS naming `id` reads the injected identity too. + #[test] + fn rls_on_id_reads_the_key_identity() { + let key = StorageKey::for_surrogate(Surrogate::new(7)); + let body = msgpack(serde_json::json!({ "body": "rust" })); + let image = text_row_image(&key, &body, SparseBodyFormatRef::Document, "id"); + let identity = key.to_identity(); + let gate = TextRowGate::new(&[], &id_filter(identity.as_str())).unwrap(); + assert!(!gate.is_open()); + assert!(gate.admits(&image).unwrap()); + } + + /// A row of a declared-key collection carries no synthesized `id`: a + /// predicate on `id` finds nothing, and one on the key finds the key. + #[test] + fn a_declared_key_row_image_has_no_id() { + let key = StorageKey::for_surrogate(Surrogate::new(9)); + let body = msgpack(serde_json::json!({ "sku": "p1", "body": "rust" })); + let image = text_row_image(&key, &body, SparseBodyFormatRef::Document, "sku"); + let identity = key.to_identity(); + let on_id = TextRowGate::new(&id_filter(identity.as_str()), &[]).unwrap(); + assert!(!on_id.admits(&image).unwrap()); + let sku = ScanFilter { + field: "sku".into(), + op: crate::bridge::scan_filter::FilterOp::Eq, + value: nodedb_types::Value::String("p1".into()), + clauses: Vec::new(), + expr: None, + }; + let on_key = TextRowGate::new(&zerompk::to_msgpack_vec(&vec![sku]).unwrap(), &[]).unwrap(); + assert!(on_key.admits(&image).unwrap()); + } +} diff --git a/nodedb/src/data/executor/handlers/text_score_columns.rs b/nodedb/src/data/executor/handlers/text_score_columns.rs new file mode 100644 index 000000000..7177b1610 --- /dev/null +++ b/nodedb/src/data/executor/handlers/text_score_columns.rs @@ -0,0 +1,128 @@ +// SPDX-License-Identifier: BUSL-1.1 + +//! Per-row `bm25_score(field, query)` columns. +//! +//! Each column is a per-document scorer: its query resolves once, then each +//! emitted row is scored by point lookups of its postings and of its +//! recorded length, read in one transaction per batch of rows. No column +//! builds a corpus-wide score map, so a `LIMIT 10` read scores ten rows. +//! +//! A row scores the BM25 value of its match. A row the index holds that the +//! query does not match scores `0.0`. A row the index does not hold (no text +//! in the field, or no text at all for the whole-document index) scores +//! `null`. The issuing transaction's staged rows are scored from their +//! staged text. The AND-mode fallback of a column is decided over the rows +//! the reading query admits, the same rows its search ranks. + +use nodedb_fts::posting::QueryMode; +use nodedb_fts::{DocScore, TextQuery}; +use nodedb_physical::physical_plan::TextScoreSpec; +use nodedb_types::{Surrogate, SurrogateBitmap}; + +use crate::data::executor::core_loop::CoreLoop; +use crate::data::executor::response_codec::DocumentRow; +use crate::data::executor::task::ExecutionTask; +use crate::engine::sparse::inverted::TextDocScorer; +use crate::types::TenantId; + +/// One score column: its output alias and its scorer. `None` scorer: the +/// collection holds no text, so every row scores `null`. +struct ScoreColumn<'a> { + alias: &'a str, + scorer: Option>, +} + +/// The score columns of one read. +pub(in crate::data::executor) struct ScoreColumns<'a> { + columns: Vec>, +} + +impl ScoreColumns<'_> { + /// Write every column into each of `rows`. `surrogates` is parallel to + /// `rows`. A row that is not an object takes none. + pub(in crate::data::executor) fn inject( + &self, + surrogates: &[Surrogate], + rows: &mut [DocumentRow], + ) -> crate::Result<()> { + for column in &self.columns { + let scores = match &column.scorer { + Some(scorer) => scorer.score(surrogates)?, + None => vec![DocScore::Absent; surrogates.len()], + }; + for (row, score) in rows.iter_mut().zip(scores) { + if let serde_json::Value::Object(map) = &mut row.data { + map.insert(column.alias.to_string(), score_value(score)); + } + } + } + Ok(()) + } +} + +/// The JSON value of a score: its number, `0.0` for a held miss, `null` for +/// a row the index does not hold. +pub(in crate::data::executor) fn score_value(score: DocScore) -> serde_json::Value { + let number = match score { + DocScore::Match(value) => Some(value), + DocScore::Miss => Some(0.0), + DocScore::Absent => None, + }; + number + .and_then(|value| serde_json::Number::from_f64(f64::from(value))) + .map_or(serde_json::Value::Null, serde_json::Value::Number) +} + +impl CoreLoop { + /// The score columns of `specs`, in order. `eligible` is the set of rows + /// the reading query admits. + pub(in crate::data::executor) fn text_score_columns<'a>( + &'a self, + task: &ExecutionTask, + tid: u64, + collection: &'a str, + specs: &'a [TextScoreSpec], + eligible: Option<&SurrogateBitmap>, + ) -> crate::Result> { + let database_id = task.request.database_id.as_u64(); + let tenant = TenantId::new(tid); + let mut columns = Vec::with_capacity(specs.len()); + for spec in specs { + let scorer = match self.text_index(task, tid, collection, spec.field.as_deref())? { + None => None, + Some(index) => { + let staged = self.text_staged_view(task, tid, index)?; + Some(self.inverted.doc_scorer( + database_id, + tenant, + index, + TextQuery { + query: &spec.query, + fuzzy_enabled: spec.fuzzy, + mode: QueryMode::from(spec.mode), + }, + eligible, + staged, + )?) + } + }; + columns.push(ScoreColumn { + alias: &spec.alias, + scorer, + }); + } + Ok(ScoreColumns { columns }) + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn a_match_scores_a_held_miss_scores_zero_and_an_unheld_row_is_null() { + assert_eq!(score_value(DocScore::Match(1.5)), serde_json::json!(1.5)); + assert_eq!(score_value(DocScore::Miss), serde_json::json!(0.0)); + assert_eq!(score_value(DocScore::Absent), serde_json::Value::Null); + } +} diff --git a/nodedb/src/data/executor/handlers/text_score_sink.rs b/nodedb/src/data/executor/handlers/text_score_sink.rs new file mode 100644 index 000000000..2e973b618 --- /dev/null +++ b/nodedb/src/data/executor/handlers/text_score_sink.rs @@ -0,0 +1,277 @@ +// SPDX-License-Identifier: BUSL-1.1 + +//! The row sink of a score scan. +//! +//! Rows stream in one at a time and are scored in batches, so memory holds +//! one batch plus the rows the scan returns. With a LIMIT and no score order +//! the sink stops the scan once it holds the limit. With a LIMIT ordered by +//! a score column it keeps the best rows in a bounded heap and never holds +//! more than the limit. With no LIMIT it returns every row. + +use std::cmp::Ordering; +use std::collections::BinaryHeap; +use std::ops::ControlFlow; + +use nodedb_physical::physical_plan::{ScoreScanBound, ScoreScanOrder}; +use nodedb_types::Surrogate; + +use crate::data::executor::handlers::text_score_columns::ScoreColumns; +use crate::data::executor::response_codec::DocumentRow; + +/// Rows scored per batch: one doc-length read transaction per column. +const SCORE_BATCH_ROWS: usize = 256; + +/// The direction of an ordered bounded scan. +#[derive(Clone, Copy)] +struct Direction { + ascending: bool, + nulls_first: bool, +} + +/// One kept row of an ordered bounded scan. +struct Ranked { + score: Option, + /// Arrival order. A later row ranks after an earlier one with the same + /// score, so the heap keeps the earliest of tied rows. + seq: u64, + direction: Direction, + row: DocumentRow, +} + +impl Ranked { + /// `Less` when `self` comes before `other` in the output. + fn output_cmp(&self, other: &Self) -> Ordering { + let Direction { + ascending, + nulls_first, + } = self.direction; + let by_score = match (self.score, other.score) { + (None, None) => Ordering::Equal, + (None, Some(_)) if nulls_first => Ordering::Less, + (None, Some(_)) => Ordering::Greater, + (Some(_), None) if nulls_first => Ordering::Greater, + (Some(_), None) => Ordering::Less, + (Some(a), Some(b)) if ascending => a.total_cmp(&b), + (Some(a), Some(b)) => b.total_cmp(&a), + }; + by_score.then(self.seq.cmp(&other.seq)) + } +} + +impl PartialEq for Ranked { + fn eq(&self, other: &Self) -> bool { + self.output_cmp(other) == Ordering::Equal + } +} + +impl Eq for Ranked {} + +impl PartialOrd for Ranked { + fn partial_cmp(&self, other: &Self) -> Option { + Some(self.cmp(other)) + } +} + +impl Ord for Ranked { + /// The heap's maximum is the row that comes last in the output. + fn cmp(&self, other: &Self) -> Ordering { + self.output_cmp(other) + } +} + +/// Where scored rows go. +enum Target { + /// Every row, in arrival order. `limit` stops the scan once reached. + All { + rows: Vec, + limit: Option, + }, + /// The first `limit` rows in `order`. + Top { + heap: BinaryHeap, + limit: usize, + order: ScoreScanOrder, + seq: u64, + }, +} + +/// Streams admitted rows of a score scan into its result. +pub(in crate::data::executor) struct ScoreScanSink<'a> { + columns: &'a ScoreColumns<'a>, + pending_keys: Vec, + pending_rows: Vec, + target: Target, +} + +impl<'a> ScoreScanSink<'a> { + pub(in crate::data::executor) fn new( + columns: &'a ScoreColumns<'a>, + bound: Option<&ScoreScanBound>, + ) -> Self { + let target = match bound { + None => Target::All { + rows: Vec::new(), + limit: None, + }, + Some(ScoreScanBound { rows, order: None }) => Target::All { + rows: Vec::with_capacity((*rows).min(SCORE_BATCH_ROWS)), + limit: Some(*rows), + }, + Some(ScoreScanBound { + rows, + order: Some(order), + }) => Target::Top { + heap: BinaryHeap::with_capacity((*rows).min(SCORE_BATCH_ROWS)), + limit: *rows, + order: order.clone(), + seq: 0, + }, + }; + Self { + columns, + pending_keys: Vec::new(), + pending_rows: Vec::new(), + target, + } + } + + /// Whether the sink takes no more rows. + pub(in crate::data::executor) fn is_full(&self) -> bool { + match &self.target { + Target::All { + rows, + limit: Some(limit), + } => rows.len() + self.pending_rows.len() >= *limit, + Target::All { limit: None, .. } => false, + Target::Top { limit, .. } => *limit == 0, + } + } + + /// Take one admitted row. `Break` once the sink is full. + pub(in crate::data::executor) fn push( + &mut self, + surrogate: Surrogate, + row: DocumentRow, + ) -> crate::Result> { + if self.is_full() { + return Ok(ControlFlow::Break(())); + } + self.pending_keys.push(surrogate); + self.pending_rows.push(row); + if self.pending_rows.len() >= SCORE_BATCH_ROWS { + self.flush()?; + } + Ok(if self.is_full() { + ControlFlow::Break(()) + } else { + ControlFlow::Continue(()) + }) + } + + /// Score the pending batch and move it into the target. + fn flush(&mut self) -> crate::Result<()> { + if self.pending_rows.is_empty() { + return Ok(()); + } + self.columns + .inject(&self.pending_keys, &mut self.pending_rows)?; + self.pending_keys.clear(); + let batch = std::mem::take(&mut self.pending_rows); + match &mut self.target { + Target::All { rows, .. } => rows.extend(batch), + Target::Top { + heap, + limit, + order, + seq, + } => { + for row in batch { + let score = match &row.data { + serde_json::Value::Object(map) => { + map.get(&order.alias).and_then(serde_json::Value::as_f64) + } + _ => None, + }; + let ranked = Ranked { + score, + seq: *seq, + direction: Direction { + ascending: order.ascending, + nulls_first: order.nulls_first, + }, + row, + }; + *seq += 1; + if heap.len() < *limit { + heap.push(ranked); + } else if heap.peek().is_some_and(|worst| ranked < *worst) { + heap.pop(); + heap.push(ranked); + } + } + } + } + Ok(()) + } + + /// The scan's rows: in arrival order, or in score order for an ordered + /// bound. + pub(in crate::data::executor) fn finish(mut self) -> crate::Result> { + self.flush()?; + Ok(match self.target { + Target::All { rows, .. } => rows, + Target::Top { heap, .. } => heap + .into_sorted_vec() + .into_iter() + .map(|ranked| ranked.row) + .collect(), + }) + } +} + +#[cfg(test)] +mod tests { + use super::*; + + fn ranked(score: Option, seq: u64, ascending: bool, nulls_first: bool) -> Ranked { + Ranked { + score, + seq, + direction: Direction { + ascending, + nulls_first, + }, + row: DocumentRow { + id: seq.to_string(), + data: serde_json::Value::Null, + }, + } + } + + #[test] + fn descending_puts_higher_scores_first_and_nulls_where_asked() { + let high = ranked(Some(2.0), 0, false, true); + let low = ranked(Some(1.0), 1, false, true); + let null = ranked(None, 2, false, true); + assert_eq!(high.output_cmp(&low), Ordering::Less); + assert_eq!(null.output_cmp(&high), Ordering::Less, "NULLS FIRST"); + + let null_last = ranked(None, 2, false, false); + let high_last = ranked(Some(2.0), 0, false, false); + assert_eq!(null_last.output_cmp(&high_last), Ordering::Greater); + } + + #[test] + fn ascending_puts_lower_scores_first() { + let low = ranked(Some(1.0), 1, true, false); + let high = ranked(Some(2.0), 0, true, false); + assert_eq!(low.output_cmp(&high), Ordering::Less); + } + + #[test] + fn ties_keep_arrival_order() { + let first = ranked(Some(1.0), 0, false, false); + let second = ranked(Some(1.0), 1, false, false); + assert_eq!(first.output_cmp(&second), Ordering::Less); + } +} diff --git a/nodedb/src/data/executor/handlers/text_search.rs b/nodedb/src/data/executor/handlers/text_search.rs index 6fd7d2919..931614a94 100644 --- a/nodedb/src/data/executor/handlers/text_search.rs +++ b/nodedb/src/data/executor/handlers/text_search.rs @@ -6,14 +6,17 @@ use tracing::debug; use nodedb_fts::FtsSearchParams; use nodedb_fts::posting::QueryMode; +use nodedb_types::{Surrogate, SurrogateBitmap}; use crate::bridge::envelope::{ErrorCode, Response}; +use nodedb_physical::physical_plan::TextScoreSpec; + use crate::data::executor::core_loop::CoreLoop; use crate::data::executor::handlers::document::read::decode::decode_scanned_row; -use crate::data::executor::handlers::transaction::overlay::FtsMergeParams; +use crate::data::executor::handlers::text_rows::{TextRowGate, combine_eligible, text_row_image}; +use crate::data::executor::handlers::text_score_columns::ScoreColumns; use crate::data::executor::response_codec::DocumentRow; -use crate::data::executor::scan_normalize::sparse_body_to_msgpack; use crate::data::executor::task::ExecutionTask; use crate::types::{DatabaseId, TenantId, TxnId}; @@ -21,11 +24,20 @@ use crate::types::{DatabaseId, TenantId, TxnId}; pub(in crate::data::executor) struct TextSearchParams<'a> { pub tid: u64, pub collection: &'a str, + /// Field index the query reads. `None` reads the whole-document index. + pub field: Option<&'a str>, pub query: &'a str, + /// Hits returned. `usize::MAX` returns every match. pub top_k: usize, + /// Boolean combination of the query terms. + pub mode: nodedb_types::text_search::QueryMode, pub fuzzy: bool, - pub prefilter: Option<&'a nodedb_types::SurrogateBitmap>, + pub prefilter: Option<&'a SurrogateBitmap>, + /// Residual WHERE predicates (`Vec`), applied before ranking. + pub filters: &'a [u8], + /// RLS filters (`Vec`), applied before ranking. pub rls_filters: &'a [u8], + pub scores: &'a [TextScoreSpec], } /// Parameters for the internal [`CoreLoop::hydrate_text_hits`] helper. @@ -42,13 +54,13 @@ pub(in crate::data::executor) struct HydrateTextHitsParams<'a> { pub database_id: u64, pub tid: u64, pub collection: &'a str, - pub top_k: usize, - pub rls_filters: &'a [u8], /// The issuing transaction, when this read runs inside `BEGIN..COMMIT`. /// A matched surrogate's body is resolved from this transaction's /// staging overlay first (a doc inserted/updated in THIS transaction is /// not yet in base storage), falling back to base storage otherwise. pub txn_id: Option, + /// Score columns injected into each hydrated row. + pub scores: &'a ScoreColumns<'a>, } impl CoreLoop { @@ -58,98 +70,88 @@ impl CoreLoop { task: &ExecutionTask, params: TextSearchParams<'_>, ) -> Response { - let TextSearchParams { - tid, - collection, - query, - top_k, - fuzzy, - prefilter, - rls_filters, - } = params; - let tenant_id = TenantId::new(tid); - debug!(core = self.core_id, tid, %collection, %query, top_k, fuzzy, "text search"); + debug!( + core = self.core_id, + tid = params.tid, + collection = %params.collection, + field = ?params.field, + query = %params.query, + top_k = params.top_k, + mode = params.mode.as_str(), + fuzzy = params.fuzzy, + "text search" + ); // Scan-quiesce gate. - let _scan_guard = match self.acquire_scan_guard(task, tid, collection) { + let _scan_guard = match self.acquire_scan_guard(task, params.tid, params.collection) { Ok(g) => g, Err(resp) => return resp, }; - // Fetch extra candidates when RLS is active. - let fetch_k = if rls_filters.is_empty() { - top_k - } else { - top_k.saturating_mul(2).max(20) + let rows = match self.text_search_rows(task, ¶ms) { + Ok(rows) => rows, + Err(e) => return self.response_error(task, e), }; - let results = match self.inverted.search( + if let Some(ref m) = self.metrics { + m.record_fts_search(0); + } + match super::super::response_codec::encode(&rows) { + Ok(payload) => self.response_with_payload(task, payload), + Err(e) => self.response_error(task, ErrorCode::from(e)), + } + } + + /// The hydrated hits of a text search, in this order: resolve the field's + /// index, restrict candidates to the rows the residual filters and RLS + /// admit, rank with the transaction's staged rows folded in, hydrate, and + /// inject the score columns. + /// + /// RLS needs no check after ranking: the eligible rows are read on this + /// core in the same task as the ranked rows are hydrated, so no write can + /// land between the two, and a row with no body is never eligible under a + /// gate. + fn text_search_rows( + &self, + task: &ExecutionTask, + p: &TextSearchParams<'_>, + ) -> crate::Result> { + let tenant_id = TenantId::new(p.tid); + let Some(index) = self.text_index(task, p.tid, p.collection, p.field)? else { + return Ok(Vec::new()); + }; + let gate = TextRowGate::new(p.filters, p.rls_filters)?; + let eligible = combine_eligible( + p.prefilter, + self.text_eligible_rows(task, p.tid, p.collection, &gate)?, + ); + let staged = self.text_staged_view(task, p.tid, index)?; + let hits = self.inverted.search_staged( task.request.database_id.as_u64(), tenant_id, - collection, + index, FtsSearchParams { - query, - top_k: fetch_k, - fuzzy_enabled: fuzzy, - mode: QueryMode::And, - prefilter, + query: p.query, + top_k: p.top_k, + fuzzy_enabled: p.fuzzy, + mode: QueryMode::from(p.mode), + prefilter: eligible.as_ref(), }, - ) { - Ok(r) => r, - Err(e) => return self.response_error(task, e), - }; - - // Read-your-own-writes for FTS: fold this transaction's staged - // document bodies into the base search result before hydration, so - // a doc inserted/updated earlier in the same transaction appears - // (and one deleted is excluded) before COMMIT. - let mut merged: Vec<(nodedb_types::Surrogate, f32, bool)> = results - .iter() - .map(|r| (r.doc_id, r.score, r.fuzzy)) - .collect(); - if let Some(txn_id) = task.request.txn_id - && let Err(e) = self.merge_fts_overlay_into_results( - FtsMergeParams { - txn_id, - database_id: task.request.database_id, - tid: tenant_id, - collection, - query, - top_k: fetch_k, - }, - &mut merged, - ) - { - return self.response_error(task, e); - } + staged.as_ref(), + )?; - let rows = match self.hydrate_text_hits( - merged, + let scores = + self.text_score_columns(task, p.tid, p.collection, p.scores, eligible.as_ref())?; + self.hydrate_text_hits( + hits.iter().map(|r| (r.doc_id, r.score, r.fuzzy)), HydrateTextHitsParams { database_id: task.request.database_id.as_u64(), - tid, - collection, - top_k, - rls_filters, + tid: p.tid, + collection: p.collection, txn_id: task.request.txn_id, + scores: &scores, }, - ) { - Ok(rows) => rows, - Err(e) => return self.response_error(task, e), - }; - - if let Some(ref m) = self.metrics { - m.record_fts_search(0); - } - match super::super::response_codec::encode(&rows) { - Ok(payload) => self.response_with_payload(task, payload), - Err(e) => self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ), - } + ) } pub(in crate::data::executor) fn strict_schema_for( @@ -170,28 +172,27 @@ impl CoreLoop { }) } + /// Resolve ranked hits to rows, in hit order, each carrying `score`, + /// `fuzzy`, its identity under the collection's identity column, and the + /// score columns. pub(in crate::data::executor) fn hydrate_text_hits( &self, hits: I, params: HydrateTextHitsParams<'_>, ) -> crate::Result> where - I: IntoIterator, + I: IntoIterator, { let HydrateTextHitsParams { database_id, tid, collection, - top_k, - rls_filters, txn_id, + scores, } = params; // Read-your-own-writes for the projection step: a matched surrogate - // added by the overlay merge (a doc inserted/updated in THIS - // transaction) has no body in base storage yet, so its columns must - // be projected from the staged `Put` bytes. Resolve every matched - // surrogate's body from the transaction's overlay first, falling - // back to base storage — mirroring the index-lookup fetch fix. + // the transaction staged has no body in base storage yet, so its + // columns come from the staged bytes. let coll_key = ( DatabaseId::new(database_id), TenantId::new(tid), @@ -203,73 +204,49 @@ impl CoreLoop { // necessarily mis-reads one of them. let format = self.sparse_body_format(DatabaseId::new(database_id), TenantId::new(tid), collection); + let identity_column = self.identity_column(database_id, tid, collection); let mut rows: Vec = Vec::new(); + let mut surrogates: Vec = Vec::new(); for (surrogate, score, fuzzy) in hits { - if rows.len() >= top_k { - break; - } let storage_key = nodedb_types::StorageKey::for_surrogate(surrogate); - let hex_key = storage_key.to_string(); - let bytes_opt = match self.overlay_or_base_body(txn_id, &coll_key, &storage_key, || { - self.sparse.get(database_id, tid, collection, &storage_key) - }) { - Ok(b) => b, - Err(e) => { - tracing::warn!( - err = %e, - %hex_key, - %collection, - "sparse store error during text hit hydration; skipping row" + let bytes_opt = self.overlay_or_base_body(txn_id, &coll_key, &storage_key, || { + self.text_base_body(database_id, tid, collection, &storage_key) + })?; + // A surrogate with no body was indexed for FTS without a document + // write (FtsIndex frames synced from Lite). Its row is the + // surrogate-derived key alone. Under a filter or RLS gate such a + // row is never eligible, so it never reaches here. + let mut value = match bytes_opt { + Some(ref bytes) => { + // The image carries the row identity under the identity + // column, the same image the document scan returns. + let image = text_row_image( + &storage_key, + bytes, + format.as_format_ref(), + &identity_column, ); - continue; + decode_scanned_row(bytes, Some(image.as_slice()), format.as_format_ref())? } - }; - // When the sparse store has no body for this surrogate the document - // was indexed for FTS without a corresponding document write (e.g. - // FtsIndex frames synced from Lite). Return a minimal row containing - // only the surrogate-derived ID so callers that only project `id` - // (the common case in sync interop tests and CDC pipelines) still - // receive a result. RLS filters are skipped when there is no body. - let mut value = if let Some(ref bytes) = bytes_opt { - // RLS is evaluated against the NORMALIZED msgpack image — the - // same bytes the projection below reads — so the gate and the - // output agree. A strict Binary Tuple is not a msgpack map at - // all and a vector-primary sidecar is a TAGGED one, so a - // predicate pushed at the stored bytes finds no field it - // recognizes and drops the row on a format mismatch rather - // than on policy. - // - // The image built for the gate is then handed to the decoder - // rather than dropped, so a sidecar row is transcoded once for - // both steps. - let normalized = if rls_filters.is_empty() { - None - } else { - let normalized = sparse_body_to_msgpack(bytes, format.as_format_ref()); - if !super::rls_eval::rls_check_msgpack_bytes(rls_filters, &normalized) { - continue; - } - Some(normalized) - }; - decode_scanned_row(bytes, normalized.as_deref(), format.as_format_ref())? - } else { - serde_json::Value::Object(serde_json::Map::new()) + None => serde_json::Value::Object(serde_json::Map::new()), }; if let serde_json::Value::Object(ref mut map) = value { map.insert( "score".to_string(), serde_json::Value::Number( - serde_json::Number::from_f64(score as f64) + serde_json::Number::from_f64(f64::from(score)) .unwrap_or_else(|| serde_json::Number::from(0)), ), ); map.insert("fuzzy".to_string(), serde_json::Value::Bool(fuzzy)); } + surrogates.push(surrogate); rows.push(DocumentRow { - id: hex_key, + id: storage_key.to_string(), data: value, }); } + scores.inject(&surrogates, &mut rows)?; Ok(rows) } } diff --git a/nodedb/src/data/executor/handlers/text_search_hybrid.rs b/nodedb/src/data/executor/handlers/text_search_hybrid.rs index f2876b67c..9752f6d43 100644 --- a/nodedb/src/data/executor/handlers/text_search_hybrid.rs +++ b/nodedb/src/data/executor/handlers/text_search_hybrid.rs @@ -10,7 +10,7 @@ use nodedb_fts::posting::QueryMode; use crate::bridge::envelope::{ErrorCode, Response}; use crate::data::executor::core_loop::CoreLoop; -use crate::data::executor::scan_normalize::sparse_body_to_msgpack; +use crate::data::executor::handlers::text_rows::{TextRowGate, combine_eligible}; use crate::data::executor::task::ExecutionTask; use crate::types::TenantId; @@ -21,10 +21,20 @@ const DEFAULT_VECTOR_WEIGHT: f32 = 0.5; pub(in crate::data::executor) struct HybridSearchParams<'a> { pub tid: u64, pub collection: &'a str, + /// Vector column the vector leg searches. Empty names the + /// collection-level index. + pub vector_field: &'a str, pub query_vector: &'a [f32], + /// Field index the text leg reads. `None` reads the whole-document index. + pub text_field: Option<&'a str>, pub query_text: &'a str, + /// Residual WHERE predicates (`Vec`), applied to both legs + /// before fusion. + pub filters: &'a [u8], pub top_k: usize, pub ef_search: usize, + /// Boolean combination of the text leg's query terms. + pub mode: nodedb_types::text_search::QueryMode, pub fuzzy: bool, pub vector_weight: f32, pub filter_bitmap: Option<&'a nodedb_types::SurrogateBitmap>, @@ -45,10 +55,14 @@ impl CoreLoop { let HybridSearchParams { tid, collection, + vector_field, query_vector, + text_field, query_text, + filters, top_k, ef_search, + mode, fuzzy, vector_weight, filter_bitmap, @@ -83,9 +97,30 @@ impl CoreLoop { // enough material to fuse. 3x is a good balance. let fetch_k = top_k.saturating_mul(3).max(20); + // The rows the residual filters and RLS admit restrict both legs + // before fusion, so the fused top-k counts only admitted rows. + let eligible = match TextRowGate::new(filters, rls_filters) + .and_then(|gate| self.text_eligible_rows(task, tid, collection, &gate)) + { + Ok(eligible) => combine_eligible(filter_bitmap, eligible), + Err(e) => return self.response_error(task, e), + }; + let filter_bitmap = eligible.as_ref(); + let text_index = match self.text_index(task, tid, collection, text_field) { + Ok(index) => index, + Err(e) => return self.response_error(task, e), + }; + // 1. Vector search. - let index_key = - CoreLoop::vector_index_key(task.request.database_id.as_u64(), tid, collection, ""); + let index_key = match self.resolve_vector_index_key( + task.request.database_id.as_u64(), + tid, + collection, + vector_field, + ) { + Ok(key) => key, + Err(e) => return self.response_error(task, e), + }; let vector_collection = self.vector_collections.get(&index_key); let vector_results = match vector_collection { Some(index) => { @@ -108,17 +143,18 @@ impl CoreLoop { None => Vec::new(), }; - // 2. Text search (no surrogate prefilter for the text leg of hybrid search). - let text_results = match self.inverted.search( - task.request.database_id.as_u64(), - tenant_id, - collection, + // 2. Text search over the field's index, restricted to the same rows, + // with the transaction's staged rows ranked in. + let text_results = match self.hybrid_text_leg( + task, + tid, + text_index, FtsSearchParams { query: query_text, top_k: fetch_k, fuzzy_enabled: fuzzy, - mode: QueryMode::And, - prefilter: None, + mode: QueryMode::from(mode), + prefilter: filter_bitmap, }, ) { Ok(results) => results, @@ -151,7 +187,7 @@ impl CoreLoop { // in via the shared overlay splice (which reuses the single-source // vector/FTS overlay merges). Outside a transaction the committed-only // construction below runs unchanged. - let (vector_ranked, text_ranked): super::hybrid_overlay::HybridRankedLegs = + let (mut vector_ranked, mut text_ranked): super::hybrid_overlay::HybridRankedLegs = if let Some(txn_id) = task.request.txn_id { match self.hybrid_ranked_with_overlay( super::hybrid_overlay::HybridOverlayParams { @@ -159,8 +195,8 @@ impl CoreLoop { database_id: task.request.database_id, tid: tenant_id, collection, + vector_field, query_vector, - query_text, fetch_k, filter_bitmap, }, @@ -196,48 +232,24 @@ impl CoreLoop { (vector_ranked, text_ranked) }; + // Each leg keeps only the rows the filters and RLS admit, judged + // before ranking on the transaction's view of each row. The fused + // top-k then counts only admitted rows. A headless vector hit has no + // row, so a gate excludes it. + super::hybrid_key::retain_admitted(&mut vector_ranked, filter_bitmap); + super::hybrid_key::retain_admitted(&mut text_ranked, filter_bitmap); + let fused = reciprocal_rank_fusion_weighted( &[vector_ranked, text_ranked], &[k_vector, k_text], top_k, ); - // Build response with per-engine rank diagnostics. - // RLS post-fusion: filter fused results by looking up each document. - // - // The predicate runs against the NORMALIZED msgpack image, never the - // stored bytes: a strict Binary Tuple is not a msgpack map at all and a - // vector-primary sidecar is a TAGGED one, so a predicate pushed at the - // stored bytes finds no field it recognizes and the row is dropped on a - // format mismatch rather than on policy. The encoding is resolved from - // the collection's registered kind — a tagged map and a plain document - // map share the same map header, so the bytes cannot answer it. - let body_format = self.sparse_body_format(task.request.database_id, tenant_id, collection); - // The fused key is rendered once per row here, at the response - // envelope; it is the only place the key becomes text. + // Build response with per-engine rank diagnostics. The fused key is + // rendered once per row here, at the response envelope; it is the + // only place the key becomes text. let rendered: Vec<(String, &crate::query::fusion::FusedResult)> = fused .iter() - .filter(|f| { - if rls_filters.is_empty() { - return true; - } - // A headless hit has no stored row to check the policy against, - // so it is treated the same as a row the lookup cannot find. - let Some(key) = f.document_id.storage_key() else { - return false; - }; - match self - .sparse - .get(task.request.database_id.as_u64(), tid, collection, &key) - { - Ok(Some(bytes)) => { - let normalized = - sparse_body_to_msgpack(&bytes, body_format.as_format_ref()); - super::rls_eval::rls_check_msgpack_bytes(rls_filters, &normalized) - } - _ => false, - } - }) .map(|f| (f.document_id.to_string(), f)) .collect(); let results: Vec<_> = rendered @@ -265,12 +277,7 @@ impl CoreLoop { } match super::super::response_codec::encode(&results) { Ok(payload) => self.response_with_payload(task, payload), - Err(e) => self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ), + Err(e) => self.response_error(task, ErrorCode::from(e)), } } } diff --git a/nodedb/src/data/executor/handlers/text_search_scan.rs b/nodedb/src/data/executor/handlers/text_search_scan.rs index 8e98ce8ea..ad844236e 100644 --- a/nodedb/src/data/executor/handlers/text_search_scan.rs +++ b/nodedb/src/data/executor/handlers/text_search_scan.rs @@ -2,30 +2,56 @@ //! Phrase search and BM25-score-scan handlers for the Data Plane CoreLoop. -use std::collections::HashMap; +use std::ops::ControlFlow; use tracing::debug; -use nodedb_fts::FtsSearchParams; -use nodedb_fts::posting::QueryMode; -use nodedb_types::StorageKey; +use nodedb_physical::physical_plan::{ScoreScanBound, TextScoreSpec}; +use nodedb_types::{StorageKey, Surrogate}; use crate::bridge::envelope::{ErrorCode, Response}; use crate::data::executor::core_loop::CoreLoop; -use crate::data::executor::handlers::document::read::decode::decode_scanned_document; +use crate::data::executor::handlers::document::read::decode::decode_scanned_row; +use crate::data::executor::handlers::text_rows::{TextRowGate, combine_eligible, text_row_image}; +use crate::data::executor::handlers::text_score_sink::ScoreScanSink; use crate::data::executor::handlers::text_search::HydrateTextHitsParams; -use crate::data::executor::handlers::transaction::overlay::FtsMergeParams; +use crate::data::executor::handlers::transaction::overlay::{Staged, staged_phrase_hits}; use crate::data::executor::response_codec::DocumentRow; use crate::data::executor::task::ExecutionTask; use crate::types::TenantId; -/// Upper bound on hits fetched by `BM25ScoreScan` to populate per-row scores. -/// The downstream BMW scorer pre-allocates `Vec::with_capacity(top_k)`, so -/// `usize::MAX` would overflow on element-size multiplication. One million is -/// well above any realistic collection size for in-process score injection -/// while staying safely allocatable. -const BM25_SCAN_MAX_HITS: usize = 1_000_000; +/// Parameters for [`CoreLoop::execute_phrase_search`]. +pub(in crate::data::executor) struct PhraseSearchParams<'a> { + pub tid: u64, + pub collection: &'a str, + /// Field index the phrase reads. `None` reads the whole-document index. + pub field: Option<&'a str>, + /// The phrase words as written. They are analyzed here, once, with the + /// collection's analyzer. + pub terms: &'a [String], + /// Hits returned. `usize::MAX` returns every match. + pub top_k: usize, + pub prefilter: Option<&'a nodedb_types::SurrogateBitmap>, + /// Residual WHERE predicates (`Vec`), applied before ranking. + pub filters: &'a [u8], + /// RLS filters (`Vec`), applied before ranking. + pub rls_filters: &'a [u8], + pub scores: &'a [TextScoreSpec], +} + +/// Parameters for [`CoreLoop::execute_bm25_score_scan`]. +pub(in crate::data::executor) struct ScoreScanParams<'a> { + pub tid: u64, + pub collection: &'a str, + /// Residual WHERE predicates (`Vec`). + pub filters: &'a [u8], + /// RLS filters (`Vec`). A row that fails one is not emitted. + pub rls_filters: &'a [u8], + pub scores: &'a [TextScoreSpec], + /// The query's LIMIT, and the score order it keeps, pushed into the scan. + pub bound: Option<&'a ScoreScanBound>, +} impl CoreLoop { /// Execute an exact phrase search. @@ -35,231 +61,222 @@ impl CoreLoop { pub(in crate::data::executor) fn execute_phrase_search( &self, task: &ExecutionTask, - tid: u64, - collection: &str, - terms: &[String], - top_k: usize, - prefilter: Option<&nodedb_types::SurrogateBitmap>, + params: PhraseSearchParams<'_>, ) -> Response { - let tenant_id = TenantId::new(tid); - debug!(core = self.core_id, tid, %collection, term_count = terms.len(), top_k, "phrase search"); + debug!( + core = self.core_id, + tid = params.tid, + collection = %params.collection, + field = ?params.field, + term_count = params.terms.len(), + top_k = params.top_k, + "phrase search" + ); - let _scan_guard = match self.acquire_scan_guard(task, tid, collection) { + let _scan_guard = match self.acquire_scan_guard(task, params.tid, params.collection) { Ok(g) => g, Err(resp) => return resp, }; - let results = match self.inverted.phrase_search( - task.request.database_id.as_u64(), - tenant_id, - collection, - crate::engine::sparse::inverted::PhraseSearchParams { - terms, - top_k, - prefilter, - }, - ) { - Ok(r) => r, + let rows = match self.phrase_search_rows(task, ¶ms) { + Ok(rows) => rows, Err(e) => return self.response_error(task, e), }; + self.text_rows_response(task, rows) + } - // Read-your-own-writes for FTS phrase search: fold staged document - // bodies into the base result with FAITHFUL phrase semantics. The - // staged doc's own analyzed token positions are self-contained, so - // the merge verifies the phrase's terms occur as a contiguous, - // in-order run (zero slop — the same adjacency the durable phrase - // search enforces on stored postings), NOT mere term presence. - let mut merged: Vec<(nodedb_types::Surrogate, f32, bool)> = - results.iter().map(|r| (r.doc_id, r.score, false)).collect(); - if let Some(txn_id) = task.request.txn_id - && let Err(e) = self.merge_fts_phrase_overlay_into_results( - FtsMergeParams { - txn_id, - database_id: task.request.database_id, - tid: tenant_id, - collection, - query: "", - top_k, - }, - terms, - &mut merged, - ) - { - return self.response_error(task, e); + /// Restrict candidates to the rows the residual filters and RLS admit, + /// match the phrase over the index with the transaction's hidden rows + /// excluded, add its staged matches, cut to `top_k`, and hydrate. + fn phrase_search_rows( + &self, + task: &ExecutionTask, + p: &PhraseSearchParams<'_>, + ) -> crate::Result> { + let tenant_id = TenantId::new(p.tid); + let database_id = task.request.database_id.as_u64(); + let Some(index) = self.text_index(task, p.tid, p.collection, p.field)? else { + return Ok(Vec::new()); + }; + let gate = TextRowGate::new(p.filters, p.rls_filters)?; + let eligible = combine_eligible( + p.prefilter, + self.text_eligible_rows(task, p.tid, p.collection, &gate)?, + ); + let staged = self.text_staged_view(task, p.tid, index)?; + // The phrase is analyzed once, with the collection's analyzer. Both + // the indexed and the staged match read this token sequence. + let phrase = self.analyze_phrase(database_id, tenant_id, p.collection, p.terms)?; + + let mut hits: Vec<(Surrogate, f32, bool)> = + if staged.as_ref().is_some_and(|view| view.hides_all()) { + Vec::new() + } else { + self.inverted + .phrase_search( + database_id, + tenant_id, + index, + crate::engine::sparse::inverted::PhraseSearchParams { + terms: &phrase, + top_k: p.top_k, + prefilter: eligible.as_ref(), + exclude: staged.as_ref().map(|view| view.hidden()), + }, + )? + .iter() + .map(|r| (r.doc_id, r.score, false)) + .collect() + }; + if let Some(view) = staged.as_ref() { + hits.extend(staged_phrase_hits(view, &phrase, eligible.as_ref())); + hits.sort_by(|a, b| { + b.1.partial_cmp(&a.1) + .unwrap_or(std::cmp::Ordering::Equal) + .then(a.0.cmp(&b.0)) + }); + hits.truncate(p.top_k); } - let rows = match self.hydrate_text_hits( - merged, + let scores = + self.text_score_columns(task, p.tid, p.collection, p.scores, eligible.as_ref())?; + self.hydrate_text_hits( + hits, HydrateTextHitsParams { - database_id: task.request.database_id.as_u64(), - tid, - collection, - top_k, - rls_filters: &[], + database_id, + tid: p.tid, + collection: p.collection, txn_id: task.request.txn_id, + scores: &scores, }, - ) { - Ok(rows) => rows, - Err(e) => return self.response_error(task, e), - }; - if let Some(ref m) = self.metrics { - m.record_fts_search(0); - } - match super::super::response_codec::encode(&rows) { - Ok(payload) => self.response_with_payload(task, payload), - Err(e) => self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ), - } + ) } - /// Execute a full-collection scan with BM25 score injected per row. - /// - /// Runs an FTS search to build a surrogate → score map, then scans every - /// document in the collection. Each document is returned with `score_alias` - /// injected as an additional field. Documents whose surrogate does not appear - /// in the score map receive `null` for the score column. + /// Execute a score scan: every row the residual filters and RLS admit, + /// each with its score columns. A row a column's index holds but its + /// query does not match carries `0.0` there. A row the index does not + /// hold carries `null`. pub(in crate::data::executor) fn execute_bm25_score_scan( &self, task: &ExecutionTask, - tid: u64, - collection: &str, - query: &str, - score_alias: &str, - fuzzy: bool, + params: ScoreScanParams<'_>, ) -> Response { - let tenant_id = TenantId::new(tid); - debug!(core = self.core_id, tid, %collection, %query, %score_alias, "bm25 score scan"); + debug!( + core = self.core_id, + tid = params.tid, + collection = %params.collection, + score_columns = params.scores.len(), + bound = ?params.bound, + "bm25 score scan" + ); - let _scan_guard = match self.acquire_scan_guard(task, tid, collection) { + let _scan_guard = match self.acquire_scan_guard(task, params.tid, params.collection) { Ok(g) => g, Err(resp) => return resp, }; - // Build a surrogate → score map from FTS hits. Bounded top_k: heap - // allocation in BMW search is `Vec::with_capacity(top_k)`, so a literal - // `usize::MAX` overflows on `top_k * size_of::()`. - let mut score_map: HashMap = match self.inverted.search( - task.request.database_id.as_u64(), - tenant_id, - collection, - FtsSearchParams { - query, - top_k: BM25_SCAN_MAX_HITS, - fuzzy_enabled: fuzzy, - mode: QueryMode::And, - prefilter: None, - }, - ) { - Ok(hits) => hits.into_iter().map(|h| (h.doc_id, h.score)).collect(), + let rows = match self.score_scan_rows(task, ¶ms) { + Ok(rows) => rows, Err(e) => return self.response_error(task, e), }; + self.text_rows_response(task, rows) + } - // Read-your-own-writes for FTS: fold staged document bodies into - // the score map (staged put re-scored/added, staged tombstone - // removed) before the collection scan renders rows below. - if let Some(txn_id) = task.request.txn_id - && let Err(e) = self.merge_fts_overlay_into_score_map( - FtsMergeParams { - txn_id, - database_id: task.request.database_id, - tid: tenant_id, - collection, - query, - top_k: BM25_SCAN_MAX_HITS, - }, - &mut score_map, - ) - { - return self.response_error(task, e); + /// Stream every current row the filters and RLS admit into the sink, + /// which scores rows in batches and keeps the bound. A row the issuing + /// transaction staged is emitted from its staged body. A staged tombstone + /// or TRUNCATE hides base rows. + /// + /// The admitted rows are found first when a gate applies: the score + /// columns decide their AND-mode fallback over exactly those rows. + fn score_scan_rows( + &self, + task: &ExecutionTask, + p: &ScoreScanParams<'_>, + ) -> crate::Result> { + let database_id = task.request.database_id; + let tenant = TenantId::new(p.tid); + let gate = TextRowGate::new(p.filters, p.rls_filters)?; + let eligible = self.text_eligible_rows(task, p.tid, p.collection, &gate)?; + let columns = + self.text_score_columns(task, p.tid, p.collection, p.scores, eligible.as_ref())?; + let mut sink = ScoreScanSink::new(&columns, p.bound); + if sink.is_full() { + return sink.finish(); } - // The body encoding of this collection's sparse rows, resolved from - // its registered kind. This scan returns EVERY row (a row with no FTS - // hit gets a null score), so it reaches vector-primary sidecars even - // when the inverted index holds nothing for the collection — and a - // sidecar decoded as a document body renders `[4,"alice"]`. - let format = self.sparse_body_format(task.request.database_id, tenant_id, collection); - - // Scan all documents and inject the score field. - let scan_result = self.sparse.scan_documents( - task.request.database_id.as_u64(), - tid, - collection, - BM25_SCAN_MAX_HITS, - ); - let mut docs: Vec<(StorageKey, Vec)> = match scan_result { - Ok(d) => d, - Err(e) => return self.response_error(task, e), + // its registered kind. A vector-primary sidecar decoded as a document + // body renders `[4,"alice"]`. + let format = self.sparse_body_format(database_id, tenant, p.collection); + let identity_column = self.identity_column(database_id.as_u64(), p.tid, p.collection); + let coll_key = (database_id, tenant, p.collection.to_string()); + let overlay = match task.request.txn_id { + Some(txn_id) => { + // Read-your-own-writes refreshes the lease (see the reaper). + self.touch_overlay(txn_id); + self.txn_overlays.get(&txn_id) + } + None => None, }; - // Read-your-own-writes: fold staged document bodies into the row - // list, gating staged-row membership on the FTS match. `score_map` - // above was already narrowed to matching surrogates (staged puts - // that match inserted, non-matches / tombstones removed), so a - // staged doc appears as a row ONLY when it is in `score_map` — - // `text_match(...)` as a predicate must not surface a staged doc - // that does not contain the query term. A staged tombstone or a - // staged update that dropped the term removes its base row too. - if let Some(txn_id) = task.request.txn_id { - self.merge_fts_rows_from_score_map( - FtsMergeParams { - txn_id, - database_id: task.request.database_id, - tid: tenant_id, - collection, - query, - top_k: BM25_SCAN_MAX_HITS, + let mut emit = |key: &StorageKey, bytes: &[u8]| -> crate::Result> { + if eligible + .as_ref() + .is_some_and(|rows| !rows.contains(key.surrogate())) + { + return Ok(ControlFlow::Continue(())); + } + let image = text_row_image(key, bytes, format.as_format_ref(), &identity_column); + let value = decode_scanned_row(bytes, Some(image.as_slice()), format.as_format_ref())?; + sink.push( + key.surrogate(), + DocumentRow { + id: key.to_string(), + data: value, }, - &mut docs, - &score_map, - ); - } + ) + }; - let mut rows: Vec = Vec::with_capacity(docs.len()); - for (key, bytes) in &docs { - let mut value = match decode_scanned_document(bytes, format.as_format_ref()) { - Ok(v) => v, - Err(e) => return self.response_error(task, e), - }; - // Inject score into the document object. - if let serde_json::Value::Object(ref mut map) = value { - let score = score_map.get(&key.surrogate()).copied(); - match score { - Some(s) => { - map.insert( - score_alias.to_string(), - serde_json::Value::Number( - serde_json::Number::from_f64(s as f64) - .unwrap_or_else(|| serde_json::Number::from(0)), - ), - ); - } - None => { - map.insert(score_alias.to_string(), serde_json::Value::Null); + let mut stopped = false; + if overlay.is_none_or(|o| o.base_visible(&coll_key)) { + self.for_each_text_base_row( + database_id.as_u64(), + p.tid, + p.collection, + |key, bytes| { + // A row the transaction staged is emitted from its staged body. + if overlay.is_some_and(|o| o.get(&coll_key, key.surrogate().as_u32()).is_some()) + { + return Ok(ControlFlow::Continue(())); } + let flow = emit(key, bytes)?; + stopped = flow.is_break(); + Ok(flow) + }, + )?; + } + if let Some(overlay) = overlay + && !stopped + { + for (surrogate, staged) in overlay.iter_for_collection(&coll_key) { + if let Staged::Put(body) = staged + && emit(&StorageKey::for_surrogate(Surrogate::new(surrogate)), body)?.is_break() + { + break; } } - rows.push(DocumentRow { - id: key.to_string(), - data: value, - }); } + sink.finish() + } + /// Encode hydrated text rows as the response payload. + fn text_rows_response(&self, task: &ExecutionTask, rows: Vec) -> Response { if let Some(ref m) = self.metrics { m.record_fts_search(0); } match super::super::response_codec::encode(&rows) { Ok(payload) => self.response_with_payload(task, payload), - Err(e) => self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ), + Err(e) => self.response_error(task, ErrorCode::from(e)), } } } diff --git a/nodedb/src/data/executor/handlers/text_search_triple.rs b/nodedb/src/data/executor/handlers/text_search_triple.rs index 5ab3cda69..89d0d0e9f 100644 --- a/nodedb/src/data/executor/handlers/text_search_triple.rs +++ b/nodedb/src/data/executor/handlers/text_search_triple.rs @@ -22,7 +22,7 @@ use super::hybrid_key::HybridFusionKey; use crate::bridge::envelope::{ErrorCode, Response}; use crate::data::executor::core_loop::CoreLoop; use crate::data::executor::handlers::graph_expansion::{GraphExpansionParams, GraphSeeds}; -use crate::data::executor::scan_normalize::sparse_body_to_msgpack; +use crate::data::executor::handlers::text_rows::{TextRowGate, combine_eligible}; use crate::data::executor::task::ExecutionTask; use crate::engine::graph::edge_store::Direction; use crate::query::fusion::{FusedResult, RankedResult, reciprocal_rank_fusion_weighted}; @@ -31,13 +31,23 @@ use crate::query::fusion::{FusedResult, RankedResult, reciprocal_rank_fusion_wei pub(in crate::data::executor) struct HybridSearchTripleParams<'a> { pub tid: u64, pub collection: &'a str, + /// Vector column the vector leg searches. Empty names the + /// collection-level index. + pub vector_field: &'a str, pub query_vector: &'a [f32], + /// Field index the text leg reads. `None` reads the whole-document index. + pub text_field: Option<&'a str>, pub query_text: &'a str, + /// Residual WHERE predicates (`Vec`), applied to every leg + /// before fusion. + pub filters: &'a [u8], pub graph_seed_id: &'a str, pub graph_depth: usize, pub graph_edge_label: Option<&'a str>, pub top_k: usize, pub ef_search: usize, + /// Boolean combination of the text leg's query terms. + pub mode: nodedb_types::text_search::QueryMode, pub fuzzy: bool, pub rrf_k: (f64, f64, f64), pub filter_bitmap: Option<&'a nodedb_types::SurrogateBitmap>, @@ -58,13 +68,17 @@ impl CoreLoop { let HybridSearchTripleParams { tid, collection, + vector_field, query_vector, + text_field, query_text, + filters, graph_seed_id, graph_depth, graph_edge_label, top_k, ef_search, + mode, fuzzy, rrf_k, filter_bitmap, @@ -90,9 +104,30 @@ impl CoreLoop { let fetch_k = top_k.saturating_mul(3).max(20); + // The rows the residual filters and RLS admit restrict every leg + // before fusion, so the fused top-k counts only admitted rows. + let eligible = match TextRowGate::new(filters, rls_filters) + .and_then(|gate| self.text_eligible_rows(task, tid, collection, &gate)) + { + Ok(eligible) => combine_eligible(filter_bitmap, eligible), + Err(e) => return self.response_error(task, e), + }; + let filter_bitmap = eligible.as_ref(); + let text_index = match self.text_index(task, tid, collection, text_field) { + Ok(index) => index, + Err(e) => return self.response_error(task, e), + }; + // 1. Vector search. - let index_key = - CoreLoop::vector_index_key(task.request.database_id.as_u64(), tid, collection, ""); + let index_key = match self.resolve_vector_index_key( + task.request.database_id.as_u64(), + tid, + collection, + vector_field, + ) { + Ok(key) => key, + Err(e) => return self.response_error(task, e), + }; let vector_collection = self.vector_collections.get(&index_key); let vector_results = match vector_collection { Some(index) => { @@ -115,17 +150,18 @@ impl CoreLoop { None => Vec::new(), }; - // 2. BM25 text search. - let text_results = match self.inverted.search( - task.request.database_id.as_u64(), - tenant_id, - collection, + // 2. BM25 text search over the field's index, restricted to the same + // rows, with the transaction's staged rows ranked in. + let text_results = match self.hybrid_text_leg( + task, + tid, + text_index, FtsSearchParams { query: query_text, top_k: fetch_k, fuzzy_enabled: fuzzy, - mode: QueryMode::And, - prefilter: None, + mode: QueryMode::from(mode), + prefilter: filter_bitmap, }, ) { Ok(results) => results, @@ -140,7 +176,7 @@ impl CoreLoop { database_id: task.request.database_id.as_u64(), tid, seeds: GraphSeeds::Names(&[graph_seed_id]), - label_filter: graph_edge_label, + label_filter: graph_edge_label.as_slice(), direction: Direction::Out, max_depth: graph_depth, max_visited: self.query_tuning.bfs_memory_budget_bytes @@ -155,7 +191,7 @@ impl CoreLoop { // vector/FTS overlay merges). The graph leg's RYOW is a separate // concern and is not folded in here. Outside a transaction the // committed-only construction below runs unchanged. - let (vector_ranked, text_ranked): super::hybrid_overlay::HybridRankedLegs = + let (mut vector_ranked, mut text_ranked): super::hybrid_overlay::HybridRankedLegs = if let Some(txn_id) = task.request.txn_id { match self.hybrid_ranked_with_overlay( super::hybrid_overlay::HybridOverlayParams { @@ -163,8 +199,8 @@ impl CoreLoop { database_id: task.request.database_id, tid: tenant_id, collection, + vector_field, query_vector, - query_text, fetch_k, filter_bitmap, }, @@ -200,7 +236,15 @@ impl CoreLoop { (vector_ranked, text_ranked) }; - let graph_ranked = graph_reached_to_ranked_keys(&expansion.reached); + // Each leg keeps only the rows the filters and RLS admit, judged + // before ranking on the transaction's view of each row. The fused + // top-k then counts only admitted rows. The graph walk reaches rows + // regardless of the filters. A headless vector hit or a graph node + // with no row is not admitted. + let mut graph_ranked = graph_reached_to_ranked_keys(&expansion.reached); + super::hybrid_key::retain_admitted(&mut vector_ranked, filter_bitmap); + super::hybrid_key::retain_admitted(&mut text_ranked, filter_bitmap); + super::hybrid_key::retain_admitted(&mut graph_ranked, filter_bitmap); let (k_vector, k_text, k_graph) = rrf_k; let fused = reciprocal_rank_fusion_weighted( @@ -209,42 +253,11 @@ impl CoreLoop { top_k, ); - // 5. Materialise results with per-engine rank diagnostics (reusing HybridSearchHit). - // - // The RLS predicate runs against the NORMALIZED msgpack image, never - // the stored bytes: a strict Binary Tuple is not a msgpack map at all - // and a vector-primary sidecar is a TAGGED one, so a predicate pushed - // at the stored bytes finds no field it recognizes and the row is - // dropped on a format mismatch rather than on policy. The encoding is - // resolved from the collection's registered kind — a tagged map and a - // plain document map share the same map header, so the bytes cannot - // answer it. - let body_format = self.sparse_body_format(task.request.database_id, tenant_id, collection); - // The fused key is rendered once per row here, at the response - // envelope; it is the only place the key becomes text. + // 5. Materialise results with per-engine rank diagnostics (reusing + // HybridSearchHit). The fused key is rendered once per row here, at + // the response envelope; it is the only place the key becomes text. let rendered: Vec<(String, &FusedResult)> = fused .iter() - .filter(|f| { - if rls_filters.is_empty() { - return true; - } - // A headless hit has no stored row to check the policy against, - // so it is treated the same as a row the lookup cannot find. - let Some(key) = f.document_id.storage_key() else { - return false; - }; - match self - .sparse - .get(task.request.database_id.as_u64(), tid, collection, &key) - { - Ok(Some(bytes)) => { - let normalized = - sparse_body_to_msgpack(&bytes, body_format.as_format_ref()); - super::rls_eval::rls_check_msgpack_bytes(rls_filters, &normalized) - } - _ => false, - } - }) .map(|f| (f.document_id.to_string(), f)) .collect(); let results: Vec<_> = rendered @@ -272,12 +285,7 @@ impl CoreLoop { } match super::super::response_codec::encode(&results) { Ok(payload) => self.response_with_payload(task, payload), - Err(e) => self.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ), + Err(e) => self.response_error(task, ErrorCode::from(e)), } } } diff --git a/nodedb/src/data/executor/handlers/transaction/overlay/fts_merge.rs b/nodedb/src/data/executor/handlers/transaction/overlay/fts_merge.rs index 52edf7587..4b8aaeae6 100644 --- a/nodedb/src/data/executor/handlers/transaction/overlay/fts_merge.rs +++ b/nodedb/src/data/executor/handlers/transaction/overlay/fts_merge.rs @@ -1,432 +1,130 @@ // SPDX-License-Identifier: BUSL-1.1 -//! Fold a transaction's staging overlay into a base full-text search result, -//! so an in-transaction FTS `SEARCH` observes the transaction's own -//! uncommitted document writes (read-your-own-writes for FTS). +//! An open transaction's staged writes as one inverted index sees them. //! -//! FTS indexing is not a stageable write in its own right — it is an inline -//! side effect of the document write -//! (`handlers/point/apply_put/core.rs::index_document_in_txn`). There is -//! therefore no staged FTS posting to read at query time. Instead, this -//! merge makes the transaction's already-staged DOCUMENT BODIES (held in -//! [`TxnOverlay`]) searchable by re-tokenizing and BM25-scoring them at -//! query time with the exact same analyzer resolution the forward indexing -//! path uses (`InvertedIndex::analyze_for_collection` — see -//! `IndexDocScope`/`index_document_in_txn`) and the SAME corpus stats -//! (`df` / `total_docs` / `avg_doc_len`) the base -//! search read, via [`InvertedIndex::corpus_stats`] / -//! [`InvertedIndex::term_df`], so a staged doc's score is directly -//! comparable to base-search scores. -//! -//! A collection's per-collection analyzer override -//! (`InvertedIndex::analyze_for_collection`, backed by -//! `FtsIndex::set_collection_analyzer`) is resolved for staged docs exactly -//! the same way the forward indexing path (`index_document_in_txn`) -//! resolves it, so a staged doc is tokenized identically whether it is -//! still staged or already committed — no second, inconsistent tokenization -//! path. - -use std::collections::HashMap; - -use nodedb_fts::posting::Bm25Params; -use nodedb_fts::search::query_parser::parse_query; -use nodedb_types::Surrogate; - +//! FTS indexing is an inline side effect of the document write, not a +//! stageable write of its own, so there is no staged posting to read. A +//! full-text read inside the transaction builds a [`StagedView`] instead: +//! every row the transaction staged hides its indexed version, a staged +//! TRUNCATE hides every indexed row, and each staged body that holds text in +//! the index enters the read with its analyzed tokens. The BM25 search, the +//! score columns, and the phrase search all read the transaction's writes +//! through this one view, before their top-k cut. + +use nodedb_fts::{IndexScope, StagedDoc, StagedView}; +use nodedb_types::{Surrogate, SurrogateBitmap}; + +use super::fts_score::earliest_contiguous_match; +use super::staged::Staged; use crate::data::executor::core_loop::CoreLoop; -use crate::data::executor::handlers::transaction::overlay::{Staged, TxnOverlay}; -use crate::engine::document::store::StorageKey; -use crate::types::{DatabaseId, TenantId, TxnId}; - -/// Scope + tuning for one FTS overlay merge: the transaction, the -/// `(database, tenant, collection)` it targets, the raw query text, and the -/// result truncation bound. Bundling these keeps the merge entry points to a -/// small argument count. -pub(in crate::data::executor) struct FtsMergeParams<'a> { - pub txn_id: TxnId, - pub database_id: DatabaseId, - pub tid: TenantId, - pub collection: &'a str, - pub query: &'a str, - pub top_k: usize, -} +use crate::data::executor::task::ExecutionTask; +use crate::types::TenantId; impl CoreLoop { - /// Merge staged writes described by `params` into `base_results` - /// (surrogate, score, fuzzy triples produced by the base BM25 search), - /// re-scoring staged puts against the query and removing staged - /// tombstones, then re-sorting by score descending and truncating to - /// `params.top_k`. - /// - /// No-op when the transaction has no overlay entries for this - /// collection. A staged put that also appears in `base_results` is - /// re-scored and replaces the base entry (an in-transaction UPDATE may - /// have changed whether/how strongly the document matches). A staged - /// put whose re-tokenized text does not match any positive query term, - /// or is excluded by a `NOT`/`-` negative term, contributes no entry - /// (and is removed if it was already present from base — e.g. an - /// UPDATE that made the document no longer match). - /// - /// Fails when a staged body will not decode: the merge exists so a - /// transaction sees its own writes, and a staged row that silently drops - /// out of the result is the exact opposite of that. - pub(in crate::data::executor) fn merge_fts_overlay_into_results( + /// The issuing transaction's view of `index`. `None` outside a + /// transaction, or when the transaction staged nothing for the + /// collection. + pub(in crate::data::executor) fn text_staged_view( &self, - params: FtsMergeParams<'_>, - base_results: &mut Vec<(Surrogate, f32, bool)>, - ) -> crate::Result<()> { - let FtsMergeParams { - txn_id, - database_id, - tid, - collection, - query, - top_k, - } = params; - let coll_key = (database_id, tid, collection.to_string()); - // Read-your-own-writes refreshes the lease (see the reaper). - self.touch_overlay(txn_id); - let Some(overlay) = self.txn_overlays.get(&txn_id) else { - return Ok(()); + task: &ExecutionTask, + tid: u64, + index: IndexScope<'_>, + ) -> crate::Result> { + let Some(txn_id) = task.request.txn_id else { + return Ok(None); }; - // A staged TRUNCATE hides every base hit; staged puts re-enter below. - if overlay.is_truncated(&coll_key) { - base_results.clear(); - } - - let (positive_terms, negative_terms) = - self.analyze_query_terms(database_id.as_u64(), tid, collection, query)?; - if positive_terms.is_empty() { - // No positive terms to score staged docs against — but staged - // tombstones still hide base rows. - remove_tombstoned(overlay, &coll_key, base_results); - return Ok(()); - } - - let config_key = (database_id, tid, collection.to_string()); - let bm25_params = Bm25Params::default(); - let ctx = self.staged_score_ctx(database_id, tid, collection, &config_key, &bm25_params)?; - - let mut seen: HashMap = base_results - .iter() - .enumerate() - .map(|(idx, (s, _, _))| (s.as_u32(), idx)) - .collect(); - - for (surrogate, staged) in overlay.iter_for_collection(&coll_key) { - match staged { - Staged::Tombstone => { - if let Some(idx) = seen.remove(&surrogate) { - base_results.remove(idx); - reindex_after_removal(&mut seen, idx); - } - } - Staged::Put(body) => { - let score = - self.score_staged_fts_doc(&ctx, body, &positive_terms, &negative_terms)?; - match (score, seen.get(&surrogate).copied()) { - (Some(s), Some(idx)) => { - base_results[idx].1 = s; - // Re-scored via the exact analyzed tokenizer, so - // this is no longer a fuzzy match regardless of - // whether the base entry was fuzzy. - base_results[idx].2 = false; - } - (Some(s), None) => { - seen.insert(surrogate, base_results.len()); - base_results.push((Surrogate::new(surrogate), s, false)); - } - (None, Some(idx)) => { - base_results.remove(idx); - seen.remove(&surrogate); - reindex_after_removal(&mut seen, idx); - } - (None, None) => {} - } - } - } - } - - base_results.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal)); - base_results.truncate(top_k); - Ok(()) - } - - /// Phrase-search variant of the overlay merge. Unlike the bag-of-words - /// BM25 merge, this honours the base phrase search's EXACT proximity - /// semantics: the phrase's analyzed terms must appear as a CONTIGUOUS, - /// in-order subsequence of the staged document's analyzed token stream - /// (adjacency with zero slop — the same `expected_pos = start + offset` - /// contiguity the durable phrase search enforces on stored postings). - /// A staged doc is included only if the phrase actually matches - /// positionally; a staged tombstone removes its base row. - /// - /// Positions come from the staged doc's own analyzed token indices — the - /// forward indexer assigns posting positions the same way (`enumerate` - /// over `InvertedIndex::analyze_for_collection` output), so no durable - /// posting list is needed to verify adjacency. Score mirrors the base - /// phrase formula at - /// rank 0 (`1 / (1 + earliest_start_pos)`), keeping staged matches - /// order-comparable with base phrase hits. - pub(in crate::data::executor) fn merge_fts_phrase_overlay_into_results( - &self, - params: FtsMergeParams<'_>, - terms: &[String], - base_results: &mut Vec<(Surrogate, f32, bool)>, - ) -> crate::Result<()> { - let FtsMergeParams { - txn_id, - database_id, - tid, - collection, - query: _query, - top_k, - } = params; - let coll_key = (database_id, tid, collection.to_string()); // Read-your-own-writes refreshes the lease (see the reaper). self.touch_overlay(txn_id); let Some(overlay) = self.txn_overlays.get(&txn_id) else { - return Ok(()); + return Ok(None); }; - // A staged TRUNCATE hides every base hit; staged puts re-enter below. - if overlay.is_truncated(&coll_key) { - base_results.clear(); - } - - // Canonicalize each phrase term through the collection's configured - // analyzer — the same resolution the base phrase search - // (`InvertedIndex::phrase_search`) uses — so the contiguity check - // compares stemmed/normalized tokens on both sides. - let db_u64 = database_id.as_u64(); - // A term the analyzer drops (a stop word) is matched as written. - let mut phrase_terms: Vec = Vec::with_capacity(terms.len()); - for t in terms { - let tokens = self - .inverted - .analyze_for_collection(db_u64, tid, collection, t)?; - phrase_terms.push(tokens.into_iter().next().unwrap_or_else(|| t.clone())); - } - - let config_key = (database_id, tid, collection.to_string()); - - let mut seen: HashMap = base_results - .iter() - .enumerate() - .map(|(idx, (s, _, _))| (s.as_u32(), idx)) - .collect(); - - for (surrogate, staged) in overlay.iter_for_collection(&coll_key) { - match staged { - Staged::Tombstone => { - if let Some(idx) = seen.remove(&surrogate) { - base_results.remove(idx); - reindex_after_removal(&mut seen, idx); - } - } - Staged::Put(body) => { - let score = - self.score_staged_phrase_doc(db_u64, &config_key, body, &phrase_terms)?; - match (score, seen.get(&surrogate).copied()) { - (Some(s), Some(idx)) => { - base_results[idx].1 = s; - base_results[idx].2 = false; - } - (Some(s), None) => { - seen.insert(surrogate, base_results.len()); - base_results.push((Surrogate::new(surrogate), s, false)); - } - (None, Some(idx)) => { - base_results.remove(idx); - seen.remove(&surrogate); - reindex_after_removal(&mut seen, idx); - } - (None, None) => {} - } - } - } - } - - base_results.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal)); - base_results.truncate(top_k); - Ok(()) - } - - /// Staged-doc scoring for `BM25ScoreScan`'s `HashMap` - /// score map, which has no top-k / ranked-Vec shape and no fuzzy flag. - pub(in crate::data::executor) fn merge_fts_overlay_into_score_map( - &self, - params: FtsMergeParams<'_>, - base_results: &mut HashMap, - ) -> crate::Result<()> { - let FtsMergeParams { - txn_id, - database_id, - tid, - collection, - query, - top_k: _top_k, - } = params; - let coll_key = (database_id, tid, collection.to_string()); - // Read-your-own-writes refreshes the lease (see the reaper). - self.touch_overlay(txn_id); - let Some(overlay) = self.txn_overlays.get(&txn_id) else { - return Ok(()); - }; - // A staged TRUNCATE hides every base hit; staged puts re-enter below. - if overlay.is_truncated(&coll_key) { - base_results.clear(); - } - let (positive_terms, negative_terms) = - self.analyze_query_terms(database_id.as_u64(), tid, collection, query)?; - - let config_key = (database_id, tid, collection.to_string()); - let bm25_params = Bm25Params::default(); - let ctx = self.staged_score_ctx(database_id, tid, collection, &config_key, &bm25_params)?; - - for (surrogate, staged) in overlay.iter_for_collection(&coll_key) { - match staged { - Staged::Tombstone => { - base_results.remove(&Surrogate::new(surrogate)); - } - Staged::Put(body) if !positive_terms.is_empty() => { - match self.score_staged_fts_doc(&ctx, body, &positive_terms, &negative_terms)? { - Some(s) => { - base_results.insert(Surrogate::new(surrogate), s); - } - None => { - base_results.remove(&Surrogate::new(surrogate)); - } - } - } - Staged::Put(_) => { - base_results.remove(&Surrogate::new(surrogate)); - } - } - } - Ok(()) - } - - /// Fold the overlay into `BM25ScoreScan`'s scanned row set - /// (`(hex_doc_id, body)` pairs) so staged-row MEMBERSHIP is gated on the - /// FTS match — a staged doc appears as a row ONLY when its surrogate is - /// in `score_map` (which the FTS score merge already narrowed to matches - /// and cleared of tombstones / non-matches). This keeps row membership - /// and scoring from ever disagreeing: a staged tombstone drops its base - /// row, a staged put that no longer matches (score-map absent) drops - /// its base row, and a staged put that matches is added with its staged - /// body. Base rows with no overlay entry are kept untouched (a - /// non-matching base row still projects with a `null` score, the - /// existing bm25-scan semantics). Run AFTER - /// [`merge_fts_overlay_into_score_map`](Self::merge_fts_overlay_into_score_map). - pub(in crate::data::executor) fn merge_fts_rows_from_score_map( - &self, - params: FtsMergeParams<'_>, - rows: &mut Vec<(StorageKey, Vec)>, - score_map: &HashMap, - ) { - let FtsMergeParams { - txn_id, + let database_id = task.request.database_id; + let coll_key = ( database_id, - tid, - collection, - .. - } = params; - let coll_key = (database_id, tid, collection.to_string()); - // Read-your-own-writes refreshes the lease (see the reaper). - self.touch_overlay(txn_id); - let Some(overlay) = self.txn_overlays.get(&txn_id) else { - return; - }; - - let mut seen: std::collections::HashSet = - rows.iter().map(|(k, _)| k.surrogate().as_u32()).collect(); - let base_visible = overlay.base_visible(&coll_key); - - rows.retain_mut(|(row_key, body)| { - let surrogate = row_key.surrogate().as_u32(); - match overlay.get(&coll_key, surrogate) { - Some(Staged::Tombstone) => false, - Some(Staged::Put(staged_body)) => { - *body = staged_body.clone(); - score_map.contains_key(&Surrogate::new(surrogate)) - } - None => base_visible, - } - }); - + TenantId::new(tid), + index.collection().to_string(), + ); + let hide_all = overlay.is_truncated(&coll_key); + let mut hidden = SurrogateBitmap::new(); + let mut docs = Vec::new(); for (surrogate, staged) in overlay.iter_for_collection(&coll_key) { - if seen.contains(&surrogate) { - continue; - } + let doc_id = Surrogate::new(surrogate); + hidden.insert(doc_id); if let Staged::Put(body) = staged - && score_map.contains_key(&Surrogate::new(surrogate)) + && let Some(tokens) = + self.tokenize_staged_body(database_id.as_u64(), &coll_key, index, body)? { - rows.push(( - StorageKey::for_surrogate(Surrogate::new(surrogate)), - body.clone(), - )); - seen.insert(surrogate); + docs.push(StagedDoc { doc_id, tokens }); } } + if !hide_all && hidden.is_empty() { + return Ok(None); + } + Ok(Some(StagedView::new(hidden, hide_all, docs))) } - /// Parse `query` and analyze its positive/negative terms with the - /// collection's configured analyzer (`InvertedIndex::analyze_for_collection` - /// — the same resolution the forward-indexing path and the base search - /// use), once per merge call (never per staged document — every staged - /// doc in the merge loop is scored against this same pair of term - /// lists). A query that fails to parse fails with `BadRequest`, as the - /// base search does. An analyzer error propagates. - fn analyze_query_terms( + /// The phrase `words` as the collection's analyzer tokenizes it: the + /// token sequence its indexed and staged text carry. A word the analyzer + /// drops (a stop word) has no token, as in the indexed text. Both the + /// indexed and the staged phrase match read this one sequence. + pub(in crate::data::executor) fn analyze_phrase( &self, database_id: u64, tid: TenantId, collection: &str, - query: &str, - ) -> crate::Result<(Vec, Vec)> { - let parsed = parse_query(query).map_err(|e| crate::Error::BadRequest { - detail: e.to_string(), - })?; - let positive_terms = self.inverted.analyze_for_collection( - database_id, - tid, - collection, - &parsed.positive.join(" "), - )?; - let negative_terms = self.inverted.analyze_for_collection( - database_id, - tid, - collection, - &parsed.negative.join(" "), - )?; - Ok((positive_terms, negative_terms)) + words: &[String], + ) -> crate::Result> { + self.inverted + .analyze_for_collection(database_id, tid, collection, &words.join(" ")) } } -/// Remove every staged tombstone's surrogate from `base_results` when there -/// are no positive query terms to score staged puts against (e.g. an -/// all-stop-word query). -fn remove_tombstoned( - overlay: &TxnOverlay, - coll_key: &(DatabaseId, TenantId, String), - base_results: &mut Vec<(Surrogate, f32, bool)>, -) { - let tombstoned: std::collections::HashSet = overlay - .iter_for_collection(coll_key) - .filter(|(_, staged)| matches!(staged, Staged::Tombstone)) - .map(|(surrogate, _)| surrogate) - .collect(); - if tombstoned.is_empty() { - return; - } - base_results.retain(|(s, _, _)| !tombstoned.contains(&s.as_u32())); +/// The staged documents `eligible` admits that hold `phrase` as a +/// contiguous, in-order run of their tokens (zero slop, the adjacency the +/// indexed phrase search enforces). Each scores `1 / (1 + earliest_start)`, +/// the indexed phrase formula at rank 0. +pub(in crate::data::executor) fn staged_phrase_hits( + staged: &StagedView, + phrase: &[String], + eligible: Option<&SurrogateBitmap>, +) -> Vec<(Surrogate, f32, bool)> { + staged + .docs() + .iter() + .filter(|doc| eligible.is_none_or(|bm| bm.contains(doc.doc_id))) + .filter_map(|doc| { + earliest_contiguous_match(&doc.tokens, phrase) + .map(|start| (doc.doc_id, 1.0 / (1.0 + start as f32), false)) + }) + .collect() } -/// After a `Vec::remove(idx)` shifts every later element left by one, shift -/// every recorded index greater than `idx` in `seen` to match. -fn reindex_after_removal(seen: &mut HashMap, removed_idx: usize) { - for idx in seen.values_mut() { - if *idx > removed_idx { - *idx -= 1; - } +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn staged_phrase_hits_respect_adjacency_and_eligibility() { + let view = StagedView::new( + SurrogateBitmap::new(), + false, + vec![ + StagedDoc { + doc_id: Surrogate::new(1), + tokens: vec!["quick".into(), "brown".into(), "fox".into()], + }, + StagedDoc { + doc_id: Surrogate::new(2), + tokens: vec!["brown".into(), "quick".into(), "fox".into()], + }, + ], + ); + let phrase = vec!["brown".to_string(), "fox".to_string()]; + let hits = staged_phrase_hits(&view, &phrase, None); + assert_eq!(hits, vec![(Surrogate::new(1), 0.5, false)]); + + let mut eligible = SurrogateBitmap::new(); + eligible.insert(Surrogate::new(2)); + assert!(staged_phrase_hits(&view, &phrase, Some(&eligible)).is_empty()); } } diff --git a/nodedb/src/data/executor/handlers/transaction/overlay/fts_score.rs b/nodedb/src/data/executor/handlers/transaction/overlay/fts_score.rs index 1772a7ebf..533fb6b30 100644 --- a/nodedb/src/data/executor/handlers/transaction/overlay/fts_score.rs +++ b/nodedb/src/data/executor/handlers/transaction/overlay/fts_score.rs @@ -1,209 +1,87 @@ // SPDX-License-Identifier: BUSL-1.1 -//! Staged-document scoring internals shared by the FTS overlay merges in -//! `fts_merge.rs`. Decodes a transaction's staged document body, re-tokenizes -//! it with the collection's configured analyzer — the SAME analyzer -//! resolution (`InvertedIndex::analyze_for_collection`) the forward indexing -//! path uses — and scores it against a query — either bag-of-words BM25 -//! (using the base index's corpus stats so scores are comparable) or -//! exact-adjacency phrase matching over the staged doc's own token positions. - -use std::collections::HashMap; - -use nodedb_fts::bm25::bm25_score; -use nodedb_fts::posting::Bm25Params; -use tracing::warn; +//! Staged-document text for the full-text reads of an open transaction. +//! +//! A transaction's document writes are not in the inverted index until it +//! commits. A full-text read inside the transaction tokenizes each staged +//! body with the collection's configured analyzer, the same resolution +//! (`InvertedIndex::analyze_for_collection`) the forward indexing path uses, +//! so a staged document is tokenized identically whether it is still staged +//! or already committed. use crate::data::executor::core_loop::CoreLoop; -use crate::data::executor::fts_text::extract_fts_text; +use crate::data::executor::fts_text::extract_fts_fields; use crate::types::{DatabaseId, TenantId}; - -/// Immutable corpus context for scoring one staged document: the collection -/// scope needed to read per-term `df`, the collection's storage config key, -/// and the shared corpus stats every staged doc is scored against. -pub(in crate::data::executor) struct StagedFtsScoreCtx<'a> { - database_id: u64, - tid: TenantId, - collection: &'a str, - config_key: &'a (DatabaseId, TenantId, String), - total_docs: u32, - avg_doc_len: f32, - bm25_params: &'a Bm25Params, -} +use nodedb_fts::IndexScope; impl CoreLoop { - /// Build the shared corpus context used to score every staged document - /// in one merge. Reads `(total_docs, avg_doc_len)` once from the same - /// stats the base search used. An empty base corpus is treated as a - /// 1-document corpus so `bm25_score`'s IDF term stays finite/positive. - pub(in crate::data::executor) fn staged_score_ctx<'a>( - &self, - database_id: DatabaseId, - tid: TenantId, - collection: &'a str, - config_key: &'a (DatabaseId, TenantId, String), - bm25_params: &'a Bm25Params, - ) -> crate::Result> { - let (total_docs, avg_doc_len) = - self.inverted - .corpus_stats(database_id.as_u64(), tid, collection)?; - Ok(StagedFtsScoreCtx { - database_id: database_id.as_u64(), - tid, - collection, - config_key, - total_docs: total_docs.max(1), - avg_doc_len: if avg_doc_len > 0.0 { avg_doc_len } else { 1.0 }, - bm25_params, - }) - } - - /// Decode a staged body, re-tokenize with the forward-indexing - /// tokenizer, and BM25-score it against `positive_terms`. - /// - /// `Ok(None)` is the "this document does not belong in the result" answer: - /// it extracts to empty text (the forward indexer never indexes such - /// documents either), matches no positive term, or is excluded by a - /// negative term. A staged body that will not decode is not that answer — - /// it is this transaction's own write becoming unreadable — and comes back - /// as `Err`. - pub(in crate::data::executor) fn score_staged_fts_doc( - &self, - ctx: &StagedFtsScoreCtx<'_>, - body: &[u8], - positive_terms: &[String], - negative_terms: &[String], - ) -> crate::Result> { - let Some(doc_tokens) = self.tokenize_staged_body(ctx.database_id, ctx.config_key, body)? - else { - return Ok(None); - }; - - if !negative_terms.is_empty() && negative_terms.iter().any(|t| doc_tokens.contains(t)) { - return Ok(None); - } - - let mut term_freq: HashMap<&str, u32> = HashMap::new(); - for token in &doc_tokens { - *term_freq.entry(token.as_str()).or_insert(0) += 1; - } - let doc_len = doc_tokens.len() as u32; - - let mut score = 0.0f32; - let mut matched_any = false; - for term in positive_terms { - let Some(&tf) = term_freq.get(term.as_str()) else { - continue; - }; - matched_any = true; - let df = self - .inverted - .term_df(ctx.database_id, ctx.tid, ctx.collection, term)? - .max(1); - score += bm25_score( - tf, - df, - doc_len, - ctx.total_docs, - ctx.avg_doc_len, - ctx.bm25_params, - ); - } - - if matched_any && score > 0.0 { - Ok(Some(score)) - } else { - Ok(None) - } - } - - /// Score a staged body for a PHRASE search: include it only when the - /// analyzed `phrase_terms` occur as a contiguous, in-order run in the - /// staged doc's analyzed token stream. Returns `1 / (1 + earliest_start)` - /// (base phrase formula at rank 0) on a match, `Ok(None)` when the phrase - /// is not present, and `Err` when the staged body will not decode. - pub(in crate::data::executor) fn score_staged_phrase_doc( - &self, - database_id: u64, - config_key: &(DatabaseId, TenantId, String), - body: &[u8], - phrase_terms: &[String], - ) -> crate::Result> { - if phrase_terms.is_empty() { - return Ok(None); - } - let Some(doc_tokens) = self.tokenize_staged_body(database_id, config_key, body)? else { - return Ok(None); - }; - Ok(earliest_contiguous_match(&doc_tokens, phrase_terms) - .map(|start| 1.0 / (1.0 + start as f32))) - } - - /// Decode a staged body via the collection's storage mode and analyze it - /// with the collection's configured analyzer — resolved through - /// [`InvertedIndex::analyze_for_collection`](crate::engine::sparse::inverted::InvertedIndex::analyze_for_collection), - /// the exact same lookup the forward indexing path - /// (`index_document_in_txn`) uses, so a staged doc is tokenized - /// identically to how it will be tokenized once committed. + /// Decode a staged body via the collection's storage mode and analyze the + /// text `index` holds (one field, or the whole document) with the + /// collection's analyzer. /// - /// `Ok(None)` means the document has nothing to tokenize: the collection is - /// unregistered, the extracted text is empty (which the forward indexer - /// also never indexes), or the analyzer produced no tokens. An undecodable - /// body is `Err`. Analyzer resolution failure stays a logged skip — it is a - /// backend metadata read error, not a statement about this document. - fn tokenize_staged_body( + /// `Ok(None)` means the document holds no text in `index`: the + /// collection is unregistered, the index's text is empty (which the + /// forward indexer never indexes either), or the analyzer produced no + /// tokens. An undecodable body or an analyzer resolution error is `Err`: + /// a staged row that silently drops out of the read is the opposite of + /// read-your-own-writes. + pub(in crate::data::executor) fn tokenize_staged_body( &self, database_id: u64, config_key: &(DatabaseId, TenantId, String), + index: IndexScope<'_>, body: &[u8], ) -> crate::Result>> { let Some(doc) = self.decode_indexed_body(config_key, body)? else { return Ok(None); }; - let text = extract_fts_text(&doc); + let fields = extract_fts_fields(&doc); + let text = fields.text_of(index); if text.is_empty() { return Ok(None); } let (_, tid, collection) = config_key; - let tokens = - match self - .inverted - .analyze_for_collection(database_id, *tid, collection, &text) - { - Ok(tokens) => tokens, - Err(e) => { - warn!( - error = %e, - %collection, - "staged FTS scoring: analyzer resolution failed; skipping doc" - ); - return Ok(None); - } - }; - if tokens.is_empty() { - Ok(None) - } else { - Ok(Some(tokens)) - } + let tokens = self + .inverted + .analyze_for_collection(database_id, *tid, collection, &text)?; + Ok((!tokens.is_empty()).then_some(tokens)) } } /// Return the earliest start index at which `phrase` occurs as a contiguous, /// in-order subsequence of `tokens`, or `None` if it never does. Adjacency /// is exact (zero slop), matching the durable phrase search. -fn earliest_contiguous_match(tokens: &[String], phrase: &[String]) -> Option { +pub(in crate::data::executor) fn earliest_contiguous_match( + tokens: &[String], + phrase: &[String], +) -> Option { if phrase.is_empty() || phrase.len() > tokens.len() { return None; } let last_start = tokens.len() - phrase.len(); - for start in 0..=last_start { - if tokens[start..start + phrase.len()] - .iter() - .zip(phrase) - .all(|(a, b)| a == b) - { - return Some(start as u32); - } + (0..=last_start) + .find(|start| { + tokens[*start..*start + phrase.len()] + .iter() + .zip(phrase) + .all(|(a, b)| a == b) + }) + .map(|start| u32::try_from(start).unwrap_or(u32::MAX)) +} + +#[cfg(test)] +mod tests { + use super::earliest_contiguous_match; + + fn words(text: &str) -> Vec { + text.split(' ').map(str::to_string).collect() + } + + #[test] + fn finds_the_earliest_contiguous_run() { + let tokens = words("a b c b c"); + assert_eq!(earliest_contiguous_match(&tokens, &words("b c")), Some(1)); + assert_eq!(earliest_contiguous_match(&tokens, &words("c a")), None); + assert_eq!(earliest_contiguous_match(&tokens, &[]), None); } - None } diff --git a/nodedb/src/data/executor/handlers/transaction/overlay/mod.rs b/nodedb/src/data/executor/handlers/transaction/overlay/mod.rs index 5023f2dba..14dc2c214 100644 --- a/nodedb/src/data/executor/handlers/transaction/overlay/mod.rs +++ b/nodedb/src/data/executor/handlers/transaction/overlay/mod.rs @@ -23,7 +23,7 @@ pub use array_staged::{ArrayTxnOverlay, StagedCellPut}; pub(in crate::data::executor) use columnar_merge::{ ColumnarMatchedRow, ColumnarOverlayMergeParams, decode_staged_row, }; -pub(in crate::data::executor) use fts_merge::FtsMergeParams; +pub(in crate::data::executor) use fts_merge::staged_phrase_hits; pub use graph_staged::{GraphCollKey, GraphTxnOverlay, NodeLabelDelta}; pub(in crate::data::executor) use merge::IndexOverlayMergeParams; pub use row_tags::BodyWrites; diff --git a/nodedb/src/engine/sparse/btree_scan_while.rs b/nodedb/src/engine/sparse/btree_scan_while.rs new file mode 100644 index 000000000..4d97c28ad --- /dev/null +++ b/nodedb/src/engine/sparse/btree_scan_while.rs @@ -0,0 +1,52 @@ +// SPDX-License-Identifier: BUSL-1.1 + +//! A streaming document scan its visitor can stop. + +use std::ops::ControlFlow; + +use nodedb_types::StorageKey; +use redb::{ReadableDatabase, ReadableTable}; + +use super::btree::{ + DOCUMENTS, KeyedTable, SparseEngine, coll_prefix, invalid_storage_key_err, redb_err, +}; + +impl SparseEngine { + /// Visit the documents of a collection in key order, one row in memory + /// at a time, until `f` returns `ControlFlow::Break`. Every redb error + /// and every `f` error is propagated. + pub fn scan_documents_while( + &self, + database_id: u64, + tenant_id: u64, + collection: &str, + mut f: F, + ) -> crate::Result<()> + where + F: FnMut(&StorageKey, &[u8]) -> crate::Result>, + { + let prefix = coll_prefix(database_id, tenant_id, collection); + let end = format!("{prefix}\u{ffff}"); + + let read_txn = self.db.begin_read().map_err(|e| redb_err("read txn", e))?; + let table = read_txn + .open_table(DOCUMENTS) + .map_err(|e| redb_err("open table", e))?; + let range = table + .range(prefix.as_str()..end.as_str()) + .map_err(|e| redb_err("doc range", e))?; + + for entry in range { + let entry = entry.map_err(|e| redb_err("doc entry", e))?; + let key = entry.0.value(); + let doc_id = key.strip_prefix(&prefix).unwrap_or(key); + let storage_key = StorageKey::parse(doc_id).ok_or_else(|| { + invalid_storage_key_err(KeyedTable::Documents, collection, doc_id) + })?; + if f(&storage_key, entry.1.value())?.is_break() { + break; + } + } + Ok(()) + } +} diff --git a/nodedb/src/engine/sparse/inverted/staged_search.rs b/nodedb/src/engine/sparse/inverted/staged_search.rs new file mode 100644 index 000000000..a93f50ea9 --- /dev/null +++ b/nodedb/src/engine/sparse/inverted/staged_search.rs @@ -0,0 +1,177 @@ +// SPDX-License-Identifier: BUSL-1.1 + +//! Transaction-aware BM25 reads: a ranked search that folds an open +//! transaction's staged documents in before the top-k cut, and per-document +//! score columns read by point lookups. + +use nodedb_fts::posting::TextSearchResult; +use nodedb_fts::{DocScore, DocScorer, FtsSearchParams, IndexScope, StagedView, TextQuery}; +use nodedb_types::{Surrogate, SurrogateBitmap, TenantId}; + +use super::core::InvertedIndex; +use super::errors::fts_index_err; +use crate::engine::sparse::fts_redb::RedbFtsBackend; + +/// Scores documents against one query on one index. +pub struct TextDocScorer<'a> { + inner: DocScorer<'a, RedbFtsBackend>, +} + +impl TextDocScorer<'_> { + /// The score of each of `docs`, parallel to `docs`. + pub fn score(&self, docs: &[Surrogate]) -> crate::Result> { + self.inner.score(docs).map_err(fts_index_err) + } +} + +impl InvertedIndex { + /// BM25 search with `staged`, the issuing transaction's view of the + /// index, folded in before the top-k cut. + pub fn search_staged<'a>( + &self, + database_id: u64, + tid: TenantId, + index: impl Into>, + params: FtsSearchParams<'_>, + staged: Option<&StagedView>, + ) -> crate::Result> { + self.inner + .search_staged(database_id, tid.as_u64(), index, params, staged) + .map_err(fts_index_err) + } + + /// A per-document scorer of `query` on `index`. `eligible` is the set of + /// rows the reading query admits, over which the AND-mode fallback is + /// decided. + pub fn doc_scorer<'a>( + &'a self, + database_id: u64, + tid: TenantId, + index: impl Into>, + query: TextQuery<'_>, + eligible: Option<&SurrogateBitmap>, + staged: Option, + ) -> crate::Result> { + self.inner + .doc_scorer(database_id, tid.as_u64(), index, query, eligible, staged) + .map(|inner| TextDocScorer { inner }) + .map_err(fts_index_err) + } +} + +#[cfg(test)] +mod tests { + use std::sync::Arc; + + use nodedb_fts::posting::QueryMode; + use nodedb_fts::{DocScore, FtsSearchParams, StagedDoc, StagedView, TextQuery}; + use nodedb_types::{Surrogate, SurrogateBitmap, TenantId}; + + use super::InvertedIndex; + use crate::engine::durability_gate::GatedDatabase; + use crate::engine::sparse::inverted::test_support::body; + + const DB: u64 = 0; + const T: TenantId = TenantId::new(1); + + fn open_temp() -> (InvertedIndex, tempfile::TempDir) { + let dir = tempfile::tempdir().unwrap(); + let path = dir.path().join("test-inverted.redb"); + let db = Arc::new(GatedDatabase::new(redb::Database::create(&path).unwrap())); + let idx = + InvertedIndex::open(db, crate::data::executor::core_loop::test_governor()).unwrap(); + (idx, dir) + } + + /// Search scores read from the redb posting table equal the point scores + /// of the same documents, and a document with no text is absent. + #[test] + fn redb_point_scores_equal_search_scores() { + let (idx, _dir) = open_temp(); + idx.index_document(DB, T, "docs", Surrogate::new(1), &body("rust systems rust")) + .unwrap(); + idx.index_document(DB, T, "docs", Surrogate::new(2), &body("rust web")) + .unwrap(); + idx.index_document(DB, T, "docs", Surrogate::new(3), &body("python")) + .unwrap(); + + let hits = idx + .search_staged( + DB, + T, + "docs", + FtsSearchParams { + query: "rust", + top_k: 10, + fuzzy_enabled: false, + mode: QueryMode::And, + prefilter: None, + }, + None, + ) + .unwrap(); + assert_eq!(hits.len(), 2); + let scorer = idx + .doc_scorer( + DB, + T, + "docs", + TextQuery { + query: "rust", + fuzzy_enabled: false, + mode: QueryMode::And, + }, + None, + None, + ) + .unwrap(); + let docs: Vec = hits.iter().map(|h| h.doc_id).collect(); + for (hit, score) in hits.iter().zip(scorer.score(&docs).unwrap()) { + assert_eq!(score, DocScore::Match(hit.score)); + } + assert_eq!( + scorer + .score(&[Surrogate::new(3), Surrogate::new(4)]) + .unwrap(), + vec![DocScore::Miss, DocScore::Absent] + ); + } + + /// A staged update that removes the top hit leaves the next one inside + /// the limit. + #[test] + fn staged_removal_of_the_top_hit_keeps_the_limit_full() { + let (idx, _dir) = open_temp(); + idx.index_document(DB, T, "docs", Surrogate::new(1), &body("rust rust rust")) + .unwrap(); + idx.index_document(DB, T, "docs", Surrogate::new(2), &body("rust lang")) + .unwrap(); + let mut hidden = SurrogateBitmap::new(); + hidden.insert(Surrogate::new(1)); + let view = StagedView::new( + hidden, + false, + vec![StagedDoc { + doc_id: Surrogate::new(1), + tokens: vec!["golang".into()], + }], + ); + let hits = idx + .search_staged( + DB, + T, + "docs", + FtsSearchParams { + query: "rust", + top_k: 1, + fuzzy_enabled: false, + mode: QueryMode::And, + prefilter: None, + }, + Some(&view), + ) + .unwrap(); + assert_eq!(hits.len(), 1); + assert_eq!(hits[0].doc_id, Surrogate::new(2)); + } +} diff --git a/nodedb/src/engine/sparse/mod.rs b/nodedb/src/engine/sparse/mod.rs index d4a5272ee..57491c7e1 100644 --- a/nodedb/src/engine/sparse/mod.rs +++ b/nodedb/src/engine/sparse/mod.rs @@ -4,6 +4,7 @@ pub mod btree; pub mod btree_index; pub mod btree_index_bulk; pub mod btree_scan; +pub mod btree_scan_while; pub mod btree_versioned; pub mod doc_cache; pub mod fts_redb; diff --git a/nodedb/tests/inproc/cases/cross_engine_three_way_fts_vector_doc.rs b/nodedb/tests/inproc/cases/cross_engine_three_way_fts_vector_doc.rs index d939cf026..1c9391db5 100644 --- a/nodedb/tests/inproc/cases/cross_engine_three_way_fts_vector_doc.rs +++ b/nodedb/tests/inproc/cases/cross_engine_three_way_fts_vector_doc.rs @@ -272,11 +272,15 @@ fn three_way_fts_vector_doc_bitmap() { nodedb_types::DatabaseId::DEFAULT, COLLECTION, ), + field: None, query: "learning".into(), top_k: 20, + mode: nodedb_types::text_search::QueryMode::And, fuzzy: false, prefilter: None, + filters: Vec::new(), rls_filters: Vec::new(), + scores: Vec::new(), }), ); diff --git a/nodedb/tests/inproc/cases/executor_tests/test_cross_engine_validation.rs b/nodedb/tests/inproc/cases/executor_tests/test_cross_engine_validation.rs index 7364c36ef..331754212 100644 --- a/nodedb/tests/inproc/cases/executor_tests/test_cross_engine_validation.rs +++ b/nodedb/tests/inproc/cases/executor_tests/test_cross_engine_validation.rs @@ -129,7 +129,7 @@ fn cross_model_query_vector_graph_relational() { &mut rx, PhysicalPlan::Graph(GraphOp::Hop { start_nodes: vec!["p0".into()], - edge_label: Some("CITES".into()), + edge_labels: vec!["CITES".into()], direction: Direction::Out, depth: 3, options: Default::default(), @@ -291,11 +291,15 @@ fn rrf_fusion_mathematically_correct() { query_text: "database systems".into(), top_k: 5, ef_search: 0, + mode: nodedb_types::text_search::QueryMode::And, fuzzy: true, vector_weight: 0.5, filter_bitmap: None, rls_filters: Vec::new(), score_alias: None, + vector_field: String::new(), + text_field: None, + filters: Vec::new(), }), ); assert_eq!(resp_equal.status, Status::Ok); @@ -316,11 +320,15 @@ fn rrf_fusion_mathematically_correct() { query_text: "database systems".into(), top_k: 5, ef_search: 0, + mode: nodedb_types::text_search::QueryMode::And, fuzzy: true, vector_weight: 0.9, filter_bitmap: None, rls_filters: Vec::new(), score_alias: None, + vector_field: String::new(), + text_field: None, + filters: Vec::new(), }), ); assert_eq!(resp_vec_heavy.status, Status::Ok); @@ -404,9 +412,13 @@ fn document_indexes_consistent_after_simulated_crash() { ), query: "database".into(), top_k: 10, + mode: nodedb_types::text_search::QueryMode::And, fuzzy: true, rls_filters: Vec::new(), prefilter: None, + field: None, + filters: Vec::new(), + scores: Vec::new(), }), ); let text_json = payload_json(&text_payload); @@ -447,9 +459,13 @@ fn document_indexes_consistent_after_simulated_crash() { ), query: "database".into(), top_k: 10, + mode: nodedb_types::text_search::QueryMode::And, fuzzy: true, rls_filters: Vec::new(), prefilter: None, + field: None, + filters: Vec::new(), + scores: Vec::new(), }), ); let text_after_json = payload_json(&text_after); @@ -470,9 +486,13 @@ fn document_indexes_consistent_after_simulated_crash() { ), query: "vector search".into(), top_k: 10, + mode: nodedb_types::text_search::QueryMode::And, fuzzy: true, rls_filters: Vec::new(), prefilter: None, + field: None, + filters: Vec::new(), + scores: Vec::new(), }), ); let text_a2_json = payload_json(&text_a2); diff --git a/nodedb/tests/inproc/cases/executor_tests/test_tenant_isolation_fulltext.rs b/nodedb/tests/inproc/cases/executor_tests/test_tenant_isolation_fulltext.rs index 3abca1590..5e032416e 100644 --- a/nodedb/tests/inproc/cases/executor_tests/test_tenant_isolation_fulltext.rs +++ b/nodedb/tests/inproc/cases/executor_tests/test_tenant_isolation_fulltext.rs @@ -58,9 +58,13 @@ fn fulltext_search_isolated() { collection: QualifiedCollection::new(DatabaseId::DEFAULT, "articles"), query: "quantum".into(), top_k: 5, + mode: nodedb_types::text_search::QueryMode::And, fuzzy: false, rls_filters: Vec::new(), prefilter: None, + field: None, + filters: Vec::new(), + scores: Vec::new(), }), ); assert_eq!(resp_a.status, Status::Ok); @@ -75,9 +79,13 @@ fn fulltext_search_isolated() { collection: QualifiedCollection::new(DatabaseId::DEFAULT, "articles"), query: "quantum".into(), top_k: 5, + mode: nodedb_types::text_search::QueryMode::And, fuzzy: false, rls_filters: Vec::new(), prefilter: None, + field: None, + filters: Vec::new(), + scores: Vec::new(), }), ); assert_eq!(resp_b.status, Status::Ok); diff --git a/nodedb/tests/inproc/cases/executor_tests/test_tenant_isolation_fulltext_negative.rs b/nodedb/tests/inproc/cases/executor_tests/test_tenant_isolation_fulltext_negative.rs index 7ea661a21..a6382c822 100644 --- a/nodedb/tests/inproc/cases/executor_tests/test_tenant_isolation_fulltext_negative.rs +++ b/nodedb/tests/inproc/cases/executor_tests/test_tenant_isolation_fulltext_negative.rs @@ -58,9 +58,13 @@ fn fulltext_cross_tenant_index_does_not_contaminate_search() { collection: QualifiedCollection::new(DatabaseId::DEFAULT, "articles"), query: "quantum".into(), top_k: 20, + mode: nodedb_types::text_search::QueryMode::And, fuzzy: false, rls_filters: Vec::new(), prefilter: None, + field: None, + filters: Vec::new(), + scores: Vec::new(), }), ); assert_eq!(resp_baseline.status, Status::Ok); @@ -102,9 +106,13 @@ fn fulltext_cross_tenant_index_does_not_contaminate_search() { collection: QualifiedCollection::new(DatabaseId::DEFAULT, "articles"), query: "quantum".into(), top_k: 20, + mode: nodedb_types::text_search::QueryMode::And, fuzzy: false, rls_filters: Vec::new(), prefilter: None, + field: None, + filters: Vec::new(), + scores: Vec::new(), }), ); assert_eq!(resp_after.status, Status::Ok); diff --git a/nodedb/tests/wire/cases/engine_surface_fts.rs b/nodedb/tests/wire/cases/engine_surface_fts.rs index 75d0becda..425a76730 100644 --- a/nodedb/tests/wire/cases/engine_surface_fts.rs +++ b/nodedb/tests/wire/cases/engine_surface_fts.rs @@ -49,7 +49,35 @@ async fn bm25_score_projection() { srv.exec("INSERT INTO fts_bm25 { id: 'b3', content: 'distributed database performance' }") .await .unwrap(); + srv.exec("INSERT INTO fts_bm25 { id: 'b4', other: 'no content field' }") + .await + .unwrap(); + + // A match scores its BM25 value. A row the `content` index holds that + // does not match scores 0. A row with no `content` text scores NULL. + let rows = srv + .query_rows( + "SELECT id, bm25_score(content, 'database') AS score \ + FROM fts_bm25 \ + ORDER BY score DESC NULLS LAST, id", + ) + .await + .unwrap(); + let mut top = ids(&rows[..2]); + top.sort_unstable(); + assert_eq!(top, vec!["b1", "b3"], "matches rank first: {rows:?}"); + for row in &rows[..2] { + let score: f64 = row[1].parse().expect("a match carries a score"); + assert!(score > 0.0, "a match scores above 0: {rows:?}"); + } + assert_eq!(rows[2][0], "b2", "the held miss follows: {rows:?}"); + let miss: f64 = rows[2][1].parse().expect("a held miss carries a score"); + assert_eq!(miss, 0.0, "a held miss scores 0: {rows:?}"); + assert_eq!(rows[3][0], "b4", "the unheld row is last: {rows:?}"); + assert!(rows[3][1].is_empty(), "an unheld row scores NULL: {rows:?}"); + // With no NULLS clause a DESC key places NULL first, as every ORDER BY + // key does. let rows = srv .query_rows( "SELECT id, bm25_score(content, 'database') AS score \ @@ -58,12 +86,23 @@ async fn bm25_score_projection() { ) .await .unwrap(); - assert!(rows.len() >= 2, "expected at least 2 rows"); - let first_id = &rows[0][0]; - assert!( - first_id == "b1" || first_id == "b3", - "expected database doc first, got {first_id}" + assert_eq!(rows[0][0], "b4", "NULL sorts first under DESC: {rows:?}"); + assert!(rows[0][1].is_empty(), "b4 scores NULL: {rows:?}"); + + // ASC places NULL last. + let rows = srv + .query_rows( + "SELECT id, bm25_score(content, 'database') AS score \ + FROM fts_bm25 \ + ORDER BY score", + ) + .await + .unwrap(); + assert_eq!( + rows[0][0], "b2", + "the 0 score sorts first under ASC: {rows:?}" ); + assert_eq!(rows[3][0], "b4", "NULL sorts last under ASC: {rows:?}"); } #[tokio::test(flavor = "multi_thread", worker_threads = 4)] @@ -140,3 +179,286 @@ async fn and_or_query_combinations() { let ids: Vec<&str> = rows.iter().map(|r| r[0].as_str()).collect(); assert!(ids.contains(&"l1"), "l1 should match: {ids:?}"); } + +/// Two documents whose `rust` sits in different fields. +async fn crossed_docs(srv: &TestServer, coll: &str) { + srv.exec(&format!( + "CREATE COLLECTION {coll} WITH (engine='document_schemaless')" + )) + .await + .unwrap(); + srv.exec(&format!( + "INSERT INTO {coll} {{ id: 'd1', title: 'rust guide', body: 'an introduction' }}" + )) + .await + .unwrap(); + srv.exec(&format!( + "INSERT INTO {coll} {{ id: 'd2', title: 'cooking', body: 'the rust belt' }}" + )) + .await + .unwrap(); +} + +fn ids(rows: &[Vec]) -> Vec<&str> { + rows.iter().map(|r| r[0].as_str()).collect() +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn text_match_reads_only_the_named_field() { + let srv = TestServer::start().await; + crossed_docs(&srv, "fts_scoped").await; + + let rows = srv + .query_rows("SELECT id FROM fts_scoped WHERE text_match(title, 'rust') ORDER BY id") + .await + .unwrap(); + assert_eq!(ids(&rows), vec!["d1"], "title scope must ignore body text"); + + let rows = srv + .query_rows("SELECT id FROM fts_scoped WHERE text_match(body, 'rust') ORDER BY id") + .await + .unwrap(); + assert_eq!(ids(&rows), vec!["d2"], "body scope must ignore title text"); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn text_match_star_reads_the_whole_document() { + let srv = TestServer::start().await; + crossed_docs(&srv, "fts_star").await; + + let rows = srv + .query_rows("SELECT id FROM fts_star WHERE text_match(*, 'rust') ORDER BY id") + .await + .unwrap(); + assert_eq!(ids(&rows), vec!["d1", "d2"]); +} + +/// Rows for the score-column tests: `a` matches both, `b` only the body, +/// `c` only the title. +async fn score_docs(srv: &TestServer, coll: &str) { + srv.exec(&format!( + "CREATE COLLECTION {coll} WITH (engine='document_schemaless')" + )) + .await + .unwrap(); + for (id, title, body) in [ + ("a", "xenon lamp", "yellow light"), + ("b", "plain", "yellow paint"), + ("c", "xenon gas", "nothing here"), + ] { + srv.exec(&format!( + "INSERT INTO {coll} {{ id: '{id}', title: '{title}', body: '{body}' }}" + )) + .await + .unwrap(); + } +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn a_score_on_another_field_is_zero_where_its_text_does_not_match() { + let srv = TestServer::start().await; + score_docs(&srv, "fts_cross_score").await; + + let rows = srv + .query_rows( + "SELECT id, bm25_score(title, 'xenon') AS s FROM fts_cross_score \ + WHERE text_match(body, 'yellow') ORDER BY id", + ) + .await + .unwrap(); + assert_eq!(ids(&rows), vec!["a", "b"], "only body matches: {rows:?}"); + let a_score: f64 = rows[0][1].parse().expect("a holds xenon in its title"); + assert!(a_score > 0.0, "a's title score must be positive: {rows:?}"); + let b_score: f64 = rows[1][1].parse().expect("b holds a title"); + assert_eq!(b_score, 0.0, "b's title lacks xenon: {rows:?}"); +} + +/// A score column beside the match it scores keeps the match shape: only +/// matching rows return. +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn a_score_beside_its_own_match_returns_only_matches() { + let srv = TestServer::start().await; + score_docs(&srv, "fts_same_score").await; + + let rows = srv + .query_rows( + "SELECT id, bm25_score(body, 'yellow') AS s FROM fts_same_score \ + WHERE text_match(body, 'yellow') ORDER BY id", + ) + .await + .unwrap(); + assert_eq!(ids(&rows), vec!["a", "b"], "c must not return: {rows:?}"); + for row in &rows { + assert!( + !row[1].is_empty(), + "every match carries its score: {rows:?}" + ); + } +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn a_residual_filter_restricts_matches_before_the_limit() { + let srv = TestServer::start().await; + srv.exec("CREATE COLLECTION fts_filtered WITH (engine='document_schemaless')") + .await + .unwrap(); + for i in 0..10 { + let tag = if i % 3 == 0 { "a" } else { "b" }; + srv.exec(&format!( + "INSERT INTO fts_filtered {{ id: 'w{i}', tag: '{tag}', body: 'widget number {i}' }}" + )) + .await + .unwrap(); + } + + let rows = srv + .query_rows( + "SELECT id, tag FROM fts_filtered \ + WHERE text_match(body, 'widget') AND tag = 'a' LIMIT 3", + ) + .await + .unwrap(); + assert_eq!(rows.len(), 3, "4 tagged matches exist, LIMIT 3: {rows:?}"); + assert!( + rows.iter().all(|r| r[1] == "a"), + "every row is tag a: {rows:?}" + ); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn a_match_with_no_limit_returns_every_match() { + let srv = TestServer::start().await; + srv.exec("CREATE COLLECTION fts_unbounded WITH (engine='document_schemaless')") + .await + .unwrap(); + for chunk in 0..3 { + let values: Vec = (0..500) + .map(|i| { + let n = chunk * 500 + i; + format!("('m{n}', 'marker row {n}')") + }) + .collect(); + srv.exec(&format!( + "INSERT INTO fts_unbounded (id, body) VALUES {}", + values.join(", ") + )) + .await + .unwrap(); + } + + let rows = srv + .query_rows("SELECT id FROM fts_unbounded WHERE text_match(body, 'marker')") + .await + .unwrap(); + assert_eq!(rows.len(), 1500, "no LIMIT must not cap the matches"); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn a_score_scan_keeps_its_filter_order_and_limit() { + let srv = TestServer::start().await; + srv.exec("CREATE COLLECTION fts_scan_tail WITH (engine='document_schemaless')") + .await + .unwrap(); + for (id, tag, body) in [ + ("k1", "b", "xylophone music"), + ("k2", "a", "plain text"), + ("k3", "a", "xylophone lesson"), + ("k4", "a", "another row"), + ] { + srv.exec(&format!( + "INSERT INTO fts_scan_tail {{ id: '{id}', tag: '{tag}', body: '{body}' }}" + )) + .await + .unwrap(); + } + + let rows = srv + .query_rows( + "SELECT id, bm25_score(body, 'xylophone') FROM fts_scan_tail \ + WHERE tag = 'a' ORDER BY id LIMIT 2", + ) + .await + .unwrap(); + assert_eq!( + ids(&rows), + vec!["k2", "k3"], + "first two tagged rows: {rows:?}" + ); + let k2: f64 = rows[0][1].parse().expect("k2 holds body text"); + assert_eq!(k2, 0.0, "k2 lacks the term: {rows:?}"); + let k3: f64 = rows[1][1].parse().expect("k3 holds body text"); + assert!(k3 > 0.0, "k3 holds the term: {rows:?}"); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn a_field_no_document_holds_is_an_undefined_column() { + let srv = TestServer::start().await; + crossed_docs(&srv, "fts_ghost").await; + + let err = srv + .query_text("SELECT id FROM fts_ghost WHERE text_match(ghost, 'rust')") + .await + .expect_err("a field no document holds must be refused"); + assert!(err.contains("42703"), "expected 42703: {err}"); + assert!(err.contains("ghost"), "the message names the field: {err}"); + assert!( + err.contains("fts_ghost"), + "the message names the collection: {err}" + ); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn a_literal_column_argument_is_a_datatype_mismatch() { + let srv = TestServer::start().await; + crossed_docs(&srv, "fts_literal").await; + + let err = srv + .query_text("SELECT id FROM fts_literal WHERE text_match('lit', 'rust')") + .await + .expect_err("a literal names no column"); + assert!(err.contains("42804"), "expected 42804: {err}"); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn an_order_by_score_with_a_foreign_qualifier_is_an_unknown_table() { + let srv = TestServer::start().await; + crossed_docs(&srv, "fts_foreign").await; + + let err = srv + .query_text("SELECT id FROM fts_foreign ORDER BY bm25_score(zz.body, 'rust')") + .await + .expect_err("zz names no relation"); + assert!(err.contains("zz"), "the error names the qualifier: {err}"); + assert!(err.contains("42P01"), "expected undefined_table: {err}"); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn order_by_score_desc_lists_matches_first() { + let srv = TestServer::start().await; + score_docs(&srv, "fts_rank_order").await; + srv.exec("INSERT INTO fts_rank_order { id: 'd', body: 'no title here' }") + .await + .unwrap(); + srv.exec("INSERT INTO fts_rank_order { id: 'e', title: 'xenon xenon', body: 'bright' }") + .await + .unwrap(); + + // `e` holds the term twice in a title as short as the others, so it + // scores highest. `d` has no title, so it scores NULL: `NULLS LAST` + // keeps it behind the matches under DESC. + let rows = srv + .query_rows( + "SELECT id FROM fts_rank_order \ + ORDER BY bm25_score(title, 'xenon') DESC NULLS LAST LIMIT 3", + ) + .await + .unwrap(); + let got = ids(&rows); + assert_eq!(got.len(), 3, "three xenon titles exist: {rows:?}"); + assert_eq!(got[0], "e", "the strongest match ranks first: {rows:?}"); + assert!( + got[1..].iter().all(|id| *id == "a" || *id == "c"), + "the tied matches follow, before the unmatched rows: {rows:?}" + ); + assert_ne!(got[1], got[2], "each tied match appears once: {rows:?}"); +} diff --git a/nodedb/tests/wire/cases/engine_surface_fts_options.rs b/nodedb/tests/wire/cases/engine_surface_fts_options.rs new file mode 100644 index 000000000..dc512ae33 --- /dev/null +++ b/nodedb/tests/wire/cases/engine_surface_fts_options.rs @@ -0,0 +1,175 @@ +// SPDX-License-Identifier: BUSL-1.1 + +//! Named query options of the text-search calls over pgwire: +//! `text_match(col, 'q', mode => 'and' | 'or', fuzzy => true | false)` and +//! the same options on `bm25_score`. Omitted options run +//! `TextSearchParams::default()`: `mode => 'or'`, `fuzzy => false`. + +use crate::harness::TestServer; + +/// `m1` holds both query terms, `m2` only `machine`, `m3` neither. +async fn mode_docs(srv: &TestServer, coll: &str) { + srv.exec(&format!( + "CREATE COLLECTION {coll} WITH (engine='document_schemaless')" + )) + .await + .unwrap(); + for (id, body) in [ + ("m1", "machine learning is everywhere"), + ("m2", "machine shop tools"), + ("m3", "garden flowers"), + ] { + srv.exec(&format!( + "INSERT INTO {coll} {{ id: '{id}', body: '{body}' }}" + )) + .await + .unwrap(); + } +} + +fn ids(rows: &[Vec]) -> Vec<&str> { + rows.iter().map(|r| r[0].as_str()).collect() +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn mode_or_and_mode_and_return_different_rows() { + let srv = TestServer::start().await; + mode_docs(&srv, "fts_mode").await; + + let and_rows = srv + .query_rows( + "SELECT id FROM fts_mode \ + WHERE text_match(body, 'machine learning', mode => 'and') ORDER BY id", + ) + .await + .unwrap(); + assert_eq!( + ids(&and_rows), + vec!["m1"], + "And keeps the row with both terms" + ); + + let or_rows = srv + .query_rows( + "SELECT id FROM fts_mode \ + WHERE text_match(body, 'machine learning', mode => 'or') ORDER BY id", + ) + .await + .unwrap(); + assert_eq!( + ids(&or_rows), + vec!["m1", "m2"], + "Or keeps a row with any term" + ); + + // No option runs the default mode, `or`. + let default_rows = srv + .query_rows( + "SELECT id FROM fts_mode WHERE text_match(body, 'machine learning') ORDER BY id", + ) + .await + .unwrap(); + assert_eq!(ids(&default_rows), vec!["m1", "m2"]); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn bm25_score_mode_scores_a_partial_match_only_under_or() { + let srv = TestServer::start().await; + mode_docs(&srv, "fts_score_mode").await; + + let rows = srv + .query_rows( + "SELECT id, bm25_score(body, 'machine learning', mode => 'and') AS s \ + FROM fts_score_mode ORDER BY id", + ) + .await + .unwrap(); + let score = |rows: &[Vec], id: &str| -> f64 { + let row = rows + .iter() + .find(|r| r[0] == id) + .unwrap_or_else(|| panic!("row {id} missing: {rows:?}")); + row[1] + .parse() + .unwrap_or_else(|e| panic!("score {:?}: {e}", row[1])) + }; + assert!(score(&rows, "m1") > 0.0, "{rows:?}"); + assert_eq!( + score(&rows, "m2"), + 0.0, + "And does not score a partial match" + ); + + let rows = srv + .query_rows( + "SELECT id, bm25_score(body, 'machine learning', mode => 'or') AS s \ + FROM fts_score_mode ORDER BY id", + ) + .await + .unwrap(); + assert!( + score(&rows, "m2") > 0.0, + "Or scores a partial match: {rows:?}" + ); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn fuzzy_on_matches_a_typo_and_fuzzy_off_does_not() { + let srv = TestServer::start().await; + mode_docs(&srv, "fts_fuzzy_opt").await; + + let off = srv + .query_rows( + "SELECT id FROM fts_fuzzy_opt WHERE text_match(body, 'machime', fuzzy => false)", + ) + .await + .unwrap(); + assert!(off.is_empty(), "no exact match for a typo: {off:?}"); + + // No option runs the default, `fuzzy => false`. + let default_rows = srv + .query_rows("SELECT id FROM fts_fuzzy_opt WHERE text_match(body, 'machime')") + .await + .unwrap(); + assert!(default_rows.is_empty(), "{default_rows:?}"); + + let on = srv + .query_rows( + "SELECT id FROM fts_fuzzy_opt \ + WHERE text_match(body, 'machime', fuzzy => true) ORDER BY id", + ) + .await + .unwrap(); + assert_eq!(ids(&on), vec!["m1", "m2"], "fuzzy matches the typo"); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn malformed_options_are_refused() { + let srv = TestServer::start().await; + mode_docs(&srv, "fts_bad_opt").await; + + for (sql, needle) in [ + ( + "SELECT id FROM fts_bad_opt WHERE text_match(body, 'x', boost => 2)", + "unknown text-search option", + ), + ( + "SELECT id FROM fts_bad_opt WHERE text_match(body, 'x', mode = 'or')", + "=>", + ), + ( + "SELECT id FROM fts_bad_opt WHERE text_match(body, 'x', 'or')", + "third positional argument", + ), + ( + "SELECT id FROM fts_bad_opt WHERE text_match(body, 'x', mode => 'xor')", + "'or' or 'and'", + ), + ] { + let err = srv + .query_text(sql) + .await + .expect_err("a malformed option must be refused"); + assert!(err.contains(needle), "{sql}: expected {needle:?} in {err}"); + } +} diff --git a/nodedb/tests/wire/cases/fts_query_semantics.rs b/nodedb/tests/wire/cases/fts_query_semantics.rs new file mode 100644 index 000000000..f62b2f60e --- /dev/null +++ b/nodedb/tests/wire/cases/fts_query_semantics.rs @@ -0,0 +1,379 @@ +// SPDX-License-Identifier: BUSL-1.1 + +//! Full-text query semantics over the wire: negation and AND matching +//! before the LIMIT, RLS folded in before ranking, a transaction's staged +//! rows ranked before the cut, `id` filters on key-only rows, redaction +//! refusal, bitemporal score scans, and TRUNCATE emptying the index. + +use crate::harness::TestServer; + +const PASSWORD: &str = "fts-semantics-secret-3"; + +fn ids(rows: &[Vec]) -> Vec<&str> { + rows.iter().map(|r| r[0].as_str()).collect() +} + +async fn create(srv: &TestServer, coll: &str) { + srv.exec(&format!( + "CREATE COLLECTION {coll} WITH (engine='document_schemaless')" + )) + .await + .unwrap(); +} + +async fn insert(srv: &TestServer, coll: &str, id: &str, extra: &str) { + srv.exec(&format!("INSERT INTO {coll} {{ id: '{id}', {extra} }}")) + .await + .unwrap(); +} + +/// Run `sql` as `user`, returning each row's first cell, or the error text. +async fn first_cells_as(srv: &TestServer, user: &str, sql: &str) -> Result, String> { + let (client, handle) = srv + .connect_as(user, PASSWORD) + .await + .unwrap_or_else(|e| panic!("connect as {user}: {e}")); + let result = match client.simple_query(sql).await { + Ok(messages) => Ok(messages + .iter() + .filter_map(|message| match message { + tokio_postgres::SimpleQueryMessage::Row(row) => { + Some(row.get(0).unwrap_or("").to_string()) + } + _ => None, + }) + .collect()), + Err(e) => Err(e.to_string()), + }; + drop(client); + handle.abort(); + result +} + +async fn create_reader(srv: &TestServer, user: &str) { + srv.exec(&format!("CREATE USER {user} PASSWORD '{PASSWORD}'")) + .await + .unwrap(); + srv.exec(&format!("GRANT ROLE readwrite TO {user}")) + .await + .unwrap(); +} + +/// The best `rust` rows hold `python`. Negation drops them before the cut, +/// so `LIMIT 2` still returns the two surviving rows. +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn negated_terms_are_excluded_before_the_limit() { + let srv = TestServer::start().await; + create(&srv, "fts_not_limit").await; + for i in 0..4 { + insert( + &srv, + "fts_not_limit", + &format!("p{i}"), + "body: 'rust rust rust python'", + ) + .await; + } + insert(&srv, "fts_not_limit", "k1", "body: 'rust tooling'").await; + insert(&srv, "fts_not_limit", "k2", "body: 'rust compiler'").await; + + let rows = srv + .query_rows("SELECT id FROM fts_not_limit WHERE text_match(body, 'rust -python') LIMIT 2") + .await + .unwrap(); + let mut got = ids(&rows); + got.sort_unstable(); + assert_eq!( + got, + vec!["k1", "k2"], + "two rows survive the negation: {rows:?}" + ); +} + +/// Many rows out-score the single AND match on one word. The AND match is +/// found under `LIMIT 3`, and the query does not fall back to OR. +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn an_and_match_is_found_past_single_word_out_scorers() { + let srv = TestServer::start().await; + create(&srv, "fts_and_limit").await; + for i in 0..30 { + insert( + &srv, + "fts_and_limit", + &format!("a{i}"), + "body: 'alpha alpha alpha'", + ) + .await; + insert( + &srv, + "fts_and_limit", + &format!("b{i}"), + "body: 'bravo bravo bravo'", + ) + .await; + } + insert( + &srv, + "fts_and_limit", + "both", + "body: 'alpha bravo and more words'", + ) + .await; + + let rows = srv + .query_rows( + "SELECT id FROM fts_and_limit \ + WHERE text_match(body, 'alpha bravo', mode => 'and') LIMIT 3", + ) + .await + .unwrap(); + assert_eq!(ids(&rows), vec!["both"], "only the AND match: {rows:?}"); +} + +/// A read policy admits half the matches. `LIMIT 3` returns three admitted +/// rows, ranked among admitted rows only. +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn rls_restricts_matches_before_the_limit() { + let srv = TestServer::start().await; + let user = "fts_rls_reader"; + create(&srv, "fts_rls_limit").await; + for i in 0..8 { + let owner = if i % 2 == 0 { user } else { "other" }; + // The rows the policy hides hold the term more often, so they would + // fill the limit if the policy ran after ranking. + let body = if i % 2 == 0 { + "widget" + } else { + "widget widget widget" + }; + insert( + &srv, + "fts_rls_limit", + &format!("w{i}"), + &format!( + "owner: '{owner}', tag: '{}', body: '{body}'", + if i < 4 { "x" } else { "y" } + ), + ) + .await; + } + create_reader(&srv, user).await; + srv.exec( + "CREATE RLS POLICY fts_rls_owner ON fts_rls_limit FOR READ \ + USING (owner = $auth.username)", + ) + .await + .unwrap(); + + let rows = first_cells_as( + &srv, + user, + "SELECT id FROM fts_rls_limit WHERE text_match(body, 'widget') LIMIT 3", + ) + .await + .unwrap(); + assert_eq!(rows.len(), 3, "three admitted matches exist: {rows:?}"); + for id in &rows { + assert!( + ["w0", "w2", "w4", "w6"].contains(&id.as_str()), + "{id} is admitted: {rows:?}" + ); + } + + // With a residual filter too, both restrict the candidates. + let rows = first_cells_as( + &srv, + user, + "SELECT id FROM fts_rls_limit WHERE text_match(body, 'widget') AND tag = 'x' LIMIT 5", + ) + .await + .unwrap(); + let mut got: Vec<&str> = rows.iter().map(String::as_str).collect(); + got.sort_unstable(); + assert_eq!(got, vec!["w0", "w2"], "admitted rows with tag x: {rows:?}"); +} + +/// A staged update that removes the best hit leaves `LIMIT 2` full, and a +/// staged AND query keeps AND semantics. +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn staged_rows_are_ranked_before_the_limit() { + let srv = TestServer::start().await; + create(&srv, "fts_txn_limit").await; + insert(&srv, "fts_txn_limit", "top", "body: 'rust rust rust rust'").await; + insert(&srv, "fts_txn_limit", "mid", "body: 'rust rust lang'").await; + insert( + &srv, + "fts_txn_limit", + "low", + "body: 'rust tooling compiler words'", + ) + .await; + + srv.exec("BEGIN").await.unwrap(); + srv.exec("UPDATE fts_txn_limit SET body = 'golang only' WHERE id = 'top'") + .await + .unwrap(); + let rows = srv + .query_rows("SELECT id FROM fts_txn_limit WHERE text_match(body, 'rust') LIMIT 2") + .await + .unwrap(); + let mut got = ids(&rows); + got.sort_unstable(); + assert_eq!(got, vec!["low", "mid"], "the limit stays full: {rows:?}"); + + srv.exec("INSERT INTO fts_txn_limit { id: 'one', body: 'rust alone' }") + .await + .unwrap(); + srv.exec("INSERT INTO fts_txn_limit { id: 'both', body: 'rust lang pairing' }") + .await + .unwrap(); + let rows = srv + .query_rows( + "SELECT id FROM fts_txn_limit \ + WHERE text_match(body, 'rust lang', mode => 'and') ORDER BY id", + ) + .await + .unwrap(); + assert_eq!( + ids(&rows), + vec!["both", "mid"], + "AND matches only: {rows:?}" + ); + srv.exec("ROLLBACK").await.unwrap(); +} + +/// A row inserted without an `id` carries its identity in its key only. A +/// filter on `id` still finds it. +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn an_id_filter_matches_a_key_only_row() { + let srv = TestServer::start().await; + create(&srv, "fts_key_id").await; + srv.exec("INSERT INTO fts_key_id { body: 'quartz crystal' }") + .await + .unwrap(); + srv.exec("INSERT INTO fts_key_id { body: 'quartz watch' }") + .await + .unwrap(); + + let rows = srv + .query_rows("SELECT id FROM fts_key_id WHERE text_match(body, 'crystal')") + .await + .unwrap(); + assert_eq!(rows.len(), 1, "one crystal row: {rows:?}"); + let id = rows[0][0].clone(); + assert!(!id.is_empty(), "the row carries its identity: {rows:?}"); + + let rows = srv + .query_rows(&format!( + "SELECT id FROM fts_key_id WHERE text_match(body, 'quartz') AND id = '{id}'" + )) + .await + .unwrap(); + assert_eq!( + ids(&rows), + vec![id.as_str()], + "the id filter admits it: {rows:?}" + ); +} + +/// A redacted column cannot be matched or scored: both are computed over +/// the stored text. +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn text_reads_of_a_redacted_column_are_refused() { + let srv = TestServer::start().await; + let user = "fts_redact_reader"; + create(&srv, "fts_redact").await; + insert(&srv, "fts_redact", "r1", "ssn: '123 456', bio: 'gardener'").await; + create_reader(&srv, user).await; + srv.exec( + "CREATE REDACTION POLICY fts_mask_ssn ON fts_redact FOR ROLE readwrite (ssn MASK '***')", + ) + .await + .unwrap(); + + for sql in [ + "SELECT id FROM fts_redact WHERE text_match(ssn, '123')", + "SELECT id, bm25_score(ssn, '123') FROM fts_redact", + "SELECT id FROM fts_redact WHERE text_match(*, '123')", + ] { + let result = first_cells_as(&srv, user, sql).await; + assert!(result.is_err(), "{sql} must be refused: {result:?}"); + } + let allowed = first_cells_as( + &srv, + user, + "SELECT id FROM fts_redact WHERE text_match(bio, 'gardener')", + ) + .await + .unwrap(); + assert_eq!(allowed, vec!["r1".to_string()]); +} + +/// A score scan over a bitemporal collection reads each row's current +/// version. +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn a_bitemporal_score_scan_reads_current_versions() { + let srv = TestServer::start().await; + srv.exec( + "CREATE COLLECTION fts_bitemp (id STRING PRIMARY KEY, body STRING) \ + WITH (engine='document_schemaless', bitemporal=true)", + ) + .await + .unwrap(); + srv.exec("INSERT INTO fts_bitemp (id, body) VALUES ('v1', 'amber stone')") + .await + .unwrap(); + srv.exec("INSERT INTO fts_bitemp (id, body) VALUES ('v2', 'granite stone')") + .await + .unwrap(); + srv.exec("UPDATE fts_bitemp SET body = 'amber resin' WHERE id = 'v2'") + .await + .unwrap(); + + let rows = srv + .query_rows( + "SELECT id, bm25_score(body, 'amber') AS s FROM fts_bitemp \ + ORDER BY s DESC NULLS LAST, id", + ) + .await + .unwrap(); + assert_eq!(rows.len(), 2, "one row per current version: {rows:?}"); + for row in &rows { + let score: f64 = row[1].parse().expect("both rows hold amber"); + assert!(score > 0.0, "both current versions match: {rows:?}"); + } +} + +/// TRUNCATE empties the collection's inverted index and keeps it usable. +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn truncate_empties_the_text_index() { + let srv = TestServer::start().await; + create(&srv, "fts_truncate").await; + for i in 0..5 { + insert( + &srv, + "fts_truncate", + &format!("t{i}"), + "body: 'basalt column'", + ) + .await; + } + srv.exec("TRUNCATE fts_truncate").await.unwrap(); + + let rows = srv + .query_rows("SELECT id FROM fts_truncate WHERE text_match(body, 'basalt')") + .await + .unwrap(); + assert!(rows.is_empty(), "no text survives TRUNCATE: {rows:?}"); + + insert(&srv, "fts_truncate", "n1", "body: 'basalt again'").await; + let rows = srv + .query_rows("SELECT id FROM fts_truncate WHERE text_match(body, 'basalt')") + .await + .unwrap(); + assert_eq!( + ids(&rows), + vec!["n1"], + "new text indexes after TRUNCATE: {rows:?}" + ); +} diff --git a/nodedb/tests/wire/cases/sql_fts_strict.rs b/nodedb/tests/wire/cases/sql_fts_strict.rs index 438f6502c..cabef1b19 100644 --- a/nodedb/tests/wire/cases/sql_fts_strict.rs +++ b/nodedb/tests/wire/cases/sql_fts_strict.rs @@ -304,3 +304,74 @@ async fn create_fulltext_index_rejects_unrecognized_trailing_tokens() { ) .await; } + +// ── Column scoping on a strict schema ────────────────────────────────────── + +const TYPED_DDL: &str = "CREATE COLLECTION docs_typed TYPE DOCUMENT STRICT (\ + id STRING PRIMARY KEY,\ + title STRING,\ + body STRING,\ + views INT\ + )"; + +async fn seed_typed(server: &TestServer) { + server.exec(TYPED_DDL).await.unwrap(); + server + .exec( + "INSERT INTO docs_typed (id, title, body, views) VALUES \ + ('t1', 'consensus primer', 'nothing else', 1), \ + ('t2', 'gardening', 'consensus in the body', 2)", + ) + .await + .unwrap(); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn an_undeclared_strict_column_is_an_undefined_column() { + let server = TestServer::start().await; + seed_typed(&server).await; + let err = server + .query_text("SELECT id FROM docs_typed WHERE text_match(summary, 'consensus')") + .await + .expect_err("summary is not a column of docs_typed"); + assert!(err.contains("42703"), "expected 42703: {err}"); + assert!( + err.contains("docs_typed"), + "the message names the collection: {err}" + ); + assert!( + err.contains("summary"), + "the message names the column: {err}" + ); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn an_int_strict_column_is_a_datatype_mismatch() { + let server = TestServer::start().await; + seed_typed(&server).await; + let err = server + .query_text("SELECT id FROM docs_typed WHERE text_match(views, 'consensus')") + .await + .expect_err("views is not a text column"); + assert!(err.contains("42804"), "expected 42804: {err}"); + assert!(err.contains("views"), "the message names the column: {err}"); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn a_text_strict_column_scopes_the_search() { + let server = TestServer::start().await; + seed_typed(&server).await; + let rows = server + .query_rows("SELECT id FROM docs_typed WHERE text_match(title, 'consensus') ORDER BY id") + .await + .expect("a declared text column must be searchable"); + let ids: Vec<&str> = rows.iter().map(|r| r[0].as_str()).collect(); + assert_eq!(ids, vec!["t1"], "only the title match returns: {rows:?}"); + + let rows = server + .query_rows("SELECT id FROM docs_typed WHERE text_match(body, 'consensus') ORDER BY id") + .await + .expect("a declared text column must be searchable"); + let ids: Vec<&str> = rows.iter().map(|r| r[0].as_str()).collect(); + assert_eq!(ids, vec!["t2"], "only the body match returns: {rows:?}"); +} diff --git a/nodedb/tests/wire/cases/sql_hybrid_search.rs b/nodedb/tests/wire/cases/sql_hybrid_search.rs index 0d0f70431..a0d5b68a0 100644 --- a/nodedb/tests/wire/cases/sql_hybrid_search.rs +++ b/nodedb/tests/wire/cases/sql_hybrid_search.rs @@ -316,3 +316,144 @@ async fn invalid_fts_query_is_a_syntax_error() { "an invalid FTS query must be syntax_error, not XX000; got: {err}" ); } + +/// `t` holds the term in its title, `b` only in its body. The query vector +/// sits nearest `b`, so the vector leg ranks `b` first and only the text leg +/// can lift `t` above it. +async fn create_scoped_hybrid_collection(server: &TestServer, name: &str) { + server + .exec(&format!("CREATE COLLECTION {name}")) + .await + .unwrap(); + server + .exec(&format!( + "CREATE VECTOR INDEX idx_{name}_emb ON {name} (embedding) METRIC cosine DIM 4" + )) + .await + .unwrap(); + for (id, tenant, title, body, emb) in [ + ( + "t", + "t1", + "consensus primer", + "plain words", + "1.0, 0.0, 0.0, 0.0", + ), + ( + "b", + "t1", + "gardening", + "consensus inside", + "0.0, 1.0, 0.0, 0.0", + ), + ( + "o", + "t2", + "consensus other", + "consensus again", + "0.0, 1.0, 0.1, 0.0", + ), + ] { + server + .exec(&format!( + "INSERT INTO {name} (id, tenant_id, title, body, embedding) \ + VALUES ('{id}', '{tenant}', '{title}', '{body}', ARRAY[{emb}])" + )) + .await + .unwrap(); + } +} + +/// The text leg of `rrf_score(..., bm25_score(title, q))` reads the title +/// index: a body-only match contributes nothing to it. +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn hybrid_text_leg_reads_only_its_column() { + let server = TestServer::start().await; + create_scoped_hybrid_collection(&server, "hs_scoped").await; + + let rows = server + .query_rows( + "SELECT id, \ + rrf_score(\ + vector_distance(embedding, ARRAY[0.0, 1.0, 0.0, 0.0]), \ + bm25_score(title, 'consensus')\ + ) AS score \ + FROM hs_scoped WHERE tenant_id = 't1' \ + ORDER BY score DESC LIMIT 5", + ) + .await + .expect("a field-scoped hybrid search must succeed"); + assert_eq!(rows.len(), 2, "both t1 rows fuse: {rows:?}"); + assert_eq!( + rows[0][0], "t", + "the title match must outrank the body-only row: {rows:?}" + ); +} + +/// A WHERE predicate restricts both legs: a row it excludes never fuses. +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn hybrid_where_filter_restricts_the_fused_rows() { + let server = TestServer::start().await; + create_scoped_hybrid_collection(&server, "hs_filtered").await; + + let rows = server + .query_rows( + "SELECT id, \ + rrf_score(\ + vector_distance(embedding, ARRAY[0.0, 1.0, 0.1, 0.0]), \ + bm25_score(body, 'consensus')\ + ) AS score \ + FROM hs_filtered WHERE tenant_id = 't2' LIMIT 5", + ) + .await + .expect("a filtered hybrid search must succeed"); + let ids: Vec<&str> = rows.iter().map(|r| r[0].as_str()).collect(); + assert_eq!(ids, vec!["o"], "only the t2 row may fuse: {rows:?}"); +} + +/// The vector leg searches the column `vector_distance` names. +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn hybrid_vector_leg_reads_the_named_column() { + let server = TestServer::start().await; + server.exec("CREATE COLLECTION hs_two_vec").await.unwrap(); + server + .exec("CREATE VECTOR INDEX idx_hs_two_vec_a ON hs_two_vec (emb_a) METRIC cosine DIM 4") + .await + .unwrap(); + server + .exec("CREATE VECTOR INDEX idx_hs_two_vec_b ON hs_two_vec (emb_b) METRIC cosine DIM 4") + .await + .unwrap(); + // `x` is nearest the query in `emb_b`, `y` in `emb_a`. + server + .exec( + "INSERT INTO hs_two_vec (id, body, emb_a, emb_b) VALUES \ + ('x', 'plain', ARRAY[0.0, 1.0, 0.0, 0.0], ARRAY[1.0, 0.0, 0.0, 0.0])", + ) + .await + .unwrap(); + server + .exec( + "INSERT INTO hs_two_vec (id, body, emb_a, emb_b) VALUES \ + ('y', 'plain', ARRAY[1.0, 0.0, 0.0, 0.0], ARRAY[0.0, 1.0, 0.0, 0.0])", + ) + .await + .unwrap(); + + let rows = server + .query_rows( + "SELECT id, \ + rrf_score(\ + vector_distance(emb_b, ARRAY[1.0, 0.0, 0.0, 0.0]), \ + bm25_score(body, 'absent')\ + ) AS score \ + FROM hs_two_vec ORDER BY score DESC LIMIT 5", + ) + .await + .expect("a hybrid search over a named vector column must succeed"); + assert!( + !rows.is_empty(), + "the emb_b index must answer the vector leg" + ); + assert_eq!(rows[0][0], "x", "emb_b ranks x first: {rows:?}"); +} diff --git a/nodedb/tests/wire/cases/sql_search_subquery_composition.rs b/nodedb/tests/wire/cases/sql_search_subquery_composition.rs index 584f79d50..9f0a89f50 100644 --- a/nodedb/tests/wire/cases/sql_search_subquery_composition.rs +++ b/nodedb/tests/wire/cases/sql_search_subquery_composition.rs @@ -546,9 +546,9 @@ async fn a_sparse_where_comparison_is_refused_at_plan_time() { ); } -/// A one-argument `vector_distance` does not route to a search -/// (`order_by/triggers.rs` returns `Ok(None)` below two arguments), so no -/// cell is declared and `s.distance` must refuse `42703`. +/// A one-argument `vector_distance` whose argument is a column names no query +/// vector, so no signature matches the call and the planner refuses it with +/// `42883` before any search routes. #[tokio::test(flavor = "multi_thread", worker_threads = 4)] async fn a_single_argument_order_by_refuses_the_distance_cell() { let server = TestServer::start().await; diff --git a/nodedb/tests/wire/cases/sql_three_source_rrf_scoped.rs b/nodedb/tests/wire/cases/sql_three_source_rrf_scoped.rs new file mode 100644 index 000000000..452703fe5 --- /dev/null +++ b/nodedb/tests/wire/cases/sql_three_source_rrf_scoped.rs @@ -0,0 +1,118 @@ +// SPDX-License-Identifier: BUSL-1.1 + +//! Three-source RRF fusion over named columns and residual filters: the text +//! leg reads the column `bm25_score` names, the vector leg the column +//! `vector_distance` names, and a WHERE predicate restricts every leg, +//! the graph walk included. + +use crate::harness::TestServer; + +/// Layout over a named vector column `emb`: +/// s1 (vector = query, title "plain", body "alpha") --hop--> s2 +/// s2 (vector far, title "alpha", body "plain", tag "drop") +/// s3 (vector far, title "plain", body "plain") +async fn create_scoped_triple_collection(server: &TestServer, name: &str) { + server + .exec(&format!("CREATE COLLECTION {name}")) + .await + .unwrap(); + server + .exec(&format!( + "CREATE VECTOR INDEX idx_{name}_emb ON {name} (emb) METRIC cosine DIM 3" + )) + .await + .unwrap(); + for (id, tag, title, body, emb) in [ + ("s1", "keep", "plain", "alpha", "1.0, 0.0, 0.0"), + ("s2", "drop", "alpha", "plain", "-1.0, 0.0, 0.0"), + ("s3", "keep", "plain", "plain", "0.0, -1.0, 0.0"), + ] { + server + .exec(&format!( + "INSERT INTO {name} (id, tag, title, body, emb) \ + VALUES ('{id}', '{tag}', '{title}', '{body}', ARRAY[{emb}])" + )) + .await + .unwrap(); + } + server + .exec(&format!( + "GRAPH INSERT EDGE IN '{name}' FROM 's1' TO 's2' TYPE 'hop'" + )) + .await + .unwrap(); +} + +/// The text leg reads the title index only: `s2` holds "alpha" in its title, +/// `s1` only in its body. With a dominant text weight `s2` ranks first. +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn rrf_score_triple_text_leg_reads_only_its_column() { + let server = TestServer::start().await; + create_scoped_triple_collection(&server, "t3_scoped").await; + + let rows = server + .query_rows( + "SELECT id, \ + rrf_score(\ + vector_distance(emb, ARRAY[0.0, 0.0, 1.0]), \ + bm25_score(title, 'alpha'), \ + graph_score(id, 's3', depth => 1, label => 'hop'), \ + 10000.0, 1.0, 10000.0 \ + ) AS score \ + FROM t3_scoped ORDER BY score DESC LIMIT 10", + ) + .await + .expect("a field-scoped three-source search must succeed"); + assert!(!rows.is_empty(), "the search must fuse rows"); + assert_eq!(rows[0][0], "s2", "the title match leads: {rows:?}"); +} + +/// A WHERE predicate restricts every leg, the graph walk included: `s2` is +/// reachable from `s1` but its tag fails the filter. +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn rrf_score_triple_where_filter_restricts_every_leg() { + let server = TestServer::start().await; + create_scoped_triple_collection(&server, "t3_filtered").await; + + let rows = server + .query_rows( + "SELECT id, \ + rrf_score(\ + vector_distance(emb, ARRAY[1.0, 0.0, 0.0]), \ + bm25_score(title, 'alpha'), \ + graph_score(id, 's1', depth => 1, label => 'hop') \ + ) AS score \ + FROM t3_filtered WHERE tag = 'keep' LIMIT 10", + ) + .await + .expect("a filtered three-source search must succeed"); + let ids: Vec<&str> = rows.iter().map(|r| r[0].as_str()).collect(); + assert!(!ids.is_empty(), "the kept rows fuse: {rows:?}"); + assert!( + !ids.contains(&"s2"), + "a row the filter excludes never fuses: {rows:?}" + ); +} + +/// The vector leg searches the column `vector_distance` names. +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn rrf_score_triple_vector_leg_reads_the_named_column() { + let server = TestServer::start().await; + create_scoped_triple_collection(&server, "t3_named_vec").await; + + let rows = server + .query_rows( + "SELECT id, \ + rrf_score(\ + vector_distance(emb, ARRAY[1.0, 0.0, 0.0]), \ + bm25_score(title, 'absent'), \ + graph_score(id, 's3', depth => 1, label => 'hop'), \ + 1.0, 10000.0, 10000.0 \ + ) AS score \ + FROM t3_named_vec ORDER BY score DESC LIMIT 10", + ) + .await + .expect("a three-source search over a named vector column must succeed"); + assert!(!rows.is_empty(), "the emb index must answer the vector leg"); + assert_eq!(rows[0][0], "s1", "emb ranks s1 first: {rows:?}"); +} diff --git a/nodedb/tests/wire/cases/sql_transactions_fts_analyzer_overlay.rs b/nodedb/tests/wire/cases/sql_transactions_fts_analyzer_overlay.rs index f7b007862..b49eb4e2f 100644 --- a/nodedb/tests/wire/cases/sql_transactions_fts_analyzer_overlay.rs +++ b/nodedb/tests/wire/cases/sql_transactions_fts_analyzer_overlay.rs @@ -191,3 +191,57 @@ async fn staged_insert_rollback_removes_match() { "ROLLBACK must leave no durable index trace of the staged insert: {after_rollback:?}" ); } + +/// The `'indonesian'` analyzer does not stem, so the index holds `running` and +/// `dogs` as written. A phrase is analyzed once with the collection's +/// analyzer, so it matches the indexed text and the staged text alike. An +/// English stemmer would turn the phrase into `run dog`, which no indexed +/// token holds. +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn a_phrase_matches_under_a_no_stem_analyzer() { + let server = TestServer::start().await; + let coll = "fts_an_phrase"; + server + .exec(&format!( + "CREATE COLLECTION {coll} WITH (engine='document_schemaless')" + )) + .await + .unwrap(); + server + .exec(&format!( + "CREATE SEARCH INDEX idx_{coll}_fts ON {coll} FIELDS body ANALYZER 'indonesian'" + )) + .await + .unwrap(); + for (id, body) in [("p1", "running dogs bark"), ("p2", "dogs running fast")] { + server + .exec(&format!( + "INSERT INTO {coll} (id, body) VALUES ('{id}', '{body}')" + )) + .await + .unwrap(); + } + let phrase = "\"running dogs\""; + + let committed = matched_ids(&server, coll, phrase).await; + assert_eq!( + committed, + vec!["p1".to_string()], + "the phrase matches its words in order, unstemmed: {committed:?}" + ); + + server.exec("BEGIN").await.unwrap(); + server + .exec(&format!( + "INSERT INTO {coll} (id, body) VALUES ('p3', 'two running dogs')" + )) + .await + .unwrap(); + let in_txn = matched_ids(&server, coll, phrase).await; + assert_eq!( + in_txn, + vec!["p1".to_string(), "p3".to_string()], + "the staged row matches the same analyzed phrase: {in_txn:?}" + ); + server.client.simple_query("ROLLBACK").await.unwrap(); +} diff --git a/nodedb/tests/wire/cases/sql_transactions_fts_overlay.rs b/nodedb/tests/wire/cases/sql_transactions_fts_overlay.rs index ed57e7554..5272248af 100644 --- a/nodedb/tests/wire/cases/sql_transactions_fts_overlay.rs +++ b/nodedb/tests/wire/cases/sql_transactions_fts_overlay.rs @@ -264,3 +264,102 @@ async fn strict_delete_hides_in_txn_then_rollback_restores() { async fn strict_update_changes_match_in_txn() { update_changes_match_in_txn("document_strict", "fts_ov_st_upd").await; } + +// ── Field scoping and residual filters inside a transaction ──────────────── + +async fn ids_of(server: &TestServer, sql: &str) -> Vec { + server + .query_rows(sql) + .await + .unwrap() + .into_iter() + .map(|r| r[0].clone()) + .collect() +} + +/// A staged insert is visible to a search scoped to the field that holds the +/// term, and invisible to a search scoped to another field. +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn staged_insert_is_visible_to_a_field_scoped_search() { + let server = TestServer::start().await; + setup(&server, "fts_ov_scoped", "document_schemaless").await; + + server.exec("BEGIN").await.unwrap(); + server + .exec( + "INSERT INTO fts_ov_scoped (id, title, body) \ + VALUES ('new1', 'elephant tales', 'nothing relevant')", + ) + .await + .unwrap(); + let by_title = ids_of( + &server, + "SELECT id FROM fts_ov_scoped WHERE text_match(title, 'elephant') ORDER BY id", + ) + .await; + let by_body = ids_of( + &server, + "SELECT id FROM fts_ov_scoped WHERE text_match(body, 'elephant') ORDER BY id", + ) + .await; + server.client.simple_query("ROLLBACK").await.unwrap(); + + assert_eq!( + by_title, + vec!["new1".to_string()], + "the title scope sees it" + ); + assert!(by_body.is_empty(), "the body scope does not: {by_body:?}"); +} + +/// A staged row that fails the WHERE filter never enters the result. +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn staged_row_failing_the_filter_is_excluded() { + let server = TestServer::start().await; + setup(&server, "fts_ov_filtered", "document_schemaless").await; + + server.exec("BEGIN").await.unwrap(); + for (id, tag) in [("keep1", "a"), ("drop1", "b")] { + server + .exec(&format!( + "INSERT INTO fts_ov_filtered (id, tag, body) \ + VALUES ('{id}', '{tag}', 'an elephant never forgets')" + )) + .await + .unwrap(); + } + let ids = ids_of( + &server, + "SELECT id FROM fts_ov_filtered \ + WHERE text_match(body, 'elephant') AND tag = 'a' ORDER BY id", + ) + .await; + server.client.simple_query("ROLLBACK").await.unwrap(); + + assert_eq!(ids, vec!["keep1".to_string()], "only the tag a row returns"); +} + +/// A field that first holds text inside the transaction is searchable there, +/// not refused as a field no document holds. +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn first_use_of_a_field_inside_a_transaction_is_not_an_error() { + let server = TestServer::start().await; + setup(&server, "fts_ov_new_field", "document_schemaless").await; + + server.exec("BEGIN").await.unwrap(); + server + .exec( + "INSERT INTO fts_ov_new_field (id, note) \ + VALUES ('n1', 'brand new elephant note')", + ) + .await + .unwrap(); + let result = server + .query_rows("SELECT id FROM fts_ov_new_field WHERE text_match(note, 'elephant')") + .await; + server.client.simple_query("ROLLBACK").await.unwrap(); + + let rows = result.expect("a field staged in this transaction must be searchable"); + let ids: Vec<&str> = rows.iter().map(|r| r[0].as_str()).collect(); + assert_eq!(ids, vec!["n1"], "the staged note matches"); +} diff --git a/nodedb/tests/wire/cases/sql_transactions_hybrid_overlay.rs b/nodedb/tests/wire/cases/sql_transactions_hybrid_overlay.rs index 9de31a0be..00794c273 100644 --- a/nodedb/tests/wire/cases/sql_transactions_hybrid_overlay.rs +++ b/nodedb/tests/wire/cases/sql_transactions_hybrid_overlay.rs @@ -300,3 +300,91 @@ async fn autocommit_hybrid_unchanged() { "autocommit hybrid must return a fused row for the committed doc matching its term: {scores:?}" ); } + +const RLS_PASSWORD: &str = "hyb-ov-rls-secret-1"; + +/// The number of rows `sql` returns on `client`. +async fn row_count(client: &tokio_postgres::Client, sql: &str) -> usize { + client + .simple_query(sql) + .await + .unwrap_or_else(|e| panic!("{sql}: {e}")) + .iter() + .filter(|m| matches!(m, tokio_postgres::SimpleQueryMessage::Row(_))) + .count() +} + +/// Under a read policy, a row the reader inserts in its own transaction is +/// judged on its staged body. The fused result holds it before COMMIT, and a +/// row the policy hides never appears. +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn staged_insert_is_fused_under_a_read_policy() { + let server = TestServer::start().await; + let coll = "hyb_ov_rls"; + let user = "hyb_ov_rls_reader"; + create_hybrid_collection(&server, coll).await; + server + .exec(&format!("CREATE USER {user} PASSWORD '{RLS_PASSWORD}'")) + .await + .unwrap(); + server + .exec(&format!("GRANT ROLE readwrite TO {user}")) + .await + .unwrap(); + for (id, owner, content, emb) in [ + ("a", user, "consensus algorithm", "0.1, 0.2, 0.3, 0.4"), + ("b", "other", "elephant herd", "0.15, 0.25, 0.35, 0.45"), + ] { + server + .exec(&format!( + "INSERT INTO {coll} (id, owner, content, embedding) \ + VALUES ('{id}', '{owner}', '{content}', ARRAY[{emb}])" + )) + .await + .unwrap(); + } + server + .exec(&format!( + "CREATE RLS POLICY hyb_ov_rls_owner ON {coll} FOR READ \ + USING (owner = $auth.username)" + )) + .await + .unwrap(); + + let (client, handle) = server + .connect_as(user, RLS_PASSWORD) + .await + .unwrap_or_else(|e| panic!("connect as {user}: {e}")); + let hybrid = format!( + "SELECT id, rrf_score(\ + vector_distance(embedding, ARRAY[0.15, 0.25, 0.35, 0.45]), \ + bm25_score(content, 'elephant')\ + ) AS score FROM {coll} LIMIT 10" + ); + + // Only 'a' is admitted. The policy hides 'b', the only text match. + assert_eq!(row_count(&client, &hybrid).await, 1, "only 'a' is admitted"); + + client.simple_query("BEGIN").await.unwrap(); + client + .simple_query(&format!( + "INSERT INTO {coll} (id, owner, content, embedding) \ + VALUES ('c', '{user}', 'elephant calf', ARRAY[0.15, 0.25, 0.35, 0.45])" + )) + .await + .unwrap(); + assert_eq!( + row_count(&client, &hybrid).await, + 2, + "the staged row the policy admits is fused before COMMIT" + ); + client.simple_query("ROLLBACK").await.unwrap(); + assert_eq!( + row_count(&client, &hybrid).await, + 1, + "the rolled-back row is gone" + ); + + drop(client); + handle.abort(); +} From 8d91af6ecba5ebae6178c550bcb5e5fa0ba8dcc2 Mon Sep 17 00:00:00 2001 From: Farhan Syah Date: Thu, 8 Oct 2026 16:51:47 +0800 Subject: [PATCH 16/24] feat(vector): restrict search candidates before the top-k cut A vector search can rank among a set of primary keys, from a key IN list in SQL or allowed_ids on the native protocol, lowered to the candidate bitmap. A native metadata filter plans as the WHERE of the collection and fills the residual-filter slot, and the candidate window widens until top_k rows pass. A read policy joins the statement's WHERE predicates instead of replacing them. A search that names no field reads the collection's only vector index. --- .../src/planner/collection_predicate.rs | 170 +++++++++++++ nodedb-sql/src/planner/mod.rs | 1 + .../src/planner/select/entry/pk_prefilter.rs | 234 ++++++++++++++++++ nodedb-sql/src/types/plan/search_cells.rs | 1 + .../sql_plan_convert/expr/inline_cte.rs | 1 + .../native/dispatch/plan_builder/mod.rs | 2 + .../native/dispatch/plan_builder/vector.rs | 27 +- .../dispatch/plan_builder/vector_filter.rs | 218 ++++++++++++++++ .../src/data/executor/core_loop/response.rs | 75 +++++- .../dispatch/bitmap/hashjoin_inline.rs | 182 +++++++++++--- .../data/executor/handlers/vector_search.rs | 12 +- .../executor/handlers/vector_search_exec.rs | 187 ++++++-------- .../executor/handlers/vector_search_window.rs | 113 +++++++++ .../wire/cases/sql_where_vector_search.rs | 79 ++++++ 14 files changed, 1148 insertions(+), 154 deletions(-) create mode 100644 nodedb-sql/src/planner/collection_predicate.rs create mode 100644 nodedb-sql/src/planner/select/entry/pk_prefilter.rs create mode 100644 nodedb/src/control/server/native/dispatch/plan_builder/vector_filter.rs create mode 100644 nodedb/src/data/executor/handlers/vector_search_window.rs diff --git a/nodedb-sql/src/planner/collection_predicate.rs b/nodedb-sql/src/planner/collection_predicate.rs new file mode 100644 index 000000000..abb005da0 --- /dev/null +++ b/nodedb-sql/src/planner/collection_predicate.rs @@ -0,0 +1,170 @@ +// SPDX-License-Identifier: Apache-2.0 + +//! Plan a predicate against one collection the way a single-collection +//! `SELECT ... WHERE` plans its `WHERE`. +//! +//! An entry point that receives a predicate as a typed tree instead of SQL +//! text (the native protocol's metadata filter) plans it here. The +//! predicate then gets the same column check and the same literal coercion +//! as the `WHERE` a SQL client writes for the same filter, so both entry +//! points match the same rows. + +use nodedb_types::DatabaseId; + +use crate::error::{Result, SqlError}; +use crate::planner::predicate_coerce::coerce_predicate_literals; +use crate::resolver::columns::{ResolvedTable, TableScope}; +use crate::types::{Filter, FilterExpr, SqlCatalog, SqlExpr}; + +/// Plan `predicate` as the `WHERE` of `SELECT * FROM `. +/// +/// A collection the catalog does not hold is `UnknownTable`. A column the +/// collection does not have is `UnknownColumn`. A literal compared against a +/// declared `TIMESTAMP` / `TIMESTAMPTZ` column becomes a typed instant. +pub fn plan_collection_predicate( + catalog: &dyn SqlCatalog, + collection: &str, + mut predicate: SqlExpr, +) -> Result> { + let info = catalog + .resolve_relation(DatabaseId::DEFAULT, collection)? + .ok_or_else(|| SqlError::UnknownTable { + name: collection.to_string(), + })?; + let scope = TableScope::single(ResolvedTable { + name: collection.to_string(), + alias: None, + info, + })?; + check_columns(&predicate, &scope)?; + coerce_predicate_literals(&mut predicate, &scope)?; + Ok(vec![Filter { + expr: FilterExpr::Expr(predicate), + }]) +} + +/// Reject a column reference that names no column of the scope. +fn check_columns(expr: &SqlExpr, scope: &TableScope) -> Result<()> { + match expr { + SqlExpr::Column { table, name } => scope.check_name(table.as_deref(), name), + SqlExpr::Literal(_) | SqlExpr::Wildcard | SqlExpr::Subquery(_) => Ok(()), + SqlExpr::BinaryOp { left, right, .. } => { + check_columns(left, scope)?; + check_columns(right, scope) + } + SqlExpr::UnaryOp { expr, .. } + | SqlExpr::Cast { expr, .. } + | SqlExpr::IsNull { expr, .. } => check_columns(expr, scope), + SqlExpr::Function { args, .. } | SqlExpr::ArrayLiteral(args) => { + args.iter().try_for_each(|arg| check_columns(arg, scope)) + } + SqlExpr::Case { + operand, + when_then, + else_expr, + } => { + if let Some(operand) = operand { + check_columns(operand, scope)?; + } + for (when, then) in when_then { + check_columns(when, scope)?; + check_columns(then, scope)?; + } + match else_expr { + Some(else_expr) => check_columns(else_expr, scope), + None => Ok(()), + } + } + SqlExpr::InList { expr, list, .. } => { + check_columns(expr, scope)?; + list.iter().try_for_each(|item| check_columns(item, scope)) + } + SqlExpr::Between { + expr, low, high, .. + } => { + check_columns(expr, scope)?; + check_columns(low, scope)?; + check_columns(high, scope) + } + SqlExpr::Like { expr, pattern, .. } => { + check_columns(expr, scope)?; + check_columns(pattern, scope) + } + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::catalog::SqlCatalogError; + use crate::types::{BinaryOp, CollectionInfo, ColumnInfo, EngineType, SqlDataType, SqlValue}; + + struct OneCollection; + + impl SqlCatalog for OneCollection { + fn get_collection( + &self, + _database_id: DatabaseId, + name: &str, + ) -> std::result::Result, SqlCatalogError> { + if name != "docs" { + return Ok(None); + } + Ok(Some(CollectionInfo { + name: "docs".into(), + engine: EngineType::DocumentStrict, + columns: vec![ColumnInfo { + name: "category".into(), + data_type: SqlDataType::String, + nullable: true, + is_primary_key: false, + default: None, + raw_type: None, + int_width: None, + float_width: None, + }], + primary_key: None, + has_auto_tier: false, + indexes: Vec::new(), + bitemporal: false, + primary: nodedb_types::PrimaryEngine::Document, + vector_primary: None, + partition_strategy: nodedb_types::PartitionStrategy::CollectionHomed, + open_schema: CollectionInfo::open_schema_for(EngineType::DocumentStrict), + })) + } + } + + fn eq(field: &str, value: &str) -> SqlExpr { + SqlExpr::BinaryOp { + left: Box::new(SqlExpr::Column { + table: None, + name: field.into(), + }), + op: BinaryOp::Eq, + right: Box::new(SqlExpr::Literal(SqlValue::String(value.into()))), + } + } + + #[test] + fn declared_column_plans_as_where_expression() { + let filters = plan_collection_predicate(&OneCollection, "docs", eq("category", "ai")) + .expect("declared column plans"); + assert_eq!(filters.len(), 1); + assert!(matches!(filters[0].expr, FilterExpr::Expr(_))); + } + + #[test] + fn undeclared_column_is_unknown_column() { + let err = plan_collection_predicate(&OneCollection, "docs", eq("ghost", "x")) + .expect_err("a strict collection rejects an undeclared column"); + assert!(matches!(err, SqlError::UnknownColumn { .. }), "{err:?}"); + } + + #[test] + fn unknown_collection_is_unknown_table() { + let err = plan_collection_predicate(&OneCollection, "nope", eq("category", "ai")) + .expect_err("no such collection"); + assert!(matches!(err, SqlError::UnknownTable { .. }), "{err:?}"); + } +} diff --git a/nodedb-sql/src/planner/mod.rs b/nodedb-sql/src/planner/mod.rs index 3a1a72da3..e97486d7d 100644 --- a/nodedb-sql/src/planner/mod.rs +++ b/nodedb-sql/src/planner/mod.rs @@ -14,6 +14,7 @@ pub mod catalog_expr_fold; pub mod catalog_fold; pub mod catalog_plan_shapes; pub mod catalog_plan_validate; +pub mod collection_predicate; pub mod const_fold; pub mod cp_projection; pub mod cte; diff --git a/nodedb-sql/src/planner/select/entry/pk_prefilter.rs b/nodedb-sql/src/planner/select/entry/pk_prefilter.rs new file mode 100644 index 000000000..2c3000c2a --- /dev/null +++ b/nodedb-sql/src/planner/select/entry/pk_prefilter.rs @@ -0,0 +1,234 @@ +// SPDX-License-Identifier: Apache-2.0 + +//! Primary-key prefilter of a vector search. +//! +//! A top-level `WHERE pk = v` or `WHERE pk IN (v, ...)` conjunct of a +//! `VectorSearch` names the only rows the search may rank. The pass moves it +//! out of the residual filters into `pk_prefilter`. The executor lowers the +//! keys to a candidate bitmap the index search honors, so the top-k is drawn +//! from those rows even when the nearest vectors lie outside them. Left as a +//! residual filter, the conjunct would run after an over-fetched top-k cut +//! and drop rows the cut never reached. + +use nodedb_types::DatabaseId; + +use crate::error::Result; +use crate::types::*; + +/// Move the primary-key conjuncts of a `VectorSearch` plan's filters into its +/// `pk_prefilter`. Other plans are untouched. +pub(super) fn apply_vector_pk_prefilter( + plan: &mut SqlPlan, + catalog: &dyn SqlCatalog, +) -> Result<()> { + let SqlPlan::VectorSearch { + collection, + filters, + pk_prefilter, + .. + } = plan + else { + return Ok(()); + }; + let Some(info) = catalog.get_collection(DatabaseId::DEFAULT, collection)? else { + return Ok(()); + }; + let Some(pk) = info.primary_key.as_deref() else { + return Ok(()); + }; + let mut keys: Option> = pk_prefilter.take(); + let mut kept: Vec = Vec::with_capacity(filters.len()); + for filter in filters.drain(..) { + match filter.expr { + FilterExpr::Comparison { + ref field, + op: CompareOp::Eq, + ref value, + } if field == pk => restrict(&mut keys, vec![value.clone()]), + FilterExpr::InList { + ref field, + ref values, + } if field == pk => restrict(&mut keys, values.clone()), + FilterExpr::Expr(expr) => { + let mut residual: Vec = Vec::new(); + for conjunct in split_conjuncts(expr) { + match pk_keys(&conjunct, pk) { + Some(found) => restrict(&mut keys, found), + None => residual.push(conjunct), + } + } + if let Some(expr) = join_conjuncts(residual) { + kept.push(Filter { + expr: FilterExpr::Expr(expr), + }); + } + } + other => kept.push(Filter { expr: other }), + } + } + *filters = kept; + *pk_prefilter = keys; + Ok(()) +} + +/// Intersect the admitted keys with `found`. The first restriction admits +/// exactly `found`. +fn restrict(keys: &mut Option>, found: Vec) { + *keys = Some(match keys.take() { + None => found, + Some(admitted) => admitted.into_iter().filter(|k| found.contains(k)).collect(), + }); +} + +/// The top-level `AND` operands of `expr`, in order. +fn split_conjuncts(expr: SqlExpr) -> Vec { + match expr { + SqlExpr::BinaryOp { + left, + op: BinaryOp::And, + right, + } => { + let mut out = split_conjuncts(*left); + out.extend(split_conjuncts(*right)); + out + } + other => vec![other], + } +} + +/// `AND` of `conjuncts`, or `None` when there are none. +fn join_conjuncts(conjuncts: Vec) -> Option { + conjuncts + .into_iter() + .reduce(|left, right| SqlExpr::BinaryOp { + left: Box::new(left), + op: BinaryOp::And, + right: Box::new(right), + }) +} + +/// The keys of `pk = literal` or `pk IN (literal, ...)`. `None` for any other +/// predicate, including an `IN` list holding a non-literal. +fn pk_keys(expr: &SqlExpr, pk: &str) -> Option> { + let is_pk = |e: &SqlExpr| matches!(e, SqlExpr::Column { name, .. } if name == pk); + match expr { + SqlExpr::BinaryOp { + left, + op: BinaryOp::Eq, + right, + } => match (left.as_ref(), right.as_ref()) { + (column, SqlExpr::Literal(value)) | (SqlExpr::Literal(value), column) + if is_pk(column) => + { + Some(vec![value.clone()]) + } + _ => None, + }, + SqlExpr::InList { + expr, + list, + negated: false, + } if is_pk(expr) => list + .iter() + .map(|item| match item { + SqlExpr::Literal(value) => Some(value.clone()), + _ => None, + }) + .collect(), + _ => None, + } +} + +#[cfg(test)] +mod tests { + use super::super::fixtures::plan_select_sql; + use crate::types::*; + + /// The key prefilter and residual filter count of a vector search plan. + fn vector_parts(sql: &str) -> (Option>, usize) { + match plan_select_sql(sql) { + SqlPlan::VectorSearch { + pk_prefilter, + filters, + .. + } => (pk_prefilter, filters.len()), + other => panic!("expected VectorSearch, got {other:?}"), + } + } + + #[test] + fn an_id_in_list_becomes_the_prefilter() { + let (keys, residual) = vector_parts( + "SELECT id FROM embeddings WHERE id IN ('a', 'b') \ + ORDER BY vector_distance(embedding, [1.0, 0.0]) LIMIT 2", + ); + assert_eq!( + keys, + Some(vec![ + SqlValue::String("a".into()), + SqlValue::String("b".into()) + ]) + ); + assert_eq!(residual, 0); + } + + #[test] + fn other_conjuncts_stay_residual_filters() { + let (keys, residual) = vector_parts( + "SELECT id FROM embeddings WHERE tag = 'x' AND id = 'a' \ + ORDER BY vector_distance(embedding, [1.0, 0.0]) LIMIT 2", + ); + assert_eq!(keys, Some(vec![SqlValue::String("a".into())])); + assert_eq!(residual, 1); + } + + #[test] + fn a_disjunction_with_the_key_is_no_prefilter() { + let (keys, residual) = vector_parts( + "SELECT id FROM embeddings WHERE id = 'a' OR tag = 'x' \ + ORDER BY vector_distance(embedding, [1.0, 0.0]) LIMIT 2", + ); + assert_eq!(keys, None); + assert_eq!(residual, 1); + } + + #[test] + fn two_key_conjuncts_intersect() { + let (keys, _) = vector_parts( + "SELECT id FROM embeddings WHERE id IN ('a', 'b') AND id IN ('b', 'c') \ + ORDER BY vector_distance(embedding, [1.0, 0.0]) LIMIT 2", + ); + assert_eq!(keys, Some(vec![SqlValue::String("b".into())])); + } + + #[test] + fn the_remote_client_form_is_a_prefiltered_search_of_the_default_column() { + let plan = plan_select_sql( + "SELECT * FROM embeddings WHERE id IN ('a') \ + ORDER BY vector_distance(ARRAY[1.0, 0.0]) LIMIT 2", + ); + let SqlPlan::VectorSearch { + field, + pk_prefilter, + top_k, + .. + } = plan + else { + panic!("expected VectorSearch, got {plan:?}"); + }; + assert_eq!( + field, "", + "no declared vector column: the collection-level index" + ); + assert_eq!(pk_prefilter, Some(vec![SqlValue::String("a".into())])); + assert_eq!(top_k, 2); + } + + #[test] + fn no_key_conjunct_is_no_prefilter() { + let (keys, _) = vector_parts( + "SELECT id FROM embeddings ORDER BY vector_distance(embedding, [1.0, 0.0]) LIMIT 2", + ); + assert_eq!(keys, None); + } +} diff --git a/nodedb-sql/src/types/plan/search_cells.rs b/nodedb-sql/src/types/plan/search_cells.rs index 51233cd5e..60880bec1 100644 --- a/nodedb-sql/src/types/plan/search_cells.rs +++ b/nodedb-sql/src/types/plan/search_cells.rs @@ -64,6 +64,7 @@ mod tests { ann_options: VectorAnnOptions::default(), skip_payload_fetch: false, payload_filters: Vec::new(), + pk_prefilter: None, projection: Vec::new(), } .carries_search_cells() diff --git a/nodedb/src/control/planner/sql_plan_convert/expr/inline_cte.rs b/nodedb/src/control/planner/sql_plan_convert/expr/inline_cte.rs index 91cd73102..26f1ebccd 100644 --- a/nodedb/src/control/planner/sql_plan_convert/expr/inline_cte.rs +++ b/nodedb/src/control/planner/sql_plan_convert/expr/inline_cte.rs @@ -364,6 +364,7 @@ mod tests { ann_options: nodedb_sql::types::VectorAnnOptions::default(), skip_payload_fetch: false, payload_filters: Vec::new(), + pk_prefilter: None, projection: Vec::new(), } } diff --git a/nodedb/src/control/server/native/dispatch/plan_builder/mod.rs b/nodedb/src/control/server/native/dispatch/plan_builder/mod.rs index 0d5555db5..27e372420 100644 --- a/nodedb/src/control/server/native/dispatch/plan_builder/mod.rs +++ b/nodedb/src/control/server/native/dispatch/plan_builder/mod.rs @@ -10,6 +10,7 @@ pub(crate) mod crdt; mod dispatch; pub(crate) mod document; pub(crate) mod document_bulk; +mod document_identity; pub(crate) mod graph; mod helpers; pub(crate) mod kv; @@ -19,6 +20,7 @@ pub(crate) mod spatial; pub(crate) mod text; pub(crate) mod timeseries; pub(crate) mod vector; +mod vector_filter; pub(crate) use dispatch::build_plan; pub(super) use helpers::{collection_type, declared_primary_key, parse_direction, require_doc_id}; diff --git a/nodedb/src/control/server/native/dispatch/plan_builder/vector.rs b/nodedb/src/control/server/native/dispatch/plan_builder/vector.rs index f1ccb99be..aa16d7ba7 100644 --- a/nodedb/src/control/server/native/dispatch/plan_builder/vector.rs +++ b/nodedb/src/control/server/native/dispatch/plan_builder/vector.rs @@ -2,9 +2,9 @@ //! Vector engine plan builders. -use nodedb_types::QualifiedCollection; use nodedb_types::protocol::TextFields; use nodedb_types::vector_distance::DistanceMetric; +use nodedb_types::{QualifiedCollection, SurrogateBitmap}; use super::super::DispatchCtx; use crate::bridge::envelope::PhysicalPlan; @@ -24,6 +24,27 @@ pub(crate) async fn build_search( let top_k = fields.top_k.unwrap_or(10) as usize; let ef_search = fields.ef_search.unwrap_or(0) as usize; let field_name = fields.field_name.clone().unwrap_or_default(); + // The allowed ids restrict the candidates before ranking: an id bound to + // no row names no candidate. + let filter_bitmap = match &fields.allowed_ids { + Some(ids) => { + let pks: Vec<&[u8]> = ids.iter().map(|id| id.as_bytes()).collect(); + let surrogates = super::helpers::existing_surrogates(ctx, collection, &pks).await?; + Some( + surrogates + .into_iter() + .flatten() + .collect::(), + ) + } + None => None, + }; + // The metadata filter fills the residual-filter slot a SQL `WHERE` + // fills, so it narrows the candidates before the top-k cut. + let rls_filters = match fields.filters.as_deref() { + Some(bytes) => super::vector_filter::residual_filters(ctx, collection, bytes)?, + None => Vec::new(), + }; Ok(PhysicalPlan::Vector(VectorOp::Search { collection: QualifiedCollection::new(ctx.database_id(), collection), @@ -34,9 +55,9 @@ pub(crate) async fn build_search( // default metric (L2 as the wire sentinel — the Data Plane will apply // the collection-configured metric when these match). metric: DistanceMetric::L2, - filter_bitmap: None, + filter_bitmap, field_name, - rls_filters: Vec::new(), + rls_filters, inline_prefilter_plan: None, // The native protocol carries primitive top_k / ef_search fields // directly; advanced ANN tuning (quantization, oversample, target diff --git a/nodedb/src/control/server/native/dispatch/plan_builder/vector_filter.rs b/nodedb/src/control/server/native/dispatch/plan_builder/vector_filter.rs new file mode 100644 index 000000000..62c30a5e2 --- /dev/null +++ b/nodedb/src/control/server/native/dispatch/plan_builder/vector_filter.rs @@ -0,0 +1,218 @@ +// SPDX-License-Identifier: BUSL-1.1 + +//! The metadata filter of a native `VectorSearch` request. +//! +//! The client sends a `MetadataFilter` as MessagePack in +//! `TextFields::filters`. The filter lowers to the expression tree the SQL +//! planner builds for the `WHERE` a pgwire client renders from the same +//! filter. It is then planned against the collection by the same column +//! check and literal coercion, and lowered to the same `ScanFilter` bytes. +//! The bytes travel in the plan's residual-filter slot, the slot a SQL +//! `WHERE` fills. The filter therefore matches the same rows on both +//! clients and narrows the candidates before the top-k cut. + +use std::sync::Arc; + +use nodedb_sql::types::{BinaryOp, SqlExpr, SqlValue, UnaryOp}; +use nodedb_types::filter::MetadataFilter; +use nodedb_types::value::Value; + +use super::super::DispatchCtx; +use crate::control::planner::catalog_adapter::OriginCatalog; +use crate::control::planner::plan_error_map::map_plan_error; +use crate::control::planner::sql_plan_convert::filter::serialize_filters; + +/// The residual-filter bytes for `filter_bytes`, the client's encoded +/// `MetadataFilter`. +/// +/// Bytes that do not decode to a `MetadataFilter`, and a filter value no SQL +/// literal can carry, are a `BadRequest`. A column the collection does not +/// have fails the same way the SQL `WHERE` fails. +pub(crate) fn residual_filters( + ctx: &DispatchCtx<'_>, + collection: &str, + filter_bytes: &[u8], +) -> crate::Result> { + let filter = decode_metadata_filter(filter_bytes)?; + let predicate = metadata_filter_predicate(&filter)?; + let state = ctx.state; + let catalog = OriginCatalog::new( + Arc::clone(&state.credentials), + state.array_catalog.clone(), + ctx.tenant_id().as_u64(), + ctx.database_id(), + Some(Arc::clone(&state.retention_policy_registry)), + ) + .with_sequence_registry(Arc::clone(&state.sequence_registry)); + let filters = nodedb_sql::planner::collection_predicate::plan_collection_predicate( + &catalog, collection, predicate, + ) + .map_err(|e| map_plan_error(e, ctx.tenant_id()))?; + serialize_filters(&filters) +} + +/// Decode the client's filter bytes as one MessagePack `MetadataFilter`. +fn decode_metadata_filter(bytes: &[u8]) -> crate::Result { + zerompk::from_msgpack(bytes).map_err(|e| crate::Error::BadRequest { + detail: format!("vector_search filter is not a MessagePack MetadataFilter: {e}"), + }) +} + +/// The `WHERE` expression a pgwire client renders for `filter`: +/// `"" `, `IN` / `NOT IN` lists, `AND` / `OR` / `NOT`. +/// An empty `AND` is `TRUE` and an empty `OR` is `FALSE`. +fn metadata_filter_predicate(filter: &MetadataFilter) -> crate::Result { + let compare = |field: &str, op: BinaryOp, value: &Value| -> crate::Result { + Ok(SqlExpr::BinaryOp { + left: Box::new(column(field)), + op, + right: Box::new(SqlExpr::Literal(literal(value)?)), + }) + }; + let in_list = |field: &str, values: &[Value], negated: bool| -> crate::Result { + let list = values + .iter() + .map(|v| literal(v).map(SqlExpr::Literal)) + .collect::>>()?; + Ok(SqlExpr::InList { + expr: Box::new(column(field)), + list, + negated, + }) + }; + match filter { + MetadataFilter::Eq { field, value } => compare(field, BinaryOp::Eq, value), + MetadataFilter::Ne { field, value } => compare(field, BinaryOp::Ne, value), + MetadataFilter::Gt { field, value } => compare(field, BinaryOp::Gt, value), + MetadataFilter::Gte { field, value } => compare(field, BinaryOp::Ge, value), + MetadataFilter::Lt { field, value } => compare(field, BinaryOp::Lt, value), + MetadataFilter::Lte { field, value } => compare(field, BinaryOp::Le, value), + MetadataFilter::In { field, values } => in_list(field, values, false), + MetadataFilter::NotIn { field, values } => in_list(field, values, true), + MetadataFilter::And(parts) => fold(parts, BinaryOp::And, true), + MetadataFilter::Or(parts) => fold(parts, BinaryOp::Or, false), + MetadataFilter::Not(inner) => Ok(SqlExpr::UnaryOp { + op: UnaryOp::Not, + expr: Box::new(metadata_filter_predicate(inner)?), + }), + other => Err(crate::Error::BadRequest { + detail: format!("vector_search filter variant has no WHERE form: {other:?}"), + }), + } +} + +/// `parts` joined by `op`, or the literal `empty` for no part. +fn fold(parts: &[MetadataFilter], op: BinaryOp, empty: bool) -> crate::Result { + let mut exprs = parts.iter().map(metadata_filter_predicate); + let Some(first) = exprs.next() else { + return Ok(SqlExpr::Literal(SqlValue::Bool(empty))); + }; + exprs.try_fold(first?, |acc, next: crate::Result| { + Ok::<_, crate::Error>(SqlExpr::BinaryOp { + left: Box::new(acc), + op, + right: Box::new(next?), + }) + }) +} + +/// A quoted identifier: the field name is kept exactly as sent. +fn column(field: &str) -> SqlExpr { + SqlExpr::Column { + table: None, + name: field.to_string(), + } +} + +/// The SQL literal a pgwire client renders for `value`. Only `NULL`, +/// booleans, integers, floats and strings have a literal form. +fn literal(value: &Value) -> crate::Result { + match value { + Value::Null => Ok(SqlValue::Null), + Value::Bool(b) => Ok(SqlValue::Bool(*b)), + Value::Integer(i) => Ok(SqlValue::Int(*i)), + Value::Float(f) => Ok(SqlValue::Float(*f)), + Value::String(s) => Ok(SqlValue::String(s.clone())), + other => Err(crate::Error::BadRequest { + detail: format!("vector_search filter value has no SQL literal form: {other:?}"), + }), + } +} + +#[cfg(test)] +mod tests { + use super::*; + + fn lowered(filter: &MetadataFilter) -> String { + format!( + "{:?}", + metadata_filter_predicate(filter).expect("filter lowers") + ) + } + + #[test] + fn malformed_bytes_are_bad_request() { + for bytes in [&b""[..], &[0xc1][..], &b"{\"Eq\":{}}"[..]] { + let err = decode_metadata_filter(bytes).expect_err("malformed filter bytes"); + assert!(matches!(err, crate::Error::BadRequest { .. }), "{err:?}"); + } + } + + #[test] + fn msgpack_filter_decodes() { + let filter = MetadataFilter::eq("category", Value::String("ai".into())); + let bytes = zerompk::to_msgpack_vec(&filter).expect("encode"); + assert_eq!(decode_metadata_filter(&bytes).expect("decode"), filter); + } + + #[test] + fn comparison_lowers_to_column_op_literal() { + let expr = metadata_filter_predicate(&MetadataFilter::Gte { + field: "Score".into(), + value: Value::Integer(3), + }) + .expect("lowers"); + match expr { + SqlExpr::BinaryOp { left, op, right } => { + assert!(matches!(*left, SqlExpr::Column { ref name, .. } if name == "Score")); + assert_eq!(op, BinaryOp::Ge); + assert!(matches!(*right, SqlExpr::Literal(SqlValue::Int(3)))); + } + other => panic!("expected a comparison, got {other:?}"), + } + } + + #[test] + fn compound_filters_lower_to_logical_operators() { + let filter = MetadataFilter::Not(Box::new(MetadataFilter::or(vec![ + MetadataFilter::eq("a", Value::Integer(1)), + MetadataFilter::NotIn { + field: "b".into(), + values: vec![Value::String("x".into())], + }, + ]))); + let text = lowered(&filter); + assert!(text.starts_with("UnaryOp { op: Not"), "{text}"); + assert!(text.contains("op: Or"), "{text}"); + assert!(text.contains("negated: true"), "{text}"); + } + + #[test] + fn empty_and_or_are_true_and_false() { + assert!(matches!( + metadata_filter_predicate(&MetadataFilter::and(vec![])).expect("lowers"), + SqlExpr::Literal(SqlValue::Bool(true)) + )); + assert!(matches!( + metadata_filter_predicate(&MetadataFilter::or(vec![])).expect("lowers"), + SqlExpr::Literal(SqlValue::Bool(false)) + )); + } + + #[test] + fn value_without_literal_form_is_bad_request() { + let filter = MetadataFilter::eq("tags", Value::Array(vec![Value::Integer(1)])); + let err = metadata_filter_predicate(&filter).expect_err("array has no literal"); + assert!(matches!(err, crate::Error::BadRequest { .. }), "{err:?}"); + } +} diff --git a/nodedb/src/data/executor/core_loop/response.rs b/nodedb/src/data/executor/core_loop/response.rs index 490ea743f..c1aecc230 100644 --- a/nodedb/src/data/executor/core_loop/response.rs +++ b/nodedb/src/data/executor/core_loop/response.rs @@ -191,6 +191,75 @@ impl CoreLoop { ) } + /// The key of the vector index a search of `field_name` reads: the + /// field's own index, or the collection-level index when the field has + /// none of its own and the collection-level one exists. Data synced from + /// NodeDB-Lite lives in the collection-level index, while SQL names the + /// field. + /// + /// A search that names no field reads the collection-level index. When + /// the collection has none and exactly one field index, it reads that + /// one: a client API with no field argument searches the collection's + /// only vector column. With several field indexes the search is + /// ambiguous and fails, naming them. + pub(in crate::data::executor) fn resolve_vector_index_key( + &self, + database_id: u64, + tenant_id: u64, + collection: &str, + field_name: &str, + ) -> Result<(DatabaseId, TenantId, String), ErrorCode> { + let index_key = Self::vector_index_key(database_id, tenant_id, collection, field_name); + if self.vector_collections.contains_key(&index_key) { + return Ok(index_key); + } + if field_name.is_empty() { + let mut field_keys = self.field_index_keys(&index_key, collection); + return match field_keys.len() { + 0 => Ok(index_key), + 1 => Ok(field_keys.remove(0)), + _ => { + let prefix_len = collection.len() + 1; + let mut fields: Vec<&str> = field_keys + .iter() + .filter_map(|(_, _, key)| key.get(prefix_len..)) + .collect(); + fields.sort_unstable(); + Err(ErrorCode::BadRequest { + detail: format!( + "vector search of '{collection}' names no field and the \ + collection has several vector indexes ({}); name one", + fields.join(", ") + ), + }) + } + }; + } + let fallback_key = Self::vector_index_key(database_id, tenant_id, collection, ""); + if self.vector_collections.contains_key(&fallback_key) { + Ok(fallback_key) + } else { + Ok(index_key) + } + } + + /// The keys of `collection`'s field indexes. `collection_key` is the + /// collection-level key of the same database and tenant. + fn field_index_keys( + &self, + collection_key: &(DatabaseId, TenantId, String), + collection: &str, + ) -> Vec<(DatabaseId, TenantId, String)> { + let prefix = format!("{collection}:"); + self.vector_collections + .keys() + .filter(|(db, tid, key)| { + *db == collection_key.0 && *tid == collection_key.1 && key.starts_with(&prefix) + }) + .cloned() + .collect() + } + /// In-memory build key for a vector collection: `"{db}:{tid}:{coll}"`. /// `coll` may itself contain `:` — parsing uses `splitn(3, ':')`. /// @@ -219,11 +288,13 @@ impl CoreLoop { let mut engine = TenantCrdtEngine::new(tenant_id, self.core_id as u64, ConstraintSet::new())?; // Rejected deltas leave no replayable record, so their - // dead-letter entries come back from storage. + // dead-letter entries come back from storage. An entry that + // does not decode or does not fit the queue fails the open: + // the engine is not installed without it. engine.restore_dead_letters( self.sparse .load_crdt_dead_letters(database_id.as_u64(), tenant_id.as_u64())?, - ); + )?; Ok(slot.insert(engine)) } } diff --git a/nodedb/src/data/executor/dispatch/bitmap/hashjoin_inline.rs b/nodedb/src/data/executor/dispatch/bitmap/hashjoin_inline.rs index 43627a8e0..33ff4b011 100644 --- a/nodedb/src/data/executor/dispatch/bitmap/hashjoin_inline.rs +++ b/nodedb/src/data/executor/dispatch/bitmap/hashjoin_inline.rs @@ -1,61 +1,67 @@ // SPDX-License-Identifier: BUSL-1.1 -//! HashJoin bitmap-prefilter injection. +//! Bitmap-prefilter sub-plan execution. //! -//! When `QueryOp::HashJoin` carries a `left_bitmap` or -//! `right_bitmap` sub-plan, the executor calls `run_bitmap_subplan` -//! to execute that sub-plan and collect the resulting surrogates into a -//! `SurrogateBitmap`. The bitmap is then used to build a prefiltered -//! `DocumentOp::Scan` for the probe side, replacing the generic -//! `scan_collection` call so non-member rows never enter the hash-join loop. -//! -//! If the sub-plan returns no decodable rows (e.g. the collection doesn't -//! exist or the index lookup yields nothing), `run_bitmap_subplan` returns an -//! empty bitmap — the probe scan proceeds without any prefilter in that case. +//! When `QueryOp::HashJoin` carries a `left_bitmap` or `right_bitmap` +//! sub-plan, or a vector search carries an inline prefilter plan, the +//! executor calls `run_bitmap_subplan` to execute that sub-plan and collect +//! the resulting surrogates into a `SurrogateBitmap`. The bitmap restricts +//! the consumer's candidates: an empty bitmap admits no row. A failing +//! sub-plan is an error, never an empty bitmap. use nodedb_types::SurrogateBitmap; -use crate::bridge::envelope::PhysicalPlan; +use crate::bridge::envelope::{ErrorCode, PhysicalPlan, Status}; use crate::data::executor::core_loop::CoreLoop; use crate::data::executor::task::ExecutionTask; use nodedb_physical::physical_plan::DocumentOp; use super::materialize::collect_surrogates; -/// Execute a bitmap-producer sub-plan and return the resulting `SurrogateBitmap`. +/// Execute a bitmap-producer sub-plan and return the resulting +/// `SurrogateBitmap`. /// -/// Returns an empty bitmap when the sub-plan produces no decodable rows — -/// the caller should treat an empty bitmap as "no prefilter" and fall back -/// to a full probe scan. +/// A sub-plan that fails returns its error. A sub-plan that succeeds with no +/// rows returns an empty bitmap, which admits no row. pub(crate) fn run_bitmap_subplan( core: &mut CoreLoop, task: &ExecutionTask, sub_plan: &PhysicalPlan, -) -> SurrogateBitmap { +) -> crate::Result { let sub_response = core.execute_plan(task, sub_plan); + if sub_response.status == Status::Error { + let code = sub_response.error_code.map_or_else( + || ErrorCode::Internal { + detail: "bitmap sub-plan failed without an error code".into(), + }, + |code| *code, + ); + return Err(crate::Error::DataPlane(code)); + } + if sub_response.payload.is_empty() { + return Ok(SurrogateBitmap::new()); + } let docs = crate::data::executor::response_codec::decode_response_to_docs(&sub_response) - .unwrap_or_default(); - collect_surrogates(&docs) + .ok_or_else(|| { + crate::Error::DataPlane(ErrorCode::Internal { + detail: "bitmap sub-plan returned an undecodable row payload".into(), + }) + })?; + Ok(collect_surrogates(&docs)) } /// Build a `DocumentOp::Scan` physical plan with a surrogate prefilter injected. /// -/// Used by the hash-join executor to replace the plain `scan_collection` call -/// for the probe side when a bitmap sub-plan was provided. The scan will skip -/// rows whose surrogate is absent from `bitmap`, pushing the filter into the -/// document engine before any msgpack decode. -/// -/// If `bitmap` is empty, returns `None` — the caller falls back to the normal -/// `scan_collection` path (no-op: an empty bitmap would block all rows). +/// Used by the hash-join executor for the probe side when a bitmap sub-plan +/// was provided. The scan skips rows whose surrogate is absent from `bitmap`, +/// pushing the filter into the document engine before any msgpack decode. An +/// empty `bitmap` scans no row. pub(crate) fn prefiltered_scan_plan( collection: &str, limit: usize, bitmap: SurrogateBitmap, -) -> Option { - if bitmap.is_empty() { - return None; - } - Some(PhysicalPlan::Document(DocumentOp::Scan { +) -> PhysicalPlan { + PhysicalPlan::Document(DocumentOp::Scan { collection: nodedb_types::QualifiedCollection::from_stored(collection.to_string()), limit, offset: 0, @@ -68,5 +74,117 @@ pub(crate) fn prefiltered_scan_plan( system_time: nodedb_types::SystemTimeScope::Current, valid_at_ms: None, prefilter: Some(bitmap), - })) + }) +} + +#[cfg(test)] +mod tests { + use std::time::{Duration, Instant}; + + use nodedb_bridge::buffer::RingBuffer; + + use super::*; + use crate::bridge::envelope::{Admission, ExemptReason, Priority, Request}; + use crate::types::{DatabaseId, ReadConsistency, RequestId, TenantId, TraceId, VShardId}; + + struct CoreHarness { + core: CoreLoop, + _req_tx: nodedb_bridge::buffer::Producer, + _resp_rx: nodedb_bridge::buffer::Consumer, + _dir: tempfile::TempDir, + } + + fn make_core() -> CoreHarness { + use crate::bridge::dispatch::{BridgeRequest, BridgeResponse}; + let dir = tempfile::tempdir().expect("tempdir"); + let (req_tx, req_rx) = RingBuffer::channel::(64); + let (resp_tx, resp_rx) = RingBuffer::channel::(64); + let core = CoreLoop::open( + 0, + req_rx, + resp_tx, + dir.path(), + std::sync::Arc::new(nodedb_types::OrdinalClock::new()), + crate::data::executor::core_loop::test_governor(), + ) + .expect("open core"); + CoreHarness { + core, + _req_tx: req_tx, + _resp_rx: resp_rx, + _dir: dir, + } + } + + fn scan(collection: &str, filters: Vec) -> PhysicalPlan { + PhysicalPlan::Document(DocumentOp::Scan { + collection: nodedb_types::QualifiedCollection::new(DatabaseId::DEFAULT, collection), + limit: 100, + offset: 0, + sort_keys: Vec::new(), + filters, + distinct: false, + projection: Vec::new(), + computed_columns: Vec::new(), + window_functions: Vec::new(), + system_time: nodedb_types::SystemTimeScope::Current, + valid_at_ms: None, + prefilter: None, + }) + } + + fn task_for(plan: PhysicalPlan) -> ExecutionTask { + ExecutionTask::new(Request { + request_id: RequestId::new(1), + tenant_id: TenantId::new(1), + database_id: DatabaseId::DEFAULT, + vshard_id: VShardId::new(0), + plan, + deadline: Instant::now() + Duration::from_secs(5), + priority: Priority::Normal, + trace_id: TraceId::ZERO, + consistency: ReadConsistency::Strong, + idempotency_key: None, + event_source: crate::event::EventSource::User, + user_roles: Vec::new(), + user_id: None, + statement_digest: None, + txn_id: None, + wal_lsn: None, + resolved_now_ms: None, + commit_hlc: None, + admission: Admission::Exempt(ExemptReason::Read), + }) + } + + /// A sub-plan that fails is an error, never an empty bitmap that would + /// silently answer with zero rows. + #[test] + fn a_failing_sub_plan_is_an_error() { + let mut harness = make_core(); + let sub_plan = scan("cells", vec![0xc1]); + let task = task_for(sub_plan.clone()); + let result = run_bitmap_subplan(&mut harness.core, &task, &sub_plan); + assert!( + matches!(result, Err(crate::Error::DataPlane(_))), + "expected the sub-plan error, got {result:?}" + ); + } + + /// A sub-plan with no rows yields an empty bitmap, and the probe scan it + /// builds admits no row. + #[test] + fn an_empty_sub_plan_admits_no_row() { + let mut harness = make_core(); + let sub_plan = scan("cells", Vec::new()); + let task = task_for(sub_plan.clone()); + let bitmap = run_bitmap_subplan(&mut harness.core, &task, &sub_plan).expect("bitmap"); + assert!(bitmap.is_empty()); + match prefiltered_scan_plan("cells", 10, bitmap) { + PhysicalPlan::Document(DocumentOp::Scan { prefilter, .. }) => { + assert!(prefilter.is_some_and(|bm| bm.is_empty())); + } + other => panic!("expected a prefiltered scan, got {other:?}"), + } + } } diff --git a/nodedb/src/data/executor/handlers/vector_search.rs b/nodedb/src/data/executor/handlers/vector_search.rs index a4bf6fa9a..4bf8f3a8b 100644 --- a/nodedb/src/data/executor/handlers/vector_search.rs +++ b/nodedb/src/data/executor/handlers/vector_search.rs @@ -124,12 +124,7 @@ pub(super) fn encode_hits_response( Ok(payload) => core.response_with_payload(task, payload), Err(e) => { warn!(core = core.core_id, error = %e, "vector search serialization failed"); - core.response_error( - task, - ErrorCode::Internal { - detail: e.to_string(), - }, - ) + core.response_error(task, ErrorCode::from(e)) } } } @@ -153,7 +148,10 @@ pub(in crate::data::executor) struct VectorSearchParams<'a> { pub metric: DistanceMetric, pub filter_bitmap: Option<&'a nodedb_types::SurrogateBitmap>, pub field_name: &'a str, - /// RLS post-candidate filters. Applied after HNSW/IVF returns candidates. + /// Residual row filters: the statement's `WHERE` conjuncts and any read + /// policy, as `ScanFilter` msgpack. The Control Plane applies them to + /// the ranked candidates. The candidate window widens until `top_k` + /// candidates pass them. pub rls_filters: &'a [u8], /// Cross-engine prefilter sub-plan: when `Some`, executed locally and /// its output rows materialized into a `SurrogateBitmap` that is diff --git a/nodedb/src/data/executor/handlers/vector_search_exec.rs b/nodedb/src/data/executor/handlers/vector_search_exec.rs index 674ddce0c..e66636a7e 100644 --- a/nodedb/src/data/executor/handlers/vector_search_exec.rs +++ b/nodedb/src/data/executor/handlers/vector_search_exec.rs @@ -6,10 +6,10 @@ use roaring::RoaringBitmap; use tracing::{debug, warn}; use super::vector_search::{ - VectorSearchParams, build_search_hit, effective_ef, encode_hits_response, - surrogate_bitmap_to_global_ids, + VectorSearchParams, encode_hits_response, surrogate_bitmap_to_global_ids, }; use super::vector_search_ann::{ResolvedAnnOptions, apply_ann_options, quantization_matches}; +use super::vector_search_window::{BodySource, SearchWindow, WindowHits}; use crate::bridge::envelope::{ErrorCode, Response}; use crate::data::executor::core_loop::CoreLoop; use crate::data::executor::task::ExecutionTask; @@ -108,17 +108,27 @@ impl CoreLoop { // `filter_bitmap`. The sub-plan emits document-shaped rows whose // `id` is the cell's surrogate as 8-char zero-padded lowercase // hex; `collect_surrogates` decodes that back into surrogate IDs. - let inline_bitmap = inline_prefilter_plan.map(|sub_plan| { - crate::data::executor::dispatch::bitmap::hashjoin_inline::run_bitmap_subplan( - self, task, sub_plan, - ) - }); + // A failing sub-plan fails the search: an empty bitmap would admit no + // candidate and silently answer with zero rows. + let inline_bitmap = match inline_prefilter_plan + .map(|sub_plan| { + crate::data::executor::dispatch::bitmap::hashjoin_inline::run_bitmap_subplan( + self, task, sub_plan, + ) + }) + .transpose() + { + Ok(bitmap) => bitmap, + Err(e) => return self.response_error(task, e), + }; let effective_filter: Option = match (filter_bitmap.cloned(), inline_bitmap) { (Some(a), Some(b)) => Some(a.intersect(&b)), (Some(a), None) => Some(a), - (None, Some(b)) if !b.is_empty() => Some(b), - _ => None, + // An empty inline bitmap admits no candidate: the prefilter + // matched nothing, so the search returns nothing. + (None, Some(b)) => Some(b), + (None, None) => None, }; let filter_bitmap = effective_filter.as_ref(); debug!(core = self.core_id, %collection, top_k, ef_search, "vector search"); @@ -130,24 +140,12 @@ impl CoreLoop { }; let database_id = task.request.database_id.as_u64(); - let index_key = CoreLoop::vector_index_key(database_id, tid, collection, field_name); - // Every index type is one `VectorCollection`: an IVF-PQ collection // answers from its exact buffer or its trained IVF-PQ index. - // If the specific field-named index does not exist, fall back to the - // empty-field index. This handles data synced from NodeDB-Lite (which - // uses collection-level storage, not named-field storage) being - // searched via a field-specific SQL query (e.g. vector_distance(embedding, ...)). let effective_key = - if !self.vector_collections.contains_key(&index_key) && !field_name.is_empty() { - let fallback_key = CoreLoop::vector_index_key(database_id, tid, collection, ""); - if self.vector_collections.contains_key(&fallback_key) { - fallback_key - } else { - index_key - } - } else { - index_key + match self.resolve_vector_index_key(database_id, tid, collection, field_name) { + Ok(key) => key, + Err(e) => return self.response_error(task, e), }; // An index with no base rows still answers a transaction's own staged // rows: the merge ranks them over an empty base result. @@ -206,15 +204,31 @@ impl CoreLoop { } } - // Over-fetch to accommodate both oversample breadth (for re-rank - // headroom) and RLS post-filter headroom. The two factors are - // multiplied so each can independently request more candidates. + // Over-fetch for oversample re-rank headroom. A residual filter + // starts from a wider window, and `rank_window` widens it until + // `top_k` candidates pass the filter. let fetch_k = if rls_filters.is_empty() { top_k.saturating_mul(oversample) } else { top_k.saturating_mul(2).saturating_mul(oversample).max(20) }; - let ef = effective_ef(ef_search, fetch_k); + // The Control Plane encoded these bytes. Bytes that do not decode + // fail the search: ranking without the filter cannot size the window. + let residual: Vec = if rls_filters.is_empty() { + Vec::new() + } else { + match zerompk::from_msgpack(rls_filters) { + Ok(filters) => filters, + Err(e) => { + return self.response_error( + task, + ErrorCode::Internal { + detail: format!("vector search residual filter decode: {e}"), + }, + ); + } + } + }; // Derive payload bitmap (node-id space) from `(field, value)` // equalities by intersecting per-field equality bitmaps. Returns @@ -273,7 +287,7 @@ impl CoreLoop { // A filter that cannot be serialized fails the search: searching // without it would return rows the filter excludes. - let searched = match combined_bm { + let bitmap_bytes = match combined_bm { Some(local_bm) => { let mut buf = Vec::with_capacity(local_bm.serialized_size()); if let Err(e) = local_bm.serialize_into(&mut buf) { @@ -284,93 +298,46 @@ impl CoreLoop { }, ); } - collection_ref.search_with_bitmap_bytes_and_metric( - query_vector, - fetch_k, - ef, - &buf, - metric, - ) + Some(buf) } - None => collection_ref.search_with_metric(query_vector, fetch_k, ef, metric), - }; - // A query of the wrong dimension is the caller's data error (22000); - // the core keeps serving. - let results = match searched { - Ok(results) => results, - Err(e) => return self.response_error(task, crate::Error::from(e)), + None => None, }; - // Pure-vector fast path: projection contains only id/distance. - // Skip the sparse-store body fetch entirely. - if skip_payload_fetch { - let mut hits: Vec<_> = results - .iter() - .map(|r| build_search_hit(Some(collection_ref), r.id, r.distance)) - .collect(); - // Read-your-own-writes for vector search: fold this - // transaction's staged vector inserts into the base HNSW/IVF - // result before truncation, so a vector inserted earlier in the - // same transaction is ranked in by true distance before COMMIT. - if let Some(txn_id) = task.request.txn_id { - if let Err(e) = self.merge_vector_overlay_into_search( - super::transaction::overlay::VectorMergeParams { - txn_id, - database_id: task.request.database_id, - tid: crate::types::TenantId::new(tid), - collection, - field_name, - query_vector, - metric, - top_k, - filter_bitmap, - payload_filters, - }, - &mut hits, - ) { - return self.response_error(task, e); - } - } else { - hits.truncate(top_k); - } - if let Some(ref m) = self.metrics { - m.record_vector_search(0); - m.record_query_by_engine("vector"); - } - return encode_hits_response(self, task, &hits); - } - - // RLS evaluation lives at the Control-Plane response boundary - // (`response_translate::vector`). DP attaches the document body - // when filters are active so CP can run the predicate without - // a follow-up round-trip; CP applies the filter and truncates to - // `top_k`. Data Plane stays pure SIMD + sparse-fetch. - // Attach body bytes whenever skip_payload_fetch is false (slow path) - // OR when RLS filters need them; the CP response translator flattens - // the bytes' fields into the hit JSON for client column projection. + // The residual filter is evaluated at the Control-Plane response + // boundary (`response_translate::vector`): the Data Plane attaches + // each body so the Control Plane runs the predicate without a second + // round-trip, then truncates to `top_k`. The slow path (no + // `skip_payload_fetch`) also attaches bodies: the response translator + // flattens their fields into the hit for column projection. The + // pure-vector fast path fetches no body. let attach = !skip_payload_fetch || !rls_filters.is_empty(); - let hits: crate::Result> = results - .iter() - .map(|r| build_search_hit(Some(collection_ref), r.id, r.distance)) - .map(|hit| { - self.attach_body( - task.request.database_id.as_u64(), - tid, - collection, - attach, - hit, - ) - }) - .collect(); - let mut hits = match hits { - Ok(hits) => hits, + let window = self.rank_window( + &SearchWindow { + collection: collection_ref, + query_vector, + ef_search, + metric, + bitmap: bitmap_bytes.as_deref(), + }, + &BodySource { + database_id: task.request.database_id.as_u64(), + tid, + collection, + attach, + }, + &residual, + top_k, + fetch_k, + ); + // A query of the wrong dimension is the caller's data error (22000); + // the core keeps serving. + let WindowHits { mut hits, fetch_k } = match window { + Ok(window) => window, Err(e) => return self.response_error(task, e), }; - let truncate_to = if rls_filters.is_empty() { - top_k - } else { - fetch_k - }; + // With a residual filter the Control Plane makes the top-k cut, so + // the whole window travels. + let truncate_to = if residual.is_empty() { top_k } else { fetch_k }; // Read-your-own-writes for vector search: fold this transaction's // staged vector inserts into the base HNSW/IVF result before // truncation, so a vector inserted earlier in the same transaction diff --git a/nodedb/src/data/executor/handlers/vector_search_window.rs b/nodedb/src/data/executor/handlers/vector_search_window.rs new file mode 100644 index 000000000..51f0fcbb8 --- /dev/null +++ b/nodedb/src/data/executor/handlers/vector_search_window.rs @@ -0,0 +1,113 @@ +// SPDX-License-Identifier: BUSL-1.1 + +//! Candidate window of a vector search. +//! +//! A search with a residual row filter (`rls_filters`: the statement's +//! `WHERE` conjuncts and any read policy) ranks more candidates than it +//! returns. The Control Plane evaluates the filter over the attached bodies +//! and keeps the `top_k` nearest survivors. A fixed over-fetch returns fewer +//! than `top_k` rows whenever the nearest candidates fail the filter. The +//! window therefore doubles until `top_k` candidates pass the filter or the +//! index holds no further candidate, so the filter narrows the candidates +//! before the top-k cut. + +use nodedb_query::scan_filter::ScanFilter; + +use super::super::response_codec::VectorSearchHit; +use super::vector_search::{build_search_hit, effective_ef}; +use crate::data::executor::core_loop::CoreLoop; +use crate::engine::vector::collection::VectorCollection; +use crate::engine::vector::distance::DistanceMetric; + +/// One ranking request against a vector index. +pub(super) struct SearchWindow<'a> { + pub collection: &'a VectorCollection, + pub query_vector: &'a [f32], + pub ef_search: usize, + pub metric: DistanceMetric, + /// Serialized candidate bitmap in the index's node-id space. + pub bitmap: Option<&'a [u8]>, +} + +/// Where the bodies of ranked candidates are read from. +pub(super) struct BodySource<'a> { + pub database_id: u64, + pub tid: u64, + pub collection: &'a str, + /// Attach each candidate's body to its hit. + pub attach: bool, +} + +/// The ranked hits and the window width that produced them. +pub(super) struct WindowHits { + pub hits: Vec, + pub fetch_k: usize, +} + +impl CoreLoop { + /// Rank `fetch_k` candidates and attach their bodies. + /// + /// With a non-empty `residual`, the window doubles until `top_k` hits + /// pass it or the index returns fewer candidates than the window asked + /// for. A hit passes when the Control Plane will keep it: it has a body + /// and every filter matches that body. The hits are not filtered here. + pub(super) fn rank_window( + &self, + window: &SearchWindow<'_>, + bodies: &BodySource<'_>, + residual: &[ScanFilter], + top_k: usize, + mut fetch_k: usize, + ) -> crate::Result { + let live = window.collection.live_count(); + loop { + let ef = effective_ef(window.ef_search, fetch_k); + let ranked = match window.bitmap { + Some(bitmap) => window.collection.search_with_bitmap_bytes_and_metric( + window.query_vector, + fetch_k, + ef, + bitmap, + window.metric, + ), + None => window.collection.search_with_metric( + window.query_vector, + fetch_k, + ef, + window.metric, + ), + } + .map_err(crate::Error::from)?; + let exhausted = ranked.len() < fetch_k || fetch_k >= live; + let hits = ranked + .iter() + .map(|r| build_search_hit(Some(window.collection), r.id, r.distance)) + .map(|hit| { + self.attach_body( + bodies.database_id, + bodies.tid, + bodies.collection, + bodies.attach, + hit, + ) + }) + .collect::>>()?; + if residual.is_empty() || exhausted || passing(&hits, residual) >= top_k { + return Ok(WindowHits { hits, fetch_k }); + } + fetch_k = fetch_k.saturating_mul(2); + } + } +} + +/// Hits the Control Plane keeps under `residual`. A filter that fails to +/// evaluate against a body drops that hit there, so it does not count here. +fn passing(hits: &[VectorSearchHit], residual: &[ScanFilter]) -> usize { + hits.iter() + .filter(|hit| { + hit.body + .as_deref() + .is_some_and(|body| ScanFilter::all_match_binary(residual, body).unwrap_or(false)) + }) + .count() +} diff --git a/nodedb/tests/wire/cases/sql_where_vector_search.rs b/nodedb/tests/wire/cases/sql_where_vector_search.rs index 258970cf2..ef50d4356 100644 --- a/nodedb/tests/wire/cases/sql_where_vector_search.rs +++ b/nodedb/tests/wire/cases/sql_where_vector_search.rs @@ -236,3 +236,82 @@ async fn arrow_distance_in_where_does_not_silently_match_none_for_delete() { remaining.len() ); } + +const RLS_PASSWORD: &str = "vec-where-rls-secret-7"; + +/// Run `sql` as `user` and return the first column of each delivered row. +async fn ids_as(server: &TestServer, user: &str, sql: &str) -> Vec { + let (client, handle) = server + .connect_as(user, RLS_PASSWORD) + .await + .unwrap_or_else(|e| panic!("connect as {user}: {e}")); + let messages = client + .simple_query(sql) + .await + .unwrap_or_else(|e| panic!("{user} runs {sql}: {e}")); + let mut out = Vec::new(); + for message in messages { + if let tokio_postgres::SimpleQueryMessage::Row(row) = message { + out.push(row.get(0).unwrap_or("").to_string()); + } + } + drop(client); + handle.abort(); + out +} + +/// The statement's own WHERE predicate and a read policy share the vector +/// search's post-filter slot. The policy joins the predicate and never +/// replaces it. +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn a_where_filter_survives_a_read_policy() { + let server = TestServer::start().await; + server.exec("CREATE COLLECTION vec_rls").await.unwrap(); + server + .exec("CREATE VECTOR INDEX idx_vec_rls_emb ON vec_rls METRIC cosine DIM 4") + .await + .unwrap(); + for (id, owner, tag, v) in [ + ("m1", "vec_rls_reader", "a", "0.10, 0.20, 0.30, 0.40"), + ("m2", "vec_rls_reader", "b", "0.11, 0.21, 0.31, 0.41"), + ("m3", "alice", "a", "0.12, 0.22, 0.32, 0.42"), + ] { + server + .exec(&format!( + "INSERT INTO vec_rls (id, owner, tag, embedding) \ + VALUES ('{id}', '{owner}', '{tag}', ARRAY[{v}])" + )) + .await + .unwrap(); + } + server + .exec(&format!( + "CREATE USER vec_rls_reader PASSWORD '{RLS_PASSWORD}'" + )) + .await + .unwrap(); + server + .exec("GRANT ROLE readwrite TO vec_rls_reader") + .await + .unwrap(); + server + .exec( + "CREATE RLS POLICY vec_rls_owner ON vec_rls FOR READ \ + USING (owner = $auth.username)", + ) + .await + .unwrap(); + + let ids = ids_as( + &server, + "vec_rls_reader", + "SELECT id FROM vec_rls \ + WHERE tag = 'a' AND embedding <-> ARRAY[0.1, 0.2, 0.3, 0.4] LIMIT 10", + ) + .await; + assert_eq!( + ids, + vec!["m1".to_string()], + "both the tag filter and the owner policy must hold" + ); +} From 5682fe69c07ca3eef5d3145514108307cbc29d10 Mon Sep 17 00:00:00 2001 From: Farhan Syah Date: Thu, 8 Oct 2026 16:51:47 +0800 Subject: [PATCH 17/24] feat(graph): follow label sets and filter walks by edge properties LABEL takes a list of labels on GRAPH TRAVERSE, GRAPH PATH and the CSR walks. A trailing EDGE WHERE clause crosses only edges whose properties match. Edge properties are stored as plain MessagePack. GRAPH TRAVERSE returns each node's depth and the crossed edges with their properties. Walks merge the session transaction's staged edges on every hop and drop a start node the graph does not hold. An unknown direction is refused. --- docs/graph.md | 29 +- docs/query-language.md | 4 + .../cases/sql_cluster_cross_node_dml.rs | 2 + .../graph_edge_predicate_cross_node.rs | 143 ++++ .../graph_rag_fusion_cross_node.rs | 2 +- .../native_implicit_edge_delete_cross_node.rs | 4 +- nodedb-graph/src/bfs_params.rs | 4 +- nodedb-graph/src/csr/index/label_filter.rs | 160 +++- nodedb-graph/src/csr/index/lookup.rs | 200 ++--- nodedb-graph/src/csr/index/restore.rs | 12 +- nodedb-graph/src/csr/index/scoped.rs | 5 +- nodedb-graph/src/csr/index/types.rs | 30 +- nodedb-graph/src/csr/persist.rs | 2 +- nodedb-graph/src/csr/rebuild/install.rs | 2 +- nodedb-graph/src/overlay_delta.rs | 64 +- nodedb-graph/src/path_overlay.rs | 428 ++++++---- nodedb-graph/src/path_params.rs | 3 +- nodedb-graph/src/test_support.rs | 26 + nodedb-graph/src/traversal.rs | 764 ------------------ nodedb-graph/src/traversal/bfs.rs | 452 +++++++++++ nodedb-graph/src/traversal/mod.rs | 14 + nodedb-graph/src/traversal/shortest_path.rs | 460 +++++++++++ nodedb-graph/src/traversal/subgraph.rs | 337 ++++++++ nodedb-graph/src/traversal_overlay.rs | 564 +++++++++---- nodedb-graph/src/traversal_surrogate.rs | 72 +- nodedb-physical/src/physical_plan/graph/op.rs | 23 +- nodedb-query/src/metadata_filter.rs | 623 ++++++++++++-- .../src/ddl_ast/graph_parse/edge_predicate.rs | 504 ++++++++++++ nodedb-sql/src/ddl_ast/graph_parse/entry.rs | 147 +++- nodedb-sql/src/ddl_ast/graph_parse/helpers.rs | 36 +- nodedb-sql/src/ddl_ast/graph_parse/mod.rs | 1 + .../src/ddl_ast/graph_parse/variants.rs | 29 +- .../src/ddl_ast/statement/types/graph.rs | 13 +- nodedb-test-support/src/tx_batch_helpers.rs | 2 +- nodedb-types/src/graph.rs | 122 ++- .../control/planner/redaction_refusal/plan.rs | 6 +- .../control/planner/rls_injection/graph.rs | 15 +- .../rls_injection/permission_tree/graph.rs | 2 +- .../src/control/server/graph_dispatch/bfs.rs | 62 +- .../src/control/server/graph_dispatch/hop.rs | 345 +++----- .../src/control/server/graph_dispatch/mod.rs | 3 + .../server/graph_dispatch/neighbor_rows.rs | 187 +++++ .../control/server/graph_dispatch/presence.rs | 62 ++ .../server/graph_dispatch/rag_fusion/coord.rs | 19 +- .../server/graph_dispatch/read_groups.rs | 4 +- .../server/graph_dispatch/shortest_path.rs | 97 ++- .../server/graph_dispatch/subgraph_edges.rs | 170 ++++ .../graph_dispatch/traverse_subgraph.rs | 265 ++++-- .../server/graph_dispatch/walk_reads.rs | 47 +- .../server/native/dispatch/graph_owner.rs | 1 + .../native/dispatch/plan_builder/graph.rs | 106 ++- .../native/dispatch/plan_builder/helpers.rs | 59 +- .../src/control/server/response_shape/walk.rs | 2 +- .../shared/ddl/neutral/graph_ops/dispatch.rs | 84 +- .../shared/ddl/neutral/graph_ops/edge.rs | 12 +- .../ddl/neutral/graph_ops/edge_parse.rs | 101 ++- .../ddl/neutral/graph_ops/edge_stage.rs | 2 +- .../shared/ddl/neutral/graph_ops/traverse.rs | 88 +- .../shared/ddl/neutral/tree_ops/children.rs | 4 +- .../server/shared/ddl/neutral/tree_ops/sum.rs | 4 +- .../src/control/server/wal_dispatch/graph.rs | 2 +- .../data/executor/cascade_journal_tests.rs | 2 +- .../src/data/executor/dispatch/graph/mod.rs | 9 + .../executor/dispatch/graph/node_labels.rs | 369 +++++++++ .../dispatch/{graph.rs => graph/route.rs} | 206 +---- .../executor/graph_label_checkpoint/load.rs | 44 +- nodedb/src/data/executor/handlers/graph.rs | 331 +++++--- .../executor/handlers/graph_edge_predicate.rs | 397 +++++++++ .../data/executor/handlers/graph_expansion.rs | 4 +- .../src/data/executor/handlers/graph_rag.rs | 9 +- .../executor/handlers/graph_rag_triple.rs | 2 +- .../data/executor/handlers/graph_traversal.rs | 83 +- .../data/executor/handlers/graph_txn_merge.rs | 582 +++++++------ nodedb/src/data/executor/handlers/mod.rs | 8 + .../data/executor/handlers/rls_write_gate.rs | 100 ++- .../transaction/overlay/graph_staged/edges.rs | 16 + .../src/data/executor/response_codec/hits.rs | 23 +- .../data/executor/wal_replay_graph_labels.rs | 41 +- .../executor/wal_replay_redo_graph_cut.rs | 2 +- .../executor/wal_replay_row_image_tests.rs | 2 +- nodedb/src/engine/graph/csr/rebuild.rs | 2 +- nodedb/src/engine/graph/edge_store/mod.rs | 8 +- .../engine/graph/edge_store/temporal/mod.rs | 1 + .../engine/graph/edge_store/temporal/read.rs | 127 ++- .../graph/pattern/executor/varlen_named.rs | 4 +- nodedb/src/lib.rs | 2 + nodedb/tests/inproc/cases/core_loop.rs | 6 + .../inproc/cases/executor_tests/test_graph.rs | 175 +++- .../cases/executor_tests/test_graph_bounds.rs | 10 +- .../test_graph_edge_predicate.rs | 210 +++++ .../test_graph_txn_overlay_no_partition.rs | 164 ++++ .../test_graph_txn_overlay_reads.rs | 372 +++++++++ .../test_security_and_isolation.rs | 2 +- .../test_tenant_isolation_graph.rs | 4 +- .../test_tenant_isolation_graph_negative.rs | 6 +- .../cases/executor_tests/test_tenant_purge.rs | 2 +- .../test_transaction_cross_engine.rs | 6 +- .../test_transaction_matrix_helpers.rs | 2 +- .../cases/graph_collection_isolation.rs | 6 +- .../inproc/cases/surrogate_round_trip.rs | 4 +- .../tests/wire/cases/graph_dsl_label_sets.rs | 157 ++++ .../tests/wire/cases/graph_edge_predicate.rs | 293 +++++++ .../wire/cases/graph_traverse_absent_start.rs | 90 +++ .../sql_transactions_graph_walk_overlay.rs | 160 ++++ 104 files changed, 8455 insertions(+), 2620 deletions(-) create mode 100644 nodedb-cluster-tests/tests/sql_cluster_cross_node_dml_tests/graph_edge_predicate_cross_node.rs delete mode 100644 nodedb-graph/src/traversal.rs create mode 100644 nodedb-graph/src/traversal/bfs.rs create mode 100644 nodedb-graph/src/traversal/mod.rs create mode 100644 nodedb-graph/src/traversal/shortest_path.rs create mode 100644 nodedb-graph/src/traversal/subgraph.rs create mode 100644 nodedb-sql/src/ddl_ast/graph_parse/edge_predicate.rs create mode 100644 nodedb/src/control/server/graph_dispatch/neighbor_rows.rs create mode 100644 nodedb/src/control/server/graph_dispatch/presence.rs create mode 100644 nodedb/src/control/server/graph_dispatch/subgraph_edges.rs create mode 100644 nodedb/src/data/executor/dispatch/graph/mod.rs create mode 100644 nodedb/src/data/executor/dispatch/graph/node_labels.rs rename nodedb/src/data/executor/dispatch/{graph.rs => graph/route.rs} (67%) create mode 100644 nodedb/src/data/executor/handlers/graph_edge_predicate.rs create mode 100644 nodedb/tests/inproc/cases/executor_tests/test_graph_edge_predicate.rs create mode 100644 nodedb/tests/inproc/cases/executor_tests/test_graph_txn_overlay_no_partition.rs create mode 100644 nodedb/tests/inproc/cases/executor_tests/test_graph_txn_overlay_reads.rs create mode 100644 nodedb/tests/wire/cases/graph_dsl_label_sets.rs create mode 100644 nodedb/tests/wire/cases/graph_edge_predicate.rs create mode 100644 nodedb/tests/wire/cases/graph_traverse_absent_start.rs create mode 100644 nodedb/tests/wire/cases/sql_transactions_graph_walk_overlay.rs 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/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/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-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/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-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-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 '