diff --git a/rust/lancedb/src/table.rs b/rust/lancedb/src/table.rs index f1bd4fd2c..c30ea1b87 100644 --- a/rust/lancedb/src/table.rs +++ b/rust/lancedb/src/table.rs @@ -91,6 +91,8 @@ mod append_generated_column_invalidation_contract; #[cfg(test)] mod delete_generated_column_invalidation_contract; #[cfg(test)] +mod merge_insert_generated_column_reject_contract; +#[cfg(test)] mod schema_metadata_updates_dependency_contract; #[cfg(test)] mod update_generated_column_invalidation_contract; diff --git a/rust/lancedb/src/table/generated_column_invalidation.rs b/rust/lancedb/src/table/generated_column_invalidation.rs index a7444287f..0202a8e94 100644 --- a/rust/lancedb/src/table/generated_column_invalidation.rs +++ b/rust/lancedb/src/table/generated_column_invalidation.rs @@ -1,12 +1,13 @@ // SPDX-License-Identifier: Apache-2.0 // SPDX-FileCopyrightText: Copyright The LanceDB Authors -//! Crate-private Native wiring for generated-column invalidation (B4b / B4c / B4d). +//! Crate-private Native wiring for generated-column invalidation (B4b / B4c / B4d / B4e). //! //! Converts the B4a pure planner into one Lance field-metadata patch for Native //! append, update, and delete commits. Planning is strict-decode/validate; -//! overwrite of a table with any generated-column definition, and direct writes -//! of generated outputs via Update, fail closed as [`Error::NotSupported`]. +//! overwrite of a table with any generated-column definition, direct writes of +//! generated outputs via Update, and Native merge-insert (standard and LSM) +//! fail closed as [`Error::NotSupported`]. use std::collections::{BTreeSet, HashMap}; @@ -116,6 +117,30 @@ pub(super) fn plan_native_delete_generated_column_invalidation( Ok(Some(planned_invalidation_to_schema_metadata_updates(plan))) } +/// Fail closed before Native `merge_insert` when any generated column is present. +/// +/// Strict-decodes and validates every present generated-column metadata value +/// through the B4a `RowSetChanged` planner against one exact dataset snapshot. +/// Malformed metadata returns the existing [`Error::InvalidInput`] validation +/// category. When at least one valid generated column is present, returns +/// [`Error::NotSupported`] before LSM dispatch or source iteration. Ordinary +/// tables (no generated metadata) return `Ok(())`. +pub(super) fn reject_native_merge_insert_if_generated_columns_present( + dataset: &Dataset, +) -> Result<()> { + let snapshot = generated_column_binding_snapshot_from_dataset(dataset)?; + let plan = plan_generated_column_invalidation( + &snapshot, + &GeneratedColumnMutationImpact::RowSetChanged, + )?; + if plan.is_empty() { + return Ok(()); + } + Err(Error::NotSupported { + message: "Merge insert is not supported on tables with generated columns".to_string(), + }) +} + /// Convert planner replacements into one non-empty Lance field-metadata patch. /// /// Each entry is keyed by stable output field ID, uses `replace: false`, and diff --git a/rust/lancedb/src/table/merge.rs b/rust/lancedb/src/table/merge.rs index 3a5b6882d..4d51ab624 100644 --- a/rust/lancedb/src/table/merge.rs +++ b/rust/lancedb/src/table/merge.rs @@ -233,10 +233,22 @@ pub(crate) async fn execute_merge_insert( params: MergeInsertBuilder, new_data: Box, ) -> Result { - match lsm::lsm_dispatch_decision(table, ¶ms).await? { + // One exact dataset supplies the generated-column fail-closed guard and + // downstream standard/LSM routing/execution. Do not refetch after the guard. + let dataset = table.dataset.get().await?; + super::generated_column_invalidation::reject_native_merge_insert_if_generated_columns_present( + dataset.as_ref(), + )?; + + match lsm::lsm_dispatch_decision(¶ms, dataset.as_ref()).await? { lsm::LsmDispatch::Lsm(plan) => { - let future = - lsm::execute_lsm_merge_insert(table, plan, params.validate_single_shard, new_data); + let future = lsm::execute_lsm_merge_insert( + table, + plan, + params.validate_single_shard, + new_data, + dataset, + ); return match params.timeout { Some(timeout) => match tokio::time::timeout(timeout, future).await { Ok(result) => result, @@ -250,7 +262,6 @@ pub(crate) async fn execute_merge_insert( lsm::LsmDispatch::Standard => {} } - let dataset = table.dataset.get().await?; let mut builder = LanceMergeInsertBuilder::try_new(dataset.clone(), params.on)?; match ( params.when_matched_update_all, diff --git a/rust/lancedb/src/table/merge/lsm.rs b/rust/lancedb/src/table/merge/lsm.rs index 87c427b3c..2c8c8c171 100644 --- a/rust/lancedb/src/table/merge/lsm.rs +++ b/rust/lancedb/src/table/merge/lsm.rs @@ -531,18 +531,18 @@ pub(crate) enum LsmDispatch { } /// Decide whether a `merge_insert` should be routed through the MemWAL write -/// path, validating the builder against the installed spec. +/// path, validating the builder against the installed spec on the exact +/// caller-supplied dataset snapshot. #[allow(clippy::redundant_pub_crate)] pub(crate) async fn lsm_dispatch_decision( - table: &NativeTable, params: &MergeInsertBuilder, + dataset: &Dataset, ) -> Result { // Explicit opt-out: use the standard path regardless of any installed spec. if params.use_lsm == Some(false) { return Ok(LsmDispatch::Standard); } - let dataset = table.dataset.get().await?; let Some(details) = dataset.mem_wal_index_details().await? else { // No write spec installed. `use_lsm(true)` demanded MemWAL routing, so // that is an error; otherwise fall back to the standard path. @@ -646,14 +646,17 @@ fn resolve_lsm_mode(details: &MemWalIndexDetails) -> Result { /// a validation failure (e.g. input spanning shards) never leaves a partial /// write behind. When `validate_single_shard` is set, every row is checked to /// route to one shard; when disabled, only the first row of the whole input is. +/// +/// `dataset` must be the same exact snapshot used for the generated-column +/// guard and [`lsm_dispatch_decision`]. #[allow(clippy::redundant_pub_crate)] pub(crate) async fn execute_lsm_merge_insert( table: &NativeTable, plan: LsmPlan, validate_single_shard: bool, new_data: Box, + dataset: Arc, ) -> Result { - let dataset = table.dataset.get().await?; let target_schema: SchemaRef = Arc::new(ArrowSchema::from(dataset.schema())); // Collect, align and shard-validate the whole input before writing diff --git a/rust/lancedb/src/table/merge_insert_generated_column_reject_contract.rs b/rust/lancedb/src/table/merge_insert_generated_column_reject_contract.rs new file mode 100644 index 000000000..ec5f73fec --- /dev/null +++ b/rust/lancedb/src/table/merge_insert_generated_column_reject_contract.rs @@ -0,0 +1,525 @@ +// SPDX-License-Identifier: Apache-2.0 +// SPDX-FileCopyrightText: Copyright The LanceDB Authors + +//! RED runtime contract tests for Native merge-insert fail-closed guard (B4e). +//! +//! Tables with generated-column definitions cannot carry dependency-epoch +//! metadata updates through Native merge-insert in this slice. Both the +//! standard and MemWAL/LSM routes must reject before consuming source input or +//! mutating the table. Ordinary tables keep existing merge-insert semantics. + +use std::sync::Arc; +use std::sync::atomic::{AtomicUsize, Ordering}; + +use arrow_array::{Int32Array, RecordBatch, RecordBatchReader, StringArray}; +use arrow_schema::{ArrowError, DataType, Field, Schema, SchemaRef}; +use futures::TryStreamExt; +use tempfile::TempDir; + +use crate::connection::ConnectBuilder; +use crate::error::Error; +use crate::function::{ + Function, FunctionArgument, FunctionCall, FunctionId, FunctionOutput, FunctionParameter, + FunctionSignature, GENERATED_COLUMN_METADATA_KEY, GeneratedColumnDefinition, + GeneratedColumnStatus, +}; +use crate::query::{ExecutableQuery, QueryBase, Select}; +use crate::table::Table; +use crate::table::schema_evolution::FieldMetadataUpdate; + +const ID: &str = "id"; +const ORDINARY: &str = "ordinary"; +const GEN_OUT: &str = "gen_out"; +const INITIAL_DEPENDENCY_EPOCH: u64 = 3; +const INITIAL_MATERIALIZED_EPOCH: u64 = 3; +const FN_ID: &str = "fn.exact.b4e.merge.literal"; +const MALFORMED_MARKER: &str = "SENSITIVE_B4E_MERGE_METADATA_MARKER_7c91_e2ab"; + +struct Fixture { + _tmp: TempDir, + table: Table, + table_name: String, + uri: String, +} + +/// RecordBatchReader that counts how many times [`Self::next`] is called. +struct ObservableReader { + inner: Box, + next_calls: Arc, +} + +impl ObservableReader { + fn wrap( + inner: Box, + next_calls: Arc, + ) -> Box { + Box::new(Self { inner, next_calls }) + } +} + +impl Iterator for ObservableReader { + type Item = Result; + + fn next(&mut self) -> Option { + self.next_calls.fetch_add(1, Ordering::SeqCst); + self.inner.next() + } +} + +impl RecordBatchReader for ObservableReader { + fn schema(&self) -> SchemaRef { + self.inner.schema() + } +} + +fn literal_only_function() -> Function { + Function::new( + FunctionId::try_new(FN_ID).unwrap(), + FunctionSignature::try_new( + vec![FunctionParameter::new("label", DataType::Utf8)], + FunctionOutput::new(DataType::Int32, true), + ) + .unwrap(), + ) +} + +fn literal_only_definition(output_field_id: i32) -> GeneratedColumnDefinition { + let function = literal_only_function(); + let call = FunctionCall::try_new( + &function, + vec![( + "label".to_string(), + FunctionArgument::try_literal( + Arc::new(StringArray::from(vec![Some("literal-only")])) as arrow_array::ArrayRef + ) + .unwrap(), + )], + ) + .unwrap(); + GeneratedColumnDefinition::try_new( + output_field_id, + call, + INITIAL_DEPENDENCY_EPOCH, + INITIAL_MATERIALIZED_EPOCH, + ) + .unwrap() +} + +fn seed_batch() -> RecordBatch { + let schema = Arc::new(Schema::new(vec![ + Field::new(ID, DataType::Int32, false), + Field::new(ORDINARY, DataType::Utf8, true), + Field::new(GEN_OUT, DataType::Int32, true), + ])); + RecordBatch::try_new( + schema, + vec![ + Arc::new(Int32Array::from(vec![1, 2])), + Arc::new(StringArray::from(vec![Some("a"), Some("b")])), + Arc::new(Int32Array::from(vec![10, 20])), + ], + ) + .unwrap() +} + +fn source_batch(ids: &[i32], ordinary: &[&str], gen_values: &[i32]) -> RecordBatch { + assert_eq!(ids.len(), ordinary.len()); + assert_eq!(ids.len(), gen_values.len()); + let schema = Arc::new(Schema::new(vec![ + Field::new(ID, DataType::Int32, false), + Field::new(ORDINARY, DataType::Utf8, true), + Field::new(GEN_OUT, DataType::Int32, true), + ])); + RecordBatch::try_new( + schema, + vec![ + Arc::new(Int32Array::from(ids.to_vec())), + Arc::new(StringArray::from( + ordinary + .iter() + .map(|value| Some(*value)) + .collect::>(), + )), + Arc::new(Int32Array::from(gen_values.to_vec())), + ], + ) + .unwrap() +} + +fn boxed_reader(batch: RecordBatch) -> Box { + let schema = batch.schema(); + Box::new(arrow_array::RecordBatchIterator::new( + vec![Ok(batch)].into_iter(), + schema, + )) +} + +fn empty_reader() -> Box { + let schema = Arc::new(Schema::new(vec![ + Field::new(ID, DataType::Int32, false), + Field::new(ORDINARY, DataType::Utf8, true), + Field::new(GEN_OUT, DataType::Int32, true), + ])); + Box::new(arrow_array::RecordBatchIterator::new( + std::iter::empty::>(), + schema, + )) +} + +async fn create_ordinary_table(name: &str) -> Fixture { + let tmp = tempfile::tempdir().unwrap(); + let uri = tmp.path().to_str().unwrap().to_string(); + let conn = ConnectBuilder::new(&uri).execute().await.unwrap(); + let table = conn + .create_table(name, seed_batch()) + .execute() + .await + .unwrap(); + Fixture { + _tmp: tmp, + table, + table_name: name.to_string(), + uri, + } +} + +async fn create_table_with_complete_literal_generated(name: &str) -> Fixture { + let fixture = create_ordinary_table(name).await; + let snapshot = fixture + .table + .generated_column_binding_snapshot() + .await + .unwrap(); + let field_id = snapshot.field(GEN_OUT).expect(GEN_OUT).field_id(); + let definition = literal_only_definition(field_id); + let json = definition.to_metadata_json().unwrap(); + fixture + .table + .update_field_metadata(&[ + FieldMetadataUpdate::new(GEN_OUT).set(GENERATED_COLUMN_METADATA_KEY, json) + ]) + .await + .unwrap(); + assert_eq!( + fixture + .table + .generated_column_status(GEN_OUT) + .await + .unwrap(), + GeneratedColumnStatus::Complete + ); + fixture +} + +async fn read_generated_definition(table: &Table) -> GeneratedColumnDefinition { + let snapshot = table.generated_column_binding_snapshot().await.unwrap(); + snapshot + .field(GEN_OUT) + .expect(GEN_OUT) + .generated_column_definition() + .expect("generated metadata must decode") + .expect("generated metadata must be present") +} + +async fn read_raw_generated_metadata(table: &Table) -> String { + let snapshot = table.generated_column_binding_snapshot().await.unwrap(); + snapshot + .field(GEN_OUT) + .expect(GEN_OUT) + .field() + .metadata() + .get(GENERATED_COLUMN_METADATA_KEY) + .expect("generated metadata key must be present") + .clone() +} + +async fn ordinary_rows(table: &Table) -> Vec<(i32, String)> { + let batches = table + .query() + .select(Select::columns(&[ID, ORDINARY])) + .execute() + .await + .unwrap() + .try_collect::>() + .await + .unwrap(); + let mut rows = Vec::new(); + for batch in batches { + let ids = batch + .column(0) + .as_any() + .downcast_ref::() + .unwrap(); + let ordinary = batch + .column(1) + .as_any() + .downcast_ref::() + .unwrap(); + for index in 0..batch.num_rows() { + rows.push((ids.value(index), ordinary.value(index).to_string())); + } + } + rows.sort_by_key(|(id, _)| *id); + rows +} + +fn assert_not_supported(err: &Error, label: &str) { + assert!( + matches!(err, Error::NotSupported { .. }), + "{label}: expected NotSupported, got {err:?}" + ); +} + +fn assert_invalid_input_redacted(err: &Error, planted_raw: &str, label: &str) { + assert!( + matches!(err, Error::InvalidInput { .. }), + "{label}: expected InvalidInput, got {err:?}" + ); + let rendered = format!("{err}\n{err:?}"); + assert!( + !rendered.contains(MALFORMED_MARKER), + "{label}: diagnostic echoed unique metadata marker: {rendered}" + ); + assert!( + !rendered.contains(FN_ID), + "{label}: diagnostic echoed Function ID: {rendered}" + ); + assert!( + !rendered.contains(GENERATED_COLUMN_METADATA_KEY), + "{label}: diagnostic echoed metadata wire key: {rendered}" + ); + assert!( + !rendered.contains(planted_raw), + "{label}: diagnostic echoed raw metadata JSON: {rendered}" + ); +} + +fn configure_standard_merge(builder: &mut crate::table::merge::MergeInsertBuilder) { + builder + .when_matched_update_all(None) + .when_not_matched_insert_all() + .when_not_matched_by_source_delete(None); +} + +#[tokio::test] +async fn standard_merge_insert_rejects_when_generated_column_present_before_input_consumption() { + let fixture = create_table_with_complete_literal_generated("b4e_standard_reject").await; + let version_before = fixture.table.version().await.unwrap(); + let rows_before = ordinary_rows(&fixture.table).await; + let definition_before = read_generated_definition(&fixture.table).await; + let raw_before = read_raw_generated_metadata(&fixture.table).await; + assert_eq!( + definition_before.function_call().function_id().as_str(), + FN_ID + ); + + let next_calls = Arc::new(AtomicUsize::new(0)); + let reader = ObservableReader::wrap( + boxed_reader(source_batch(&[1, 3], &["updated", "inserted"], &[11, 30])), + next_calls.clone(), + ); + + let mut builder = fixture.table.merge_insert(&[ID]); + configure_standard_merge(&mut builder); + let err = builder + .execute(reader) + .await + .expect_err("generated-column table must reject standard merge_insert"); + assert_not_supported(&err, "standard merge_insert generated reject"); + assert_eq!( + next_calls.load(Ordering::SeqCst), + 0, + "rejection must occur before consuming the RecordBatchReader" + ); + + assert_eq!(fixture.table.version().await.unwrap(), version_before); + assert_eq!(ordinary_rows(&fixture.table).await, rows_before); + assert_eq!( + read_generated_definition(&fixture.table).await, + definition_before + ); + assert_eq!( + read_raw_generated_metadata(&fixture.table).await, + raw_before + ); + assert_eq!( + fixture + .table + .generated_column_status(GEN_OUT) + .await + .unwrap(), + GeneratedColumnStatus::Complete + ); +} + +#[tokio::test] +async fn empty_standard_merge_insert_rejects_when_generated_column_present() { + let fixture = create_table_with_complete_literal_generated("b4e_empty_reject").await; + let version_before = fixture.table.version().await.unwrap(); + let rows_before = ordinary_rows(&fixture.table).await; + let raw_before = read_raw_generated_metadata(&fixture.table).await; + + let next_calls = Arc::new(AtomicUsize::new(0)); + let reader = ObservableReader::wrap(empty_reader(), next_calls.clone()); + + let mut builder = fixture.table.merge_insert(&[ID]); + configure_standard_merge(&mut builder); + let err = builder + .execute(reader) + .await + .expect_err("empty merge_insert must still reject on generated-column tables"); + assert_not_supported(&err, "empty standard merge_insert generated reject"); + assert_eq!(next_calls.load(Ordering::SeqCst), 0); + + assert_eq!(fixture.table.version().await.unwrap(), version_before); + assert_eq!(ordinary_rows(&fixture.table).await, rows_before); + assert_eq!( + read_raw_generated_metadata(&fixture.table).await, + raw_before + ); +} + +#[tokio::test] +async fn forced_lsm_without_spec_rejects_generated_before_missing_spec_and_input() { + let fixture = create_table_with_complete_literal_generated("b4e_lsm_force_reject").await; + let version_before = fixture.table.version().await.unwrap(); + let rows_before = ordinary_rows(&fixture.table).await; + let raw_before = read_raw_generated_metadata(&fixture.table).await; + + let next_calls = Arc::new(AtomicUsize::new(0)); + let reader = ObservableReader::wrap( + boxed_reader(source_batch(&[1], &["must-not-land"], &[11])), + next_calls.clone(), + ); + + let mut builder = fixture.table.merge_insert(&[ID]); + builder + .when_matched_update_all(None) + .when_not_matched_insert_all() + .use_lsm(true); + let err = builder + .execute(reader) + .await + .expect_err("generated-column guard must run before LSM missing-spec validation"); + assert_not_supported(&err, "forced LSM generated reject"); + assert_eq!(next_calls.load(Ordering::SeqCst), 0); + + assert_eq!(fixture.table.version().await.unwrap(), version_before); + assert_eq!(ordinary_rows(&fixture.table).await, rows_before); + assert_eq!( + read_raw_generated_metadata(&fixture.table).await, + raw_before + ); +} + +#[tokio::test] +async fn malformed_generated_metadata_rejects_merge_insert_before_mutation_and_redacts() { + let fixture = create_ordinary_table("b4e_malformed_preflight").await; + let snapshot = fixture + .table + .generated_column_binding_snapshot() + .await + .unwrap(); + let field_id = snapshot.field(GEN_OUT).expect(GEN_OUT).field_id(); + let planted_raw = format!( + r#"{{"format_version":1,"output_field_id":{field_id},"function_call":{{"function_id":"{FN_ID}","marker":"{MALFORMED_MARKER}"}},"dependency_epoch":1,"materialized_epoch":1}}"# + ); + assert!(planted_raw.contains(MALFORMED_MARKER)); + assert!(planted_raw.contains(FN_ID)); + fixture + .table + .update_field_metadata(&[FieldMetadataUpdate::new(GEN_OUT) + .set(GENERATED_COLUMN_METADATA_KEY, planted_raw.clone())]) + .await + .unwrap(); + assert_eq!( + read_raw_generated_metadata(&fixture.table).await, + planted_raw, + "planted malformed raw metadata must round-trip byte-for-byte" + ); + + let version_before = fixture.table.version().await.unwrap(); + let rows_before = ordinary_rows(&fixture.table).await; + let next_calls = Arc::new(AtomicUsize::new(0)); + let reader = ObservableReader::wrap( + boxed_reader(source_batch(&[1], &["must-not-land"], &[99])), + next_calls.clone(), + ); + + let mut builder = fixture.table.merge_insert(&[ID]); + configure_standard_merge(&mut builder); + let err = builder + .execute(reader) + .await + .expect_err("malformed generated metadata must fail closed before merge_insert"); + assert_invalid_input_redacted(&err, &planted_raw, "malformed merge_insert preflight"); + assert_eq!(next_calls.load(Ordering::SeqCst), 0); + + let fresh = ConnectBuilder::new(&fixture.uri) + .execute() + .await + .unwrap() + .open_table(&fixture.table_name) + .execute() + .await + .unwrap(); + assert_eq!(fresh.version().await.unwrap(), version_before); + assert_eq!(ordinary_rows(&fresh).await, rows_before); + assert_eq!(read_raw_generated_metadata(&fresh).await, planted_raw); +} + +#[tokio::test] +async fn ordinary_table_standard_merge_insert_preserves_result_semantics() { + let fixture = create_ordinary_table("b4e_ordinary_standard").await; + let mut builder = fixture.table.merge_insert(&[ID]); + configure_standard_merge(&mut builder); + let result = builder + .execute(boxed_reader(source_batch( + &[1, 3], + &["updated", "inserted"], + &[11, 30], + ))) + .await + .expect("ordinary-table standard merge_insert must succeed"); + + assert_eq!(result.num_inserted_rows, 1); + assert_eq!(result.num_updated_rows, 1); + assert_eq!(result.num_deleted_rows, 1); + assert_eq!(result.num_attempts, 1); + assert_eq!(result.num_rows, 2); + assert!(result.version > 0); + + assert_eq!( + ordinary_rows(&fixture.table).await, + vec![(1, "updated".to_string()), (3, "inserted".to_string()),] + ); +} + +#[tokio::test] +async fn ordinary_table_forced_lsm_without_spec_keeps_missing_spec_error() { + let fixture = create_ordinary_table("b4e_ordinary_lsm_missing_spec").await; + let version_before = fixture.table.version().await.unwrap(); + let rows_before = ordinary_rows(&fixture.table).await; + + let mut builder = fixture.table.merge_insert(&[ID]); + builder + .when_matched_update_all(None) + .when_not_matched_insert_all() + .use_lsm(true); + let err = builder + .execute(boxed_reader(source_batch(&[1], &["x"], &[1]))) + .await + .expect_err("ordinary table without MemWAL spec must keep missing-spec InvalidInput"); + match err { + Error::InvalidInput { message } => { + assert!( + message.contains("no MemWAL write spec"), + "expected missing-spec message, got {message}" + ); + } + other => panic!("expected InvalidInput missing-spec, got {other:?}"), + } + + assert_eq!(fixture.table.version().await.unwrap(), version_before); + assert_eq!(ordinary_rows(&fixture.table).await, rows_before); +}