diff --git a/rust/lancedb/Cargo.toml b/rust/lancedb/Cargo.toml index 5f35d5ec0..0939a41eb 100644 --- a/rust/lancedb/Cargo.toml +++ b/rust/lancedb/Cargo.toml @@ -44,7 +44,6 @@ lance-io = { workspace = true } lance-index = { workspace = true, features = ["tokenizer-jieba", "tokenizer-lindera"] } lance-table = { workspace = true } lance-linalg = { workspace = true } -lance-testing = { workspace = true } lance-encoding = { workspace = true } lance-arrow = { workspace = true } lance-namespace = { workspace = true } @@ -95,6 +94,7 @@ semver = { workspace = true } [dev-dependencies] anyhow = "1" +lance-testing = { workspace = true } tempfile = "3.5.0" random_word = { version = "0.4.3", features = ["en"] } tokio = { version = "1.23", features = ["io-util", "macros", "net", "rt-multi-thread", "sync"] } diff --git a/rust/lancedb/src/embeddings.rs b/rust/lancedb/src/embeddings.rs index ac665d76b..58e037bdd 100644 --- a/rust/lancedb/src/embeddings.rs +++ b/rust/lancedb/src/embeddings.rs @@ -198,28 +198,36 @@ fn compute_embedding_arrays( batch: &RecordBatch, embeddings: &[(EmbeddingDefinition, Arc)], ) -> Result>> { - if embeddings.len() == 1 { - let (fld, func) = &embeddings[0]; - let src_column = - batch - .column_by_name(&fld.source_column) - .ok_or_else(|| Error::InvalidInput { - message: format!("Source column '{}' not found", fld.source_column), - })?; + let input_columns = embeddings + .iter() + .map(|(fld, func)| { + let src_column = + batch + .column_by_name(&fld.source_column) + .ok_or_else(|| Error::InvalidInput { + message: format!("Source column '{}' not found", fld.source_column), + })?; + Ok((src_column.clone(), func)) + }) + .collect::>>()?; + + if batch.num_rows() == 0 { + return input_columns + .iter() + .map(|(_, func)| Ok(arrow_array::new_empty_array(func.dest_type()?.as_ref()))) + .collect(); + } + + if input_columns.len() == 1 { + let (src_column, func) = &input_columns[0]; return Ok(vec![func.compute_source_embeddings(src_column.clone())?]); } // Parallel path: multiple embeddings std::thread::scope(|s| { - let handles: Vec<_> = embeddings + let handles: Vec<_> = input_columns .iter() - .map(|(fld, func)| { - let src_column = batch.column_by_name(&fld.source_column).ok_or_else(|| { - Error::InvalidInput { - message: format!("Source column '{}' not found", fld.source_column), - } - })?; - + .map(|(src_column, func)| { let handle = s.spawn(move || func.compute_source_embeddings(src_column.clone())); Ok(handle) @@ -392,3 +400,104 @@ impl RecordBatchReader for WithEmbeddings { .into_rich_schema() } } + +#[cfg(test)] +mod tests { + use std::sync::{ + Arc, + atomic::{AtomicUsize, Ordering}, + }; + + use arrow_array::{Array, ArrayRef, FixedSizeListArray, RecordBatch, StringArray}; + use arrow_schema::DataType; + + use super::*; + + #[derive(Debug)] + struct FailingEmbedding { + calls: AtomicUsize, + } + + impl EmbeddingFunction for FailingEmbedding { + fn name(&self) -> &str { + "failing" + } + + fn source_type(&self) -> Result> { + Ok(Cow::Owned(DataType::Utf8)) + } + + fn dest_type(&self) -> Result> { + Ok(Cow::Owned(DataType::new_fixed_size_list( + DataType::Float32, + 3, + false, + ))) + } + + fn compute_source_embeddings(&self, _source: Arc) -> Result> { + self.calls.fetch_add(1, Ordering::SeqCst); + Err(Error::Runtime { + message: "embedding function must not receive an empty batch".to_string(), + }) + } + + fn compute_query_embeddings(&self, _input: Arc) -> Result> { + unreachable!("query embeddings are not exercised by this test") + } + } + + #[test] + fn empty_batch_skips_embedding_functions() { + let embedding_function = Arc::new(FailingEmbedding { + calls: AtomicUsize::new(0), + }); + let source: ArrayRef = Arc::new(StringArray::from(Vec::<&str>::new())); + let batch = RecordBatch::try_from_iter([("text", source)]).unwrap(); + let embeddings = vec![( + EmbeddingDefinition::new("text", "failing", Some("text_embedding")), + embedding_function.clone() as Arc, + )]; + + let result = compute_embeddings_for_batch(batch, &embeddings).unwrap(); + + assert_eq!(embedding_function.calls.load(Ordering::SeqCst), 0); + assert_eq!(result.num_rows(), 0); + + let embedding = result.column_by_name("text_embedding").unwrap(); + assert_eq!( + embedding.data_type(), + &DataType::new_fixed_size_list(DataType::Float32, 3, false) + ); + assert_eq!(embedding.null_count(), 0); + + let embedding = embedding + .as_any() + .downcast_ref::() + .unwrap(); + assert_eq!(embedding.len(), 0); + assert_eq!(embedding.value_length(), 3); + assert_eq!(embedding.values().len(), 0); + } + + #[test] + fn empty_batch_still_validates_source_column() { + let embedding_function = Arc::new(FailingEmbedding { + calls: AtomicUsize::new(0), + }); + let source: ArrayRef = Arc::new(StringArray::from(Vec::<&str>::new())); + let batch = RecordBatch::try_from_iter([("text", source)]).unwrap(); + let embeddings = vec![( + EmbeddingDefinition::new("missing_column", "failing", Some("text_embedding")), + embedding_function.clone() as Arc, + )]; + + let result = compute_embeddings_for_batch(batch, &embeddings); + assert!(result.is_err()); + assert!( + matches!(result.unwrap_err(), Error::InvalidInput { .. }), + "expected InvalidInput error when source column is missing" + ); + assert_eq!(embedding_function.calls.load(Ordering::SeqCst), 0); + } +}