From 7a46da2e67fd041ff68abe220a8baf5403cd9a1a Mon Sep 17 00:00:00 2001 From: Xuanwo Date: Wed, 12 Aug 2026 18:52:21 +0800 Subject: [PATCH] feat: guard generated metadata on add columns --- rust/lancedb/src/function/schema_admission.rs | 13 +- rust/lancedb/src/table.rs | 2 + ...erated_column_schema_admission_contract.rs | 441 ++++++++++++++++++ rust/lancedb/src/table/schema_evolution.rs | 29 ++ 4 files changed, 479 insertions(+), 6 deletions(-) create mode 100644 rust/lancedb/src/table/add_columns_generated_column_schema_admission_contract.rs diff --git a/rust/lancedb/src/function/schema_admission.rs b/rust/lancedb/src/function/schema_admission.rs index 90087194d..bdff2242f 100644 --- a/rust/lancedb/src/function/schema_admission.rs +++ b/rust/lancedb/src/function/schema_admission.rs @@ -3,10 +3,11 @@ //! Schema admission for caller-authored generated-column definition ingress. //! -//! General-purpose table-schema inputs (for example `Database::create_table`) -//! must not invent or mutate Job-owned `lancedb::generated_column` top-level -//! field metadata. Only generated-column create/change/refresh Job publication -//! may create or change that reserved key. +//! General-purpose table-schema inputs (for example `Database::create_table` +//! and Native `add_columns` schema-bearing transforms) must not invent or +//! mutate Job-owned `lancedb::generated_column` top-level field metadata. Only +//! generated-column create/change/refresh Job publication may create or change +//! that reserved key. //! //! This helper checks raw key presence on top-level fields only. It does not //! recurse into nested children, inspect schema-level metadata, decode the @@ -20,8 +21,8 @@ use crate::{Error, Result}; /// Reject a caller-authored Arrow schema that carries reserved generated-column /// definition metadata on any top-level field. /// -/// Safe to call at the start of create-table (and later overwrite / add-columns) -/// paths before source consumption, namespace mutation, or HTTP. +/// Safe to call at the start of create-table and Native add-columns paths +/// before source consumption, namespace mutation, or HTTP. pub fn reject_caller_authored_generated_column_schema(schema: &Schema) -> Result<()> { for field in schema.fields() { if field.metadata().contains_key(GENERATED_COLUMN_METADATA_KEY) { diff --git a/rust/lancedb/src/table.rs b/rust/lancedb/src/table.rs index 0321513a1..6b4398654 100644 --- a/rust/lancedb/src/table.rs +++ b/rust/lancedb/src/table.rs @@ -86,6 +86,8 @@ pub mod schema_evolution; pub mod update; pub mod write_progress; +#[cfg(test)] +mod add_columns_generated_column_schema_admission_contract; #[cfg(test)] mod append_generated_column_invalidation_contract; #[cfg(test)] diff --git a/rust/lancedb/src/table/add_columns_generated_column_schema_admission_contract.rs b/rust/lancedb/src/table/add_columns_generated_column_schema_admission_contract.rs new file mode 100644 index 000000000..c64879817 --- /dev/null +++ b/rust/lancedb/src/table/add_columns_generated_column_schema_admission_contract.rs @@ -0,0 +1,441 @@ +// SPDX-License-Identifier: Apache-2.0 +// SPDX-FileCopyrightText: Copyright The LanceDB Authors + +//! RED runtime contract tests for Native add-columns schema admission (B4h). +//! +//! Caller-authored Arrow field metadata under +//! [`crate::function::GENERATED_COLUMN_METADATA_KEY`] must not enter table +//! schema state through general-purpose Native `add_columns`. Generated +//! definitions are Job-owned. Schema-bearing transforms (`BatchUDF`, `Stream`, +//! `Reader`, `AllNulls`) currently accept and persist reserved top-level field +//! metadata; these tests pin the missing pre-consumption admission guard. + +use std::collections::HashMap; +use std::sync::Arc; +use std::sync::atomic::{AtomicUsize, Ordering}; + +use arrow_array::{Int32Array, RecordBatch, RecordBatchIterator, RecordBatchReader, StringArray}; +use arrow_schema::{ArrowError, DataType, Field, Schema, SchemaRef}; +use datafusion_physical_plan::stream::RecordBatchStreamAdapter; +use futures::{TryStreamExt, stream}; +use lance::dataset::{BatchUDF, NewColumnTransform}; +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, +}; +use crate::query::{ExecutableQuery, QueryBase, Select}; +use crate::table::Table; + +const ID: &str = "id"; +const ORDINARY: &str = "ordinary"; +const GEN_OUT: &str = "gen_out"; +const ORDINARY_META_KEY: &str = "unit"; +const ORDINARY_META_VALUE: &str = "label"; +const FN_ID: &str = "fn.exact.b4h.add_columns.literal"; +const MALFORMED_MARKER: &str = "SENSITIVE_B4H_ADD_COLUMNS_METADATA_MARKER_4f8a_c3e2"; + +struct Fixture { + _tmp: TempDir, + table: Table, +} + +/// Counts [`RecordBatchReader::next`] calls. [`RecordBatchReader::schema`] is free. +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_definition(output_field_id: i32) -> GeneratedColumnDefinition { + let function = Function::new( + FunctionId::try_new(FN_ID).unwrap(), + FunctionSignature::try_new( + vec![FunctionParameter::new("label", DataType::Utf8)], + FunctionOutput::new(DataType::Int32, true), + ) + .unwrap(), + ); + 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, 1, 1).unwrap() +} + +fn valid_reserved_payload() -> String { + literal_definition(1).to_metadata_json().unwrap() +} + +fn malformed_reserved_payload() -> String { + format!( + r#"{{"format_version":1,"output_field_id":1,"function_call":"{MALFORMED_MARKER}","dependency_epoch":1,"materialized_epoch":1}}"# + ) +} + +fn seed_batch() -> RecordBatch { + let schema = Arc::new(Schema::new(vec![ + Field::new(ID, DataType::Int32, false), + Field::new(ORDINARY, DataType::Utf8, true), + ])); + RecordBatch::try_new( + schema, + vec![ + Arc::new(Int32Array::from(vec![1, 2])), + Arc::new(StringArray::from(vec![Some("a"), Some("b")])), + ], + ) + .unwrap() +} + +fn field_with_metadata(metadata: HashMap) -> Field { + Field::new(GEN_OUT, DataType::Int32, true).with_metadata(metadata) +} + +fn reserved_field(payload: &str) -> Field { + field_with_metadata( + [( + GENERATED_COLUMN_METADATA_KEY.to_string(), + payload.to_string(), + )] + .into(), + ) +} + +fn ordinary_metadata_field() -> Field { + field_with_metadata( + [( + ORDINARY_META_KEY.to_string(), + ORDINARY_META_VALUE.to_string(), + )] + .into(), + ) +} + +fn values_batch(schema: SchemaRef, values: Vec) -> RecordBatch { + RecordBatch::try_new(schema, vec![Arc::new(Int32Array::from(values))]).unwrap() +} + +fn boxed_reader(batch: RecordBatch) -> Box { + let schema = batch.schema(); + Box::new(RecordBatchIterator::new( + vec![Ok(batch)].into_iter(), + schema, + )) +} + +fn observable_stream( + batch: RecordBatch, + yield_calls: Arc, +) -> datafusion_physical_plan::SendableRecordBatchStream { + let schema = batch.schema(); + let counter = yield_calls.clone(); + Box::pin(RecordBatchStreamAdapter::new( + schema, + stream::once(async move { + counter.fetch_add(1, Ordering::SeqCst); + Ok(batch) + }), + )) +} + +fn assert_not_supported_redacted(err: &Error, label: &str, payload: &str) { + match err { + Error::NotSupported { message } => { + let rendered = format!("{err}\n{err:?}\n{message}"); + assert!( + !rendered.contains(GENERATED_COLUMN_METADATA_KEY), + "{label}: leaked metadata wire key: {rendered}" + ); + assert!( + !rendered.contains(payload), + "{label}: leaked raw payload: {rendered}" + ); + assert!( + !rendered.contains(FN_ID), + "{label}: leaked Function ID: {rendered}" + ); + assert!( + !rendered.contains(GEN_OUT), + "{label}: leaked output field name: {rendered}" + ); + assert!( + !rendered.contains(MALFORMED_MARKER), + "{label}: leaked malformed marker: {rendered}" + ); + assert!( + message.to_lowercase().contains("generated") + || message.to_lowercase().contains("job"), + "{label}: message must describe Job-owned generated-column boundary: {message}" + ); + } + other => panic!("{label}: expected Error::NotSupported, got {other:?}"), + } +} + +async fn create_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 } +} + +async fn snapshot_rows(table: &Table) -> Vec<(i32, String)> { + let batches: Vec = 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 i in 0..batch.num_rows() { + rows.push((ids.value(i), ordinary.value(i).to_string())); + } + } + rows.sort_by_key(|(id, _)| *id); + rows +} + +async fn assert_table_unchanged( + table: &Table, + version_before: u64, + schema_before: &Schema, + rows_before: &[(i32, String)], +) { + assert_eq!(table.version().await.unwrap(), version_before); + let schema_after = table.schema().await.unwrap(); + assert_eq!(schema_after.as_ref(), schema_before); + assert!( + schema_after.field_with_name(GEN_OUT).is_err(), + "rejected add_columns must leave column `{GEN_OUT}` absent" + ); + assert_eq!(snapshot_rows(table).await, rows_before); +} + +#[tokio::test] +async fn batch_udf_rejects_valid_reserved_before_mapper() { + let fixture = create_table("b4h_batch_udf").await; + let table = &fixture.table; + let version_before = table.version().await.unwrap(); + let schema_before = table.schema().await.unwrap(); + let rows_before = snapshot_rows(table).await; + + let payload = valid_reserved_payload(); + let output_schema = Arc::new(Schema::new(vec![reserved_field(&payload)])); + let mapper_schema = output_schema.clone(); + let mapper_calls = Arc::new(AtomicUsize::new(0)); + let calls = mapper_calls.clone(); + let udf = BatchUDF { + mapper: Box::new(move |batch: &RecordBatch| { + calls.fetch_add(1, Ordering::SeqCst); + let values = Int32Array::from(vec![Some(10); batch.num_rows()]); + Ok(RecordBatch::try_new( + mapper_schema.clone(), + vec![Arc::new(values)], + )?) + }), + output_schema, + result_checkpoint: None, + }; + + // Public Table::add_columns builder path. + let err = table + .add_columns() + .transform(NewColumnTransform::BatchUDF(udf)) + .execute() + .await + .expect_err("BatchUDF must reject reserved generated-column metadata"); + assert_not_supported_redacted(&err, "BatchUDF reserved admission", &payload); + assert_eq!( + mapper_calls.load(Ordering::SeqCst), + 0, + "rejection must occur before invoking the BatchUDF mapper" + ); + assert_table_unchanged(table, version_before, schema_before.as_ref(), &rows_before).await; +} + +#[tokio::test] +async fn stream_rejects_malformed_reserved_before_yield() { + let fixture = create_table("b4h_stream").await; + let table = &fixture.table; + let version_before = table.version().await.unwrap(); + let schema_before = table.schema().await.unwrap(); + let rows_before = snapshot_rows(table).await; + + let payload = malformed_reserved_payload(); + assert!(payload.contains(MALFORMED_MARKER)); + let output_schema = Arc::new(Schema::new(vec![reserved_field(&payload)])); + let yield_calls = Arc::new(AtomicUsize::new(0)); + let stream = observable_stream( + values_batch(output_schema, vec![10, 20]), + yield_calls.clone(), + ); + + // Direct experimental BaseTable::add_columns path. + let err = table + .base_table() + .add_columns(NewColumnTransform::Stream(stream), None) + .await + .expect_err("Stream must reject reserved generated-column metadata"); + assert_not_supported_redacted(&err, "Stream reserved admission", &payload); + assert_eq!( + yield_calls.load(Ordering::SeqCst), + 0, + "rejection must occur before polling/yielding the user Stream" + ); + assert_table_unchanged(table, version_before, schema_before.as_ref(), &rows_before).await; +} + +#[tokio::test] +async fn reader_rejects_valid_reserved_before_next() { + let fixture = create_table("b4h_reader").await; + let table = &fixture.table; + let version_before = table.version().await.unwrap(); + let schema_before = table.schema().await.unwrap(); + let rows_before = snapshot_rows(table).await; + + let payload = valid_reserved_payload(); + let output_schema = Arc::new(Schema::new(vec![reserved_field(&payload)])); + let next_calls = Arc::new(AtomicUsize::new(0)); + let reader = ObservableReader::wrap( + boxed_reader(values_batch(output_schema, vec![10, 20])), + next_calls.clone(), + ); + + let err = table + .base_table() + .add_columns(NewColumnTransform::Reader(reader), None) + .await + .expect_err("Reader must reject reserved generated-column metadata"); + assert_not_supported_redacted(&err, "Reader reserved admission", &payload); + assert_eq!( + next_calls.load(Ordering::SeqCst), + 0, + "rejection must occur before RecordBatchReader::next" + ); + assert_table_unchanged(table, version_before, schema_before.as_ref(), &rows_before).await; +} + +#[tokio::test] +async fn all_nulls_rejects_malformed_reserved_before_commit() { + let fixture = create_table("b4h_all_nulls").await; + let table = &fixture.table; + let version_before = table.version().await.unwrap(); + let schema_before = table.schema().await.unwrap(); + let rows_before = snapshot_rows(table).await; + + let payload = malformed_reserved_payload(); + let output_schema = Arc::new(Schema::new(vec![reserved_field(&payload)])); + + let err = table + .add_columns() + .transform(NewColumnTransform::AllNulls(output_schema)) + .execute() + .await + .expect_err("AllNulls must reject reserved generated-column metadata"); + assert_not_supported_redacted(&err, "AllNulls reserved admission", &payload); + assert_table_unchanged(table, version_before, schema_before.as_ref(), &rows_before).await; +} + +#[tokio::test] +async fn sql_expressions_add_columns_still_succeeds() { + let fixture = create_table("b4h_sql_control").await; + let table = &fixture.table; + + table + .add_columns() + .transform(NewColumnTransform::SqlExpressions(vec![( + "doubled".into(), + "id * 2".into(), + )])) + .execute() + .await + .expect("ordinary SqlExpressions add_columns must remain supported"); + + let schema = table.schema().await.unwrap(); + assert!(schema.field_with_name("doubled").is_ok()); + assert!(schema.field_with_name(GEN_OUT).is_err()); + assert!( + !schema + .field_with_name("doubled") + .unwrap() + .metadata() + .contains_key(GENERATED_COLUMN_METADATA_KEY) + ); +} + +#[tokio::test] +async fn schema_bearing_ordinary_metadata_is_preserved() { + let fixture = create_table("b4h_ordinary_meta").await; + let table = &fixture.table; + let output_schema = Arc::new(Schema::new(vec![ordinary_metadata_field()])); + + // AllNulls is schema-bearing and metadata-only; proves ordinary metadata + // remains accepted so a later guard cannot reject every field metadata map. + table + .add_columns() + .transform(NewColumnTransform::AllNulls(output_schema)) + .execute() + .await + .expect("ordinary non-reserved field metadata must remain accepted"); + + let schema = table.schema().await.unwrap(); + let md = schema.field_with_name(GEN_OUT).unwrap().metadata(); + assert_eq!( + md.get(ORDINARY_META_KEY).map(String::as_str), + Some(ORDINARY_META_VALUE) + ); + assert!(!md.contains_key(GENERATED_COLUMN_METADATA_KEY)); +} diff --git a/rust/lancedb/src/table/schema_evolution.rs b/rust/lancedb/src/table/schema_evolution.rs index 322cc5724..c01259844 100644 --- a/rust/lancedb/src/table/schema_evolution.rs +++ b/rust/lancedb/src/table/schema_evolution.rs @@ -8,14 +8,42 @@ //! - [`alter_columns`](execute_alter_columns): Rename columns, change types, or modify nullability //! - [`drop_columns`](execute_drop_columns): Remove columns from the table +use arrow_array::RecordBatchReader; use lance::dataset::{ColumnAlteration, NewColumnTransform}; use serde::{Deserialize, Serialize}; use std::collections::HashMap; use super::NativeTable; use crate::function::GENERATED_COLUMN_METADATA_KEY; +use crate::function::schema_admission::reject_caller_authored_generated_column_schema; use crate::{Error, Result}; +/// Reject caller-authored schema-bearing `add_columns` transforms that carry +/// reserved generated-column top-level field metadata. +/// +/// Borrows without consuming the transform: Stream is not polled, Reader is +/// not iterated, and BatchUDF mapper is not invoked. `SqlExpressions` cannot +/// carry an Arrow output schema and is accepted. +pub(crate) fn reject_caller_authored_generated_column_add_columns_transform( + transforms: &NewColumnTransform, +) -> Result<()> { + match transforms { + NewColumnTransform::BatchUDF(udf) => { + reject_caller_authored_generated_column_schema(udf.output_schema.as_ref()) + } + NewColumnTransform::Stream(stream) => { + reject_caller_authored_generated_column_schema(stream.schema().as_ref()) + } + NewColumnTransform::Reader(reader) => { + reject_caller_authored_generated_column_schema(reader.schema().as_ref()) + } + NewColumnTransform::AllNulls(schema) => { + reject_caller_authored_generated_column_schema(schema.as_ref()) + } + NewColumnTransform::SqlExpressions(_) => Ok(()), + } +} + /// Shared rejection for general-purpose field-metadata updates that name the /// reserved generated-column definition key. /// @@ -153,6 +181,7 @@ pub(crate) async fn execute_add_columns( transforms: NewColumnTransform, read_columns: Option>, ) -> Result { + reject_caller_authored_generated_column_add_columns_transform(&transforms)?; table.dataset.ensure_mutable()?; let mut dataset = (*table.dataset.get().await?).clone(); dataset.add_columns(transforms, read_columns, None).await?;