diff --git a/Cargo.lock b/Cargo.lock index ec6b7cbfb..627915c47 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -5432,6 +5432,7 @@ dependencies = [ "aws-sdk-kms", "aws-sdk-s3", "aws-smithy-runtime", + "base64 0.22.1", "bytes", "candle-core", "candle-nn", diff --git a/rust/lancedb/Cargo.toml b/rust/lancedb/Cargo.toml index e33b86b12..f1ac0eb52 100644 --- a/rust/lancedb/Cargo.toml +++ b/rust/lancedb/Cargo.toml @@ -12,6 +12,7 @@ rust-version.workspace = true # See more keys and their definitions at https://doc.rust-lang.org/cargo/reference/manifest.html [dependencies] ahash = { workspace = true } +base64 = "0.22" arrow = { workspace = true } arrow-array = { workspace = true } arrow-buffer = { workspace = true } diff --git a/rust/lancedb/src/function.rs b/rust/lancedb/src/function.rs new file mode 100644 index 000000000..c11bbc9f9 --- /dev/null +++ b/rust/lancedb/src/function.rs @@ -0,0 +1,933 @@ +// SPDX-License-Identifier: Apache-2.0 +// SPDX-FileCopyrightText: Copyright The LanceDB Authors + +//! Immutable first-class Function and generated-column value model (B1a). +//! +//! This module defines transport and metadata value types only. It does not +//! provide catalogs, job execution, query planning, or generated-column runtime. + +use std::collections::HashSet; +use std::io::Cursor; +use std::sync::Arc; + +use arrow_array::{ArrayRef, RecordBatch}; +use arrow_ipc::reader::FileReader; +use arrow_schema::{DataType, Field, Schema}; +use base64::Engine; +use base64::engine::general_purpose::STANDARD as BASE64; +use serde::de::Error as DeError; +use serde::{Deserialize, Deserializer, Serialize, Serializer}; +use serde_json::Value; + +use crate::ipc::{batches_to_ipc_file, schema_to_ipc_file}; +use crate::{Error, Result}; + +/// Field metadata key used to store a [`GeneratedColumnDefinition`] JSON document. +pub const GENERATED_COLUMN_METADATA_KEY: &str = "lancedb::generated_column"; + +const FORMAT_VERSION_V1: u32 = 1; +const TYPE_IPC_FIELD_NAME: &str = ""; +const LITERAL_IPC_FIELD_NAME: &str = "value"; + +fn invalid_input(message: impl Into) -> Error { + Error::InvalidInput { + message: message.into(), + } +} + +fn encode_type_ipc(data_type: &DataType) -> Result> { + let schema = Schema::new(vec![Field::new( + TYPE_IPC_FIELD_NAME, + data_type.clone(), + true, + )]); + schema_to_ipc_file(&schema) +} + +fn decode_type_ipc(bytes: &[u8]) -> Result { + let reader = FileReader::try_new(Cursor::new(bytes), None) + .map_err(|e| invalid_input(format!("invalid Arrow IPC for data type: {e}")))?; + let schema = reader.schema(); + if schema.fields().len() != 1 { + return Err(invalid_input( + "type data_type_ipc must be a schema-only Arrow IPC file with exactly one field", + )); + } + if reader.num_batches() != 0 { + return Err(invalid_input( + "type data_type_ipc must be schema-only Arrow IPC (no record batches)", + )); + } + let data_type = schema.field(0).data_type().clone(); + let canonical = encode_type_ipc(&data_type)?; + if canonical.as_slice() != bytes { + return Err(invalid_input( + "type data_type_ipc must be canonical schema-only Arrow IPC with no trailing bytes", + )); + } + Ok(data_type) +} + +fn encode_type_ipc_b64(data_type: &DataType) -> Result { + Ok(BASE64.encode(encode_type_ipc(data_type)?)) +} + +fn decode_type_ipc_b64(encoded: &str) -> Result { + let bytes = BASE64 + .decode(encoded.as_bytes()) + .map_err(|e| invalid_input(format!("invalid base64 data_type_ipc: {e}")))?; + decode_type_ipc(&bytes) +} + +fn encode_literal_ipc(array: &ArrayRef) -> Result> { + if array.len() != 1 { + return Err(invalid_input( + "literal argument must contain exactly one row", + )); + } + let schema = Arc::new(Schema::new(vec![Field::new( + LITERAL_IPC_FIELD_NAME, + array.data_type().clone(), + true, + )])); + let batch = RecordBatch::try_new(schema, vec![array.clone()])?; + batches_to_ipc_file(&[batch]) +} + +fn decode_literal_ipc(bytes: &[u8]) -> Result { + let mut reader = FileReader::try_new(Cursor::new(bytes), None) + .map_err(|e| invalid_input(format!("invalid Arrow IPC for literal: {e}")))?; + let schema = reader.schema(); + if schema.fields().len() != 1 { + return Err(invalid_input("literal ipc must contain exactly one column")); + } + if reader.num_batches() != 1 { + return Err(invalid_input( + "literal ipc must contain exactly one record batch", + )); + } + let batch = reader + .next() + .ok_or_else(|| invalid_input("literal ipc must contain exactly one record batch"))? + .map_err(|e| invalid_input(format!("invalid Arrow IPC for literal batch: {e}")))?; + if batch.num_columns() != 1 { + return Err(invalid_input("literal ipc must contain exactly one column")); + } + if batch.num_rows() != 1 { + return Err(invalid_input("literal ipc must contain exactly one row")); + } + let array = batch.column(0).clone(); + let canonical = encode_literal_ipc(&array)?; + if canonical.as_slice() != bytes { + return Err(invalid_input( + "literal ipc must be canonical one-batch one-field one-row Arrow IPC with no trailing bytes", + )); + } + Ok(array) +} + +fn encode_literal_ipc_b64(array: &ArrayRef) -> Result { + Ok(BASE64.encode(encode_literal_ipc(array)?)) +} + +fn decode_literal_ipc_b64(encoded: &str) -> Result { + let bytes = BASE64 + .decode(encoded.as_bytes()) + .map_err(|e| invalid_input(format!("invalid base64 literal ipc: {e}")))?; + decode_literal_ipc(&bytes) +} + +/// Opaque, non-empty identifier for a [`Function`]. +#[derive(Debug, Clone, PartialEq, Eq, Hash)] +pub struct FunctionId { + value: String, +} + +impl FunctionId { + /// Create a new opaque function id. + /// + /// Returns an error when `value` is empty. + pub fn try_new(value: impl Into) -> Result { + let value = value.into(); + if value.is_empty() { + return Err(invalid_input("FunctionId must be non-empty")); + } + Ok(Self { value }) + } + + /// Borrow the opaque id string. + pub fn as_str(&self) -> &str { + &self.value + } +} + +/// A named typed parameter in a [`FunctionSignature`]. +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct FunctionParameter { + name: String, + data_type: DataType, +} + +impl FunctionParameter { + /// Create a parameter with the given name and Arrow data type. + pub fn new(name: impl Into, data_type: DataType) -> Self { + Self { + name: name.into(), + data_type, + } + } + + /// Parameter name. + pub fn name(&self) -> &str { + &self.name + } + + /// Parameter Arrow data type. + pub fn data_type(&self) -> &DataType { + &self.data_type + } +} + +/// Output type description for a [`FunctionSignature`]. +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct FunctionOutput { + data_type: DataType, + nullable: bool, +} + +impl FunctionOutput { + /// Create an output type with Arrow data type and nullability. + pub fn new(data_type: DataType, nullable: bool) -> Self { + Self { + data_type, + nullable, + } + } + + /// Output Arrow data type. + pub fn data_type(&self) -> &DataType { + &self.data_type + } + + /// Whether the output may be null. + pub fn nullable(&self) -> bool { + self.nullable + } +} + +/// Ordered parameter list plus a single output type for a [`Function`]. +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct FunctionSignature { + parameters: Vec, + output: FunctionOutput, +} + +impl FunctionSignature { + /// Create a signature with unique non-empty parameter names. + pub fn try_new(parameters: Vec, output: FunctionOutput) -> Result { + let mut seen = HashSet::with_capacity(parameters.len()); + for parameter in ¶meters { + if parameter.name.is_empty() { + return Err(invalid_input( + "FunctionSignature parameter names must be non-empty", + )); + } + if !seen.insert(parameter.name.as_str()) { + return Err(invalid_input(format!( + "duplicate FunctionSignature parameter name `{}`", + parameter.name + ))); + } + } + Ok(Self { parameters, output }) + } + + /// Ordered parameters. + pub fn parameters(&self) -> &[FunctionParameter] { + &self.parameters + } + + /// Output type. + pub fn output(&self) -> &FunctionOutput { + &self.output + } +} + +/// Immutable first-class function value (format version 1). +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct Function { + id: FunctionId, + signature: FunctionSignature, +} + +impl Function { + /// Create a function from an opaque id and signature. + pub fn new(id: FunctionId, signature: FunctionSignature) -> Self { + Self { id, signature } + } + + /// Opaque function id. + pub fn id(&self) -> &FunctionId { + &self.id + } + + /// Function signature. + pub fn signature(&self) -> &FunctionSignature { + &self.signature + } +} + +#[derive(Debug, Serialize, Deserialize)] +#[serde(deny_unknown_fields)] +struct FunctionWire { + format_version: u32, + id: String, + signature: SignatureWire, +} + +#[derive(Debug, Serialize, Deserialize)] +#[serde(deny_unknown_fields)] +struct SignatureWire { + parameters: Vec, + output: OutputWire, +} + +#[derive(Debug, Serialize, Deserialize)] +#[serde(deny_unknown_fields)] +struct ParameterWire { + name: String, + data_type_ipc: String, +} + +#[derive(Debug, Serialize, Deserialize)] +#[serde(deny_unknown_fields)] +struct OutputWire { + data_type_ipc: String, + nullable: bool, +} + +impl Function { + fn to_wire(&self) -> Result { + let parameters = self + .signature + .parameters + .iter() + .map(|parameter| { + Ok(ParameterWire { + name: parameter.name.clone(), + data_type_ipc: encode_type_ipc_b64(¶meter.data_type)?, + }) + }) + .collect::>>()?; + Ok(FunctionWire { + format_version: FORMAT_VERSION_V1, + id: self.id.value.clone(), + signature: SignatureWire { + parameters, + output: OutputWire { + data_type_ipc: encode_type_ipc_b64(&self.signature.output.data_type)?, + nullable: self.signature.output.nullable, + }, + }, + }) + } + + fn from_wire(wire: FunctionWire) -> Result { + if wire.format_version != FORMAT_VERSION_V1 { + return Err(invalid_input(format!( + "unsupported Function format_version {}", + wire.format_version + ))); + } + let id = FunctionId::try_new(wire.id)?; + let parameters = wire + .signature + .parameters + .into_iter() + .map(|parameter| { + Ok(FunctionParameter::new( + parameter.name, + decode_type_ipc_b64(¶meter.data_type_ipc)?, + )) + }) + .collect::>>()?; + let output = FunctionOutput::new( + decode_type_ipc_b64(&wire.signature.output.data_type_ipc)?, + wire.signature.output.nullable, + ); + let signature = FunctionSignature::try_new(parameters, output)?; + Ok(Self { id, signature }) + } +} + +impl Serialize for Function { + fn serialize(&self, serializer: S) -> std::result::Result + where + S: Serializer, + { + self.to_wire() + .map_err(serde::ser::Error::custom)? + .serialize(serializer) + } +} + +impl<'de> Deserialize<'de> for Function { + fn deserialize(deserializer: D) -> std::result::Result + where + D: Deserializer<'de>, + { + let wire = FunctionWire::deserialize(deserializer)?; + Self::from_wire(wire).map_err(D::Error::custom) + } +} + +/// A typed argument bound to a function parameter. +#[derive(Debug, Clone)] +pub struct FunctionArgument { + kind: FunctionArgumentKind, +} + +#[derive(Debug, Clone)] +enum FunctionArgumentKind { + Field { field_id: i32, data_type: DataType }, + Literal { array: ArrayRef }, +} + +impl PartialEq for FunctionArgument { + fn eq(&self, other: &Self) -> bool { + match (&self.kind, &other.kind) { + ( + FunctionArgumentKind::Field { + field_id: a_id, + data_type: a_ty, + }, + FunctionArgumentKind::Field { + field_id: b_id, + data_type: b_ty, + }, + ) => a_id == b_id && a_ty == b_ty, + ( + FunctionArgumentKind::Literal { array: a }, + FunctionArgumentKind::Literal { array: b }, + ) => a.as_ref() == b.as_ref(), + _ => false, + } + } +} + +impl Eq for FunctionArgument {} + +impl FunctionArgument { + /// Create a field argument referencing a non-negative stable field id. + pub fn try_field(field_id: i32, data_type: DataType) -> Result { + if field_id < 0 { + return Err(invalid_input( + "field argument field_id must be non-negative", + )); + } + Ok(Self { + kind: FunctionArgumentKind::Field { + field_id, + data_type, + }, + }) + } + + /// Create a literal argument from a one-row Arrow array (including typed NULL). + pub fn try_literal(array: ArrayRef) -> Result { + if array.len() != 1 { + return Err(invalid_input( + "literal argument must contain exactly one row", + )); + } + // Validate that the value has a canonical IPC encoding. + let _ = encode_literal_ipc(&array)?; + Ok(Self { + kind: FunctionArgumentKind::Literal { array }, + }) + } + + /// Stable field id when this argument is a field reference. + pub fn field_id(&self) -> Option { + match &self.kind { + FunctionArgumentKind::Field { field_id, .. } => Some(*field_id), + FunctionArgumentKind::Literal { .. } => None, + } + } + + /// Arrow data type expected or carried by this argument. + pub fn data_type(&self) -> &DataType { + match &self.kind { + FunctionArgumentKind::Field { data_type, .. } => data_type, + FunctionArgumentKind::Literal { array } => array.data_type(), + } + } + + /// One-row literal array when this argument is a literal. + pub fn literal_array(&self) -> Option<&ArrayRef> { + match &self.kind { + FunctionArgumentKind::Literal { array } => Some(array), + FunctionArgumentKind::Field { .. } => None, + } + } + + /// Whether this argument is a typed NULL literal. + pub fn is_typed_null(&self) -> bool { + match &self.kind { + FunctionArgumentKind::Literal { array } => array.is_null(0), + FunctionArgumentKind::Field { .. } => false, + } + } + + fn to_value_wire(&self) -> Result { + match &self.kind { + FunctionArgumentKind::Field { + field_id, + data_type, + } => Ok(ArgumentValueWire::Field { + field_id: *field_id, + data_type_ipc: encode_type_ipc_b64(data_type)?, + }), + FunctionArgumentKind::Literal { array } => Ok(ArgumentValueWire::Literal { + ipc: encode_literal_ipc_b64(array)?, + }), + } + } + + fn from_value_wire(wire: ArgumentValueWire) -> Result { + match wire { + ArgumentValueWire::Field { + field_id, + data_type_ipc, + } => Self::try_field(field_id, decode_type_ipc_b64(&data_type_ipc)?), + ArgumentValueWire::Literal { ipc } => Self::try_literal(decode_literal_ipc_b64(&ipc)?), + } + } +} + +/// A function id plus named argument bindings. +/// +/// [`FunctionCall::try_new`] validates against a [`Function`] and normalizes +/// bindings to signature parameter order. Structural decode via +/// [`Deserialize`] or [`GeneratedColumnDefinition::from_metadata_json`] does +/// not perform that catalog validation: decoded calls must pass +/// [`FunctionCall::validate_against`] before execution. +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct FunctionCall { + function_id: FunctionId, + arguments: Vec<(String, FunctionArgument)>, +} + +impl FunctionCall { + /// Bind arguments to `function`, normalizing to signature parameter order. + pub fn try_new(function: &Function, bindings: Vec<(String, FunctionArgument)>) -> Result { + let parameters = function.signature().parameters(); + if bindings.len() != parameters.len() { + return Err(invalid_input(format!( + "FunctionCall requires exactly {} bindings, got {}", + parameters.len(), + bindings.len() + ))); + } + + let mut seen = HashSet::with_capacity(bindings.len()); + let mut by_name = std::collections::HashMap::with_capacity(bindings.len()); + for (name, argument) in bindings { + if !seen.insert(name.clone()) { + return Err(invalid_input(format!( + "duplicate FunctionCall binding for parameter `{name}`" + ))); + } + by_name.insert(name, argument); + } + + let mut arguments = Vec::with_capacity(parameters.len()); + for parameter in parameters { + let Some(argument) = by_name.remove(parameter.name()) else { + return Err(invalid_input(format!( + "missing FunctionCall binding for parameter `{}`", + parameter.name() + ))); + }; + if argument.data_type() != parameter.data_type() { + return Err(invalid_input(format!( + "FunctionCall argument type mismatch for parameter `{}`", + parameter.name() + ))); + } + arguments.push((parameter.name().to_string(), argument)); + } + + if let Some((unknown, _)) = by_name.into_iter().next() { + return Err(invalid_input(format!( + "unknown FunctionCall parameter `{unknown}`" + ))); + } + + Ok(Self { + function_id: function.id().clone(), + arguments, + }) + } + + /// Validate this call against a catalog [`Function`] identity and signature. + /// + /// Requires exact [`FunctionId`] equality, exact argument count, each + /// parameter name in exact signature order, and exact Arrow type equality + /// for every binding. This does not check that field arguments exist or + /// match types in a table schema. + /// + /// Calls produced by [`Self::try_new`] always pass for the same + /// `function`. Structurally decoded calls (for example from + /// [`GeneratedColumnDefinition::from_metadata_json`]) must pass this check + /// before execution. + pub fn validate_against(&self, function: &Function) -> Result<()> { + if self.function_id != *function.id() { + return Err(invalid_input(format!( + "FunctionCall function_id mismatch: call has `{}`, function has `{}`", + self.function_id.as_str(), + function.id().as_str() + ))); + } + + let parameters = function.signature().parameters(); + if self.arguments.len() != parameters.len() { + return Err(invalid_input(format!( + "FunctionCall requires exactly {} bindings, got {}", + parameters.len(), + self.arguments.len() + ))); + } + + for (index, parameter) in parameters.iter().enumerate() { + let (bound_name, argument) = &self.arguments[index]; + if bound_name != parameter.name() { + return Err(invalid_input(format!( + "FunctionCall parameter name mismatch at position {index}: call has `{bound_name}`, signature has `{}`", + parameter.name() + ))); + } + if argument.data_type() != parameter.data_type() { + return Err(invalid_input(format!( + "FunctionCall argument type mismatch for parameter `{}`", + parameter.name() + ))); + } + } + + Ok(()) + } + + /// Opaque function id stored on the call. + pub fn function_id(&self) -> &FunctionId { + &self.function_id + } + + /// Bindings as `(parameter_name, argument)`. + /// + /// After [`Self::try_new`], bindings are in signature parameter order. + /// Structurally decoded calls may not be ordered or complete until + /// [`Self::validate_against`] succeeds. + pub fn arguments(&self) -> &[(String, FunctionArgument)] { + &self.arguments + } + + fn to_wire(&self) -> Result { + let arguments = self + .arguments + .iter() + .map(|(parameter, argument)| { + Ok(ArgumentBindingWire { + parameter: parameter.clone(), + value: argument.to_value_wire()?, + }) + }) + .collect::>>()?; + Ok(FunctionCallWire { + function_id: self.function_id.value.clone(), + arguments, + }) + } + + fn from_wire(wire: FunctionCallWire) -> Result { + let function_id = FunctionId::try_new(wire.function_id)?; + let mut seen = HashSet::with_capacity(wire.arguments.len()); + let mut arguments = Vec::with_capacity(wire.arguments.len()); + for binding in wire.arguments { + if binding.parameter.is_empty() { + return Err(invalid_input( + "FunctionCall parameter names must be non-empty", + )); + } + if !seen.insert(binding.parameter.clone()) { + return Err(invalid_input(format!( + "duplicate FunctionCall binding for parameter `{}`", + binding.parameter + ))); + } + arguments.push(( + binding.parameter, + FunctionArgument::from_value_wire(binding.value)?, + )); + } + Ok(Self { + function_id, + arguments, + }) + } +} + +#[derive(Debug, Serialize, Deserialize)] +#[serde(deny_unknown_fields)] +struct FunctionCallWire { + function_id: String, + arguments: Vec, +} + +#[derive(Debug, Serialize, Deserialize)] +#[serde(deny_unknown_fields)] +struct ArgumentBindingWire { + parameter: String, + value: ArgumentValueWire, +} + +#[derive(Debug, Serialize, Deserialize)] +#[serde(tag = "kind", deny_unknown_fields)] +enum ArgumentValueWire { + #[serde(rename = "field")] + Field { + field_id: i32, + data_type_ipc: String, + }, + #[serde(rename = "literal")] + Literal { ipc: String }, +} + +impl Serialize for FunctionCall { + fn serialize(&self, serializer: S) -> std::result::Result + where + S: Serializer, + { + self.to_wire() + .map_err(serde::ser::Error::custom)? + .serialize(serializer) + } +} + +impl<'de> Deserialize<'de> for FunctionCall { + fn deserialize(deserializer: D) -> std::result::Result + where + D: Deserializer<'de>, + { + let wire = FunctionCallWire::deserialize(deserializer)?; + Self::from_wire(wire).map_err(D::Error::custom) + } +} + +/// Projection completeness of a generated column relative to its dependency epoch. +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +#[non_exhaustive] +pub enum GeneratedColumnStatus { + /// `materialized_epoch == dependency_epoch`. + Complete, + /// `materialized_epoch < dependency_epoch`. + Incomplete, +} + +/// Strict version-1 generated column metadata value. +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct GeneratedColumnDefinition { + output_field_id: i32, + function_call: FunctionCall, + dependency_epoch: u64, + materialized_epoch: u64, +} + +impl GeneratedColumnDefinition { + /// Create a generated-column definition. + /// + /// Rejects a negative `output_field_id` and `materialized_epoch > dependency_epoch`. + pub fn try_new( + output_field_id: i32, + function_call: FunctionCall, + dependency_epoch: u64, + materialized_epoch: u64, + ) -> Result { + if output_field_id < 0 { + return Err(invalid_input( + "generated column output_field_id must be non-negative", + )); + } + if materialized_epoch > dependency_epoch { + return Err(invalid_input( + "materialized_epoch must not be greater than dependency_epoch", + )); + } + Ok(Self { + output_field_id, + function_call, + dependency_epoch, + materialized_epoch, + }) + } + + /// Wire format version (always 1 for this type). + pub fn format_version(&self) -> u32 { + FORMAT_VERSION_V1 + } + + /// Output field id this definition applies to. + pub fn output_field_id(&self) -> i32 { + self.output_field_id + } + + /// Embedded function call. + pub fn function_call(&self) -> &FunctionCall { + &self.function_call + } + + /// Dependency epoch. + pub fn dependency_epoch(&self) -> u64 { + self.dependency_epoch + } + + /// Materialized epoch. + pub fn materialized_epoch(&self) -> u64 { + self.materialized_epoch + } + + /// Completeness status derived from the epochs. + pub fn status(&self) -> GeneratedColumnStatus { + if self.materialized_epoch == self.dependency_epoch { + GeneratedColumnStatus::Complete + } else { + GeneratedColumnStatus::Incomplete + } + } + + /// Serialize to the strict metadata JSON string. + pub fn to_metadata_json(&self) -> Result { + let wire = GeneratedColumnWire { + format_version: FORMAT_VERSION_V1, + output_field_id: self.output_field_id, + function_call: self.function_call.to_wire()?, + dependency_epoch: self.dependency_epoch, + materialized_epoch: self.materialized_epoch, + }; + serde_json::to_string(&wire).map_err(|e| invalid_input(format!("serialize metadata: {e}"))) + } + + /// Deserialize metadata JSON and require `expected_output_field_id`. + /// + /// Decoding is structural: the embedded [`FunctionCall`] is not validated + /// against a catalog [`Function`]. Callers that will execute the call must + /// look up the function and run [`FunctionCall::validate_against`] first. + /// Structural decode still allows query/invalidation paths to inspect + /// stored field ids without a catalog lookup. + pub fn from_metadata_json(json: &str, expected_output_field_id: i32) -> Result { + let value: Value = serde_json::from_str(json) + .map_err(|e| invalid_input(format!("invalid generated column metadata JSON: {e}")))?; + let wire: GeneratedColumnWire = serde_json::from_value(value) + .map_err(|e| invalid_input(format!("invalid generated column metadata: {e}")))?; + if wire.format_version != FORMAT_VERSION_V1 { + return Err(invalid_input(format!( + "unsupported generated column format_version {}", + wire.format_version + ))); + } + if wire.output_field_id < 0 { + return Err(invalid_input( + "generated column output_field_id must be non-negative", + )); + } + if wire.output_field_id != expected_output_field_id { + return Err(invalid_input(format!( + "generated column output_field_id mismatch: metadata has {}, expected {}", + wire.output_field_id, expected_output_field_id + ))); + } + if wire.materialized_epoch > wire.dependency_epoch { + return Err(invalid_input( + "materialized_epoch must not be greater than dependency_epoch", + )); + } + let function_call = FunctionCall::from_wire(wire.function_call)?; + Ok(Self { + output_field_id: wire.output_field_id, + function_call, + dependency_epoch: wire.dependency_epoch, + materialized_epoch: wire.materialized_epoch, + }) + } + + /// Increment `dependency_epoch` by one using checked arithmetic. + /// + /// On overflow, returns an error and leaves the definition unchanged. + pub fn invalidate(&mut self) -> Result<()> { + let next = self + .dependency_epoch + .checked_add(1) + .ok_or_else(|| invalid_input("dependency_epoch overflow"))?; + self.dependency_epoch = next; + Ok(()) + } + + /// Set `materialized_epoch` to the current `dependency_epoch`. + pub fn mark_materialized(&mut self) { + self.materialized_epoch = self.dependency_epoch; + } +} + +#[derive(Debug, Serialize, Deserialize)] +#[serde(deny_unknown_fields)] +struct GeneratedColumnWire { + format_version: u32, + output_field_id: i32, + function_call: FunctionCallWire, + dependency_epoch: u64, + materialized_epoch: u64, +} + +#[cfg(test)] +mod tests { + use super::*; + use arrow_array::Int32Array; + + #[test] + fn type_ipc_rejects_non_canonical_field_name() { + let schema = Schema::new(vec![Field::new("value", DataType::Int32, true)]); + let bytes = schema_to_ipc_file(&schema).expect("schema ipc"); + let err = decode_type_ipc(&bytes).expect_err("non-canonical field name"); + let message = err.to_string().to_lowercase(); + assert!( + message.contains("canonical") || message.contains("schema-only"), + "unexpected error: {message}" + ); + } + + #[test] + fn type_ipc_round_trip_is_byte_identical() { + let data_type = DataType::Timestamp(arrow_schema::TimeUnit::Nanosecond, Some("UTC".into())); + let bytes = encode_type_ipc(&data_type).expect("encode"); + let decoded = decode_type_ipc(&bytes).expect("decode"); + assert_eq!(decoded, data_type); + assert_eq!(encode_type_ipc(&decoded).expect("re-encode"), bytes); + } + + #[test] + fn literal_ipc_rejects_non_canonical_field_name() { + let schema = Arc::new(Schema::new(vec![Field::new("x", DataType::Int32, true)])); + let batch = RecordBatch::try_new( + schema, + vec![Arc::new(Int32Array::from(vec![Some(1)])) as ArrayRef], + ) + .expect("batch"); + let bytes = batches_to_ipc_file(&[batch]).expect("literal ipc"); + let err = decode_literal_ipc(&bytes).expect_err("non-canonical field name"); + let message = err.to_string().to_lowercase(); + assert!( + message.contains("canonical") || message.contains("literal"), + "unexpected error: {message}" + ); + } +} diff --git a/rust/lancedb/src/lib.rs b/rust/lancedb/src/lib.rs index 70d023ccc..55cb06e13 100644 --- a/rust/lancedb/src/lib.rs +++ b/rust/lancedb/src/lib.rs @@ -181,6 +181,7 @@ pub mod dataloader; pub mod embeddings; pub mod error; pub mod expr; +pub mod function; pub mod index; pub mod io; pub mod ipc; diff --git a/rust/lancedb/tests/first_class_function_contract.rs b/rust/lancedb/tests/first_class_function_contract.rs new file mode 100644 index 000000000..3119726e6 --- /dev/null +++ b/rust/lancedb/tests/first_class_function_contract.rs @@ -0,0 +1,1095 @@ +// SPDX-License-Identifier: Apache-2.0 +// SPDX-FileCopyrightText: Copyright The LanceDB Authors + +//! Contract tests for LanceDB Enterprise first-class functions. +//! +//! These tests pin the intended public surface under [`lancedb::function`]. +//! They intentionally fail to compile until that module exists. + +use std::sync::Arc; + +use arrow_array::{ArrayRef, Int32Array, RecordBatch, StringArray}; +use arrow_schema::{DataType, Field, Schema}; +use lancedb::Result; +use lancedb::function::{ + Function, FunctionArgument, FunctionCall, FunctionId, FunctionOutput, FunctionParameter, + FunctionSignature, GENERATED_COLUMN_METADATA_KEY, GeneratedColumnDefinition, + GeneratedColumnStatus, +}; +use lancedb::ipc::{batches_to_ipc_file, schema_to_ipc_file}; +use serde_json::Value; + +fn sample_output() -> FunctionOutput { + FunctionOutput::new(DataType::Int32, true) +} + +fn sample_signature() -> Result { + FunctionSignature::try_new( + vec![ + FunctionParameter::new("x", DataType::Int32), + FunctionParameter::new("label", DataType::Utf8), + ], + sample_output(), + ) +} + +fn sample_function() -> Result { + let id = FunctionId::try_new("fn.exact.example")?; + Ok(Function::new(id, sample_signature()?)) +} + +fn field_arg(field_id: i32, data_type: DataType) -> Result { + FunctionArgument::try_field(field_id, data_type) +} + +fn int_literal(value: Option) -> Result { + FunctionArgument::try_literal(Arc::new(Int32Array::from(vec![value])) as ArrayRef) +} + +fn utf8_literal(value: Option<&str>) -> Result { + FunctionArgument::try_literal(Arc::new(StringArray::from(vec![value])) as ArrayRef) +} + +fn sample_call(function: &Function) -> Result { + FunctionCall::try_new( + function, + vec![ + ("x".to_string(), field_arg(7, DataType::Int32)?), + ("label".to_string(), utf8_literal(Some("ok"))?), + ], + ) +} + +fn assert_json_object_keys_subset(value: &Value, allowed: &[&str]) { + let object = value + .as_object() + .unwrap_or_else(|| panic!("expected JSON object, got {value}")); + for key in object.keys() { + assert!( + allowed.contains(&key.as_str()), + "unexpected JSON key `{key}` in {value}" + ); + } +} + +fn assert_forbidden_function_keys_absent(value: &Value) { + // Function transport forbids catalog/name/version/lineage-style fields on the + // function (or function_call) object itself. Parameter objects still use `name`. + let forbidden = [ + "name", + "catalog", + "catalog_name", + "version", + "function_version", + "FunctionVersion", + "lineage", + "user_version", + "idempotency_key", + "digest", + "storage", + "storage_location", + "deterministic", + "null_policy", + "nullPolicy", + ]; + let object = value + .as_object() + .unwrap_or_else(|| panic!("expected JSON object, got {value}")); + for key in forbidden { + assert!( + !object.contains_key(key), + "function JSON must not contain forbidden key `{key}`: {value}" + ); + } +} + +/// Minimal RFC 4648 base64 encoder for test-only wire mutation. +fn base64_encode(input: &[u8]) -> String { + const TABLE: &[u8; 64] = b"ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz0123456789+/"; + let mut out = String::with_capacity(input.len().div_ceil(3) * 4); + for chunk in input.chunks(3) { + let b0 = chunk[0] as u32; + let b1 = chunk.get(1).copied().unwrap_or(0) as u32; + let b2 = chunk.get(2).copied().unwrap_or(0) as u32; + let triple = (b0 << 16) | (b1 << 8) | b2; + out.push(TABLE[((triple >> 18) & 0x3F) as usize] as char); + out.push(TABLE[((triple >> 12) & 0x3F) as usize] as char); + if chunk.len() > 1 { + out.push(TABLE[((triple >> 6) & 0x3F) as usize] as char); + } else { + out.push('='); + } + if chunk.len() > 2 { + out.push(TABLE[(triple & 0x3F) as usize] as char); + } else { + out.push('='); + } + } + out +} + +/// Minimal RFC 4648 base64 decoder for test-only wire mutation. +fn base64_decode(input: &str) -> std::result::Result, String> { + fn decode_char(c: u8) -> std::result::Result { + match c { + b'A'..=b'Z' => Ok(c - b'A'), + b'a'..=b'z' => Ok(c - b'a' + 26), + b'0'..=b'9' => Ok(c - b'0' + 52), + b'+' => Ok(62), + b'/' => Ok(63), + _ => Err(format!("invalid base64 byte: {c}")), + } + } + + let bytes = input.as_bytes(); + if !bytes.len().is_multiple_of(4) { + return Err("base64 length must be a multiple of 4".into()); + } + let mut out = Vec::with_capacity(bytes.len() / 4 * 3); + for chunk in bytes.chunks(4) { + let pad = chunk.iter().filter(|&&b| b == b'=').count(); + let c0 = decode_char(chunk[0])?; + let c1 = decode_char(chunk[1])?; + let c2 = if chunk[2] == b'=' { + 0 + } else { + decode_char(chunk[2])? + }; + let c3 = if chunk[3] == b'=' { + 0 + } else { + decode_char(chunk[3])? + }; + let triple = ((c0 as u32) << 18) | ((c1 as u32) << 12) | ((c2 as u32) << 6) | (c3 as u32); + out.push(((triple >> 16) & 0xFF) as u8); + if pad < 2 { + out.push(((triple >> 8) & 0xFF) as u8); + } + if pad < 1 { + out.push((triple & 0xFF) as u8); + } + } + Ok(out) +} + +fn first_literal_ipc(metadata: &Value) -> String { + let args = metadata + .pointer("/function_call/arguments") + .and_then(Value::as_array) + .expect("function_call.arguments array"); + for arg in args { + let value = arg + .get("value") + .and_then(Value::as_object) + .expect("argument.value object"); + if value.get("kind").and_then(Value::as_str) == Some("literal") { + return value + .get("ipc") + .and_then(Value::as_str) + .expect("literal value.ipc string") + .to_string(); + } + } + panic!("expected at least one literal argument with ipc"); +} + +fn set_first_literal_ipc(metadata: &mut Value, ipc: String) { + let args = metadata + .pointer_mut("/function_call/arguments") + .and_then(Value::as_array_mut) + .expect("function_call.arguments array"); + for arg in args { + let value = arg + .get_mut("value") + .and_then(Value::as_object_mut) + .expect("argument.value object"); + if value.get("kind").and_then(Value::as_str) == Some("literal") { + value.insert("ipc".into(), Value::String(ipc)); + return; + } + } + panic!("expected at least one literal argument with ipc"); +} + +fn metadata_with_literal_call() -> Result<(String, i32)> { + let function = sample_function()?; + let call = FunctionCall::try_new( + &function, + vec![ + ("x".to_string(), int_literal(Some(7))?), + ("label".to_string(), utf8_literal(Some("ok"))?), + ], + )?; + let definition = GeneratedColumnDefinition::try_new(13, call, 1, 1)?; + Ok((definition.to_metadata_json()?, 13)) +} + +fn int32_batch_ipc(values: &[Option]) -> Result> { + let schema = Arc::new(Schema::new(vec![Field::new( + "value", + DataType::Int32, + true, + )])); + let array = Arc::new(Int32Array::from(values.to_vec())) as ArrayRef; + let batch = RecordBatch::try_new(schema, vec![array]).expect("record batch"); + batches_to_ipc_file(&[batch]) +} + +fn int32_two_one_row_batches_ipc() -> Result> { + let schema = Arc::new(Schema::new(vec![Field::new( + "value", + DataType::Int32, + true, + )])); + let batch_a = RecordBatch::try_new( + schema.clone(), + vec![Arc::new(Int32Array::from(vec![Some(1)])) as ArrayRef], + ) + .expect("batch a"); + let batch_b = RecordBatch::try_new( + schema, + vec![Arc::new(Int32Array::from(vec![Some(2)])) as ArrayRef], + ) + .expect("batch b"); + batches_to_ipc_file(&[batch_a, batch_b]) +} + +fn int32_two_column_one_row_ipc() -> Result> { + let schema = Arc::new(Schema::new(vec![ + Field::new("a", DataType::Int32, true), + Field::new("b", DataType::Int32, true), + ])); + let batch = RecordBatch::try_new( + schema, + vec![ + Arc::new(Int32Array::from(vec![Some(1)])) as ArrayRef, + Arc::new(Int32Array::from(vec![Some(2)])) as ArrayRef, + ], + ) + .expect("two-column batch"); + batches_to_ipc_file(&[batch]) +} + +/// Mutate the first signature `data_type_ipc` string in a Function JSON value. +fn set_first_parameter_data_type_ipc(function_json: &mut Value, ipc: String) { + let parameters = function_json + .pointer_mut("/signature/parameters") + .and_then(Value::as_array_mut) + .expect("signature.parameters array"); + let first = parameters + .first_mut() + .expect("at least one signature parameter"); + first + .as_object_mut() + .expect("parameter object") + .insert("data_type_ipc".into(), Value::String(ipc)); +} + +#[test] +fn function_id_rejects_empty() { + let err = FunctionId::try_new("").expect_err("empty FunctionId must be rejected"); + let message = err.to_string(); + assert!( + message.to_lowercase().contains("empty") || message.to_lowercase().contains("non-empty"), + "unexpected error: {message}" + ); +} + +#[test] +fn function_id_preserves_exact_opaque_value() -> Result<()> { + let id = FunctionId::try_new("fn.exact.opaque-value")?; + assert_eq!(id.as_str(), "fn.exact.opaque-value"); + Ok(()) +} + +#[test] +fn function_signature_rejects_duplicate_or_empty_parameter_names() { + let duplicate = FunctionSignature::try_new( + vec![ + FunctionParameter::new("x", DataType::Int32), + FunctionParameter::new("x", DataType::Utf8), + ], + sample_output(), + ); + assert!( + duplicate.is_err(), + "duplicate parameter names must be rejected" + ); + + let empty = FunctionSignature::try_new( + vec![FunctionParameter::new("", DataType::Int32)], + sample_output(), + ); + assert!(empty.is_err(), "empty parameter names must be rejected"); +} + +#[test] +fn function_json_round_trip_pins_output_and_ipc_wire_shape() -> Result<()> { + let function = sample_function()?; + let json = serde_json::to_value(&function).expect("serialize Function"); + // Function transport JSON is a strict format_version = 1 envelope (wire mechanism). + assert_json_object_keys_subset(&json, &["format_version", "id", "signature"]); + assert_eq!(json["format_version"], 1); + assert_forbidden_function_keys_absent(&json); + + let signature = json + .get("signature") + .and_then(Value::as_object) + .expect("signature object"); + assert_json_object_keys_subset(&Value::Object(signature.clone()), &["parameters", "output"]); + + let parameters = signature + .get("parameters") + .and_then(Value::as_array) + .expect("parameters array"); + assert_eq!(parameters.len(), 2); + for parameter in parameters { + assert_json_object_keys_subset(parameter, &["name", "data_type_ipc"]); + let ipc = parameter + .get("data_type_ipc") + .and_then(Value::as_str) + .expect("parameter data_type_ipc"); + assert!(!ipc.is_empty(), "parameter data_type_ipc must be base64"); + base64_decode(ipc).expect("parameter data_type_ipc must be valid base64"); + } + + let output = signature + .get("output") + .and_then(Value::as_object) + .expect("output object"); + assert_json_object_keys_subset( + &Value::Object(output.clone()), + &["data_type_ipc", "nullable"], + ); + let output_ipc = output + .get("data_type_ipc") + .and_then(Value::as_str) + .expect("output data_type_ipc"); + base64_decode(output_ipc).expect("output data_type_ipc must be valid base64"); + assert_eq!(output.get("nullable"), Some(&Value::Bool(true))); + + assert_eq!(json["id"], Value::String("fn.exact.example".to_string())); + + let restored: Function = serde_json::from_value(json.clone()).expect("deserialize Function"); + assert_eq!(restored.id().as_str(), function.id().as_str()); + assert_eq!( + restored.signature().parameters().len(), + function.signature().parameters().len() + ); + assert_eq!( + restored.signature().parameters()[0].name(), + function.signature().parameters()[0].name() + ); + assert_eq!( + restored.signature().parameters()[0].data_type(), + function.signature().parameters()[0].data_type() + ); + assert_eq!( + restored.signature().parameters()[1].name(), + function.signature().parameters()[1].name() + ); + assert_eq!( + restored.signature().parameters()[1].data_type(), + function.signature().parameters()[1].data_type() + ); + assert_eq!( + restored.signature().output().data_type(), + function.signature().output().data_type() + ); + assert_eq!( + restored.signature().output().nullable(), + function.signature().output().nullable() + ); + assert_eq!(restored.signature().output().data_type(), &DataType::Int32); + assert!(restored.signature().output().nullable()); + + // Handle / wire identity is the opaque FunctionId only. + assert_eq!( + serde_json::to_value(&restored).expect("re-serialize Function"), + json + ); + + let mut unknown = json.clone(); + unknown + .as_object_mut() + .unwrap() + .insert("unexpected_field".into(), Value::Bool(true)); + assert!( + serde_json::from_value::(unknown).is_err(), + "Function JSON must reject unknown fields" + ); + + let mut unknown_version = json.clone(); + unknown_version["format_version"] = Value::from(2); + assert!( + serde_json::from_value::(unknown_version).is_err(), + "Function JSON must reject unknown format_version" + ); + Ok(()) +} + +#[test] +fn function_signature_data_type_ipc_decode_is_fail_closed() -> Result<()> { + let function = sample_function()?; + let json = serde_json::to_value(&function).expect("serialize Function"); + let original_ipc = json["signature"]["parameters"][0]["data_type_ipc"] + .as_str() + .expect("parameter data_type_ipc") + .to_string(); + let original_bytes = base64_decode(&original_ipc).expect("parameter data_type_ipc base64"); + + let mut invalid_base64 = json.clone(); + set_first_parameter_data_type_ipc(&mut invalid_base64, "!!!".into()); + assert!( + serde_json::from_value::(invalid_base64).is_err(), + "invalid base64 data_type_ipc must be rejected" + ); + + let multi_field_schema = Schema::new(vec![ + Field::new("a", DataType::Int32, true), + Field::new("b", DataType::Utf8, true), + ]); + let multi_field_ipc = schema_to_ipc_file(&multi_field_schema)?; + let mut multi_field = json.clone(); + set_first_parameter_data_type_ipc(&mut multi_field, base64_encode(&multi_field_ipc)); + assert!( + serde_json::from_value::(multi_field).is_err(), + "schema-only IPC with more than one field must be rejected" + ); + + let mut trailing = original_bytes; + trailing.extend_from_slice(b"extra"); + let mut trailing_ipc = json.clone(); + set_first_parameter_data_type_ipc(&mut trailing_ipc, base64_encode(&trailing)); + assert!( + serde_json::from_value::(trailing_ipc).is_err(), + "single-field IPC with trailing bytes must be rejected" + ); + + // Type encoding is schema-only: a valid one-field IPC file that also carries a + // record batch must be rejected (data-bearing IPC is not a type encoding). + let data_bearing_schema = Arc::new(Schema::new(vec![Field::new( + "value", + DataType::Int32, + true, + )])); + let data_bearing_batch = RecordBatch::try_new( + data_bearing_schema, + vec![Arc::new(Int32Array::from(vec![Some(1)])) as ArrayRef], + ) + .expect("one-field one-row batch"); + let data_bearing_ipc = batches_to_ipc_file(&[data_bearing_batch])?; + let mut data_bearing = json.clone(); + set_first_parameter_data_type_ipc(&mut data_bearing, base64_encode(&data_bearing_ipc)); + assert!( + serde_json::from_value::(data_bearing).is_err(), + "one-field IPC that contains a record batch must be rejected for data_type_ipc" + ); + Ok(()) +} + +fn sample_definition_metadata() -> Result<(Function, String, i32)> { + let function = sample_function()?; + let call = sample_call(&function)?; + let definition = GeneratedColumnDefinition::try_new(11, call, 3, 3)?; + Ok((function, definition.to_metadata_json()?, 11)) +} + +fn schema_only_type_ipc_b64(data_type: DataType) -> Result { + let schema = Schema::new(vec![Field::new("", data_type, true)]); + Ok(base64_encode(&schema_to_ipc_file(&schema)?)) +} + +#[test] +fn function_call_validate_against_requires_exact_identity_and_signature() -> Result<()> { + let (function, metadata_json, output_field_id) = sample_definition_metadata()?; + + // try_new and a normal metadata round-trip must validate against the catalog Function. + let call = sample_call(&function)?; + call.validate_against(&function)?; + let restored = GeneratedColumnDefinition::from_metadata_json(&metadata_json, output_field_id)?; + restored.function_call().validate_against(&function)?; + + let value: Value = serde_json::from_str(&metadata_json).expect("metadata JSON"); + + // Reordered wire bindings decode structurally, then fail validation. + let mut reordered = value.clone(); + let args = reordered + .pointer_mut("/function_call/arguments") + .and_then(Value::as_array_mut) + .expect("function_call.arguments"); + args.swap(0, 1); + let reordered_def = + GeneratedColumnDefinition::from_metadata_json(&reordered.to_string(), output_field_id)?; + assert!( + reordered_def + .function_call() + .validate_against(&function) + .is_err(), + "reordered wire bindings must fail validate_against" + ); + + // Changed parameter name. + let mut renamed = value.clone(); + renamed["function_call"]["arguments"][0]["parameter"] = Value::String("renamed".into()); + let renamed_def = + GeneratedColumnDefinition::from_metadata_json(&renamed.to_string(), output_field_id)?; + assert!( + renamed_def + .function_call() + .validate_against(&function) + .is_err(), + "changed parameter name must fail validate_against" + ); + + // Missing binding. + let mut missing = value.clone(); + missing["function_call"]["arguments"] + .as_array_mut() + .expect("arguments") + .pop(); + let missing_def = + GeneratedColumnDefinition::from_metadata_json(&missing.to_string(), output_field_id)?; + assert!( + missing_def + .function_call() + .validate_against(&function) + .is_err(), + "missing binding must fail validate_against" + ); + + // Extra binding. + let mut extra = value.clone(); + let extra_binding = serde_json::json!({ + "parameter": "extra", + "value": { + "kind": "field", + "field_id": 99, + "data_type_ipc": schema_only_type_ipc_b64(DataType::Int32)? + } + }); + extra["function_call"]["arguments"] + .as_array_mut() + .expect("arguments") + .push(extra_binding); + let extra_def = + GeneratedColumnDefinition::from_metadata_json(&extra.to_string(), output_field_id)?; + assert!( + extra_def + .function_call() + .validate_against(&function) + .is_err(), + "extra binding must fail validate_against" + ); + + // Changed argument Arrow type (field argument data_type_ipc). + let mut type_changed = value.clone(); + type_changed["function_call"]["arguments"][0]["value"]["data_type_ipc"] = + Value::String(schema_only_type_ipc_b64(DataType::Utf8)?); + let type_changed_def = + GeneratedColumnDefinition::from_metadata_json(&type_changed.to_string(), output_field_id)?; + assert!( + type_changed_def + .function_call() + .validate_against(&function) + .is_err(), + "changed argument type must fail validate_against" + ); + + // Different Function ID. + let mut different_id = value.clone(); + different_id["function_call"]["function_id"] = Value::String("fn.other".into()); + let different_id_def = + GeneratedColumnDefinition::from_metadata_json(&different_id.to_string(), output_field_id)?; + assert!( + different_id_def + .function_call() + .validate_against(&function) + .is_err(), + "different FunctionId must fail validate_against" + ); + + Ok(()) +} + +#[test] +fn function_call_requires_complete_unique_ordered_typed_bindings() -> Result<()> { + let function = sample_function()?; + + let missing = FunctionCall::try_new( + &function, + vec![("x".to_string(), field_arg(1, DataType::Int32)?)], + ); + assert!(missing.is_err(), "missing parameter must fail"); + + let duplicate = FunctionCall::try_new( + &function, + vec![ + ("x".to_string(), field_arg(1, DataType::Int32)?), + ("label".to_string(), utf8_literal(Some("a"))?), + ("x".to_string(), field_arg(2, DataType::Int32)?), + ], + ); + assert!(duplicate.is_err(), "duplicate parameter name must fail"); + + let unknown = FunctionCall::try_new( + &function, + vec![ + ("x".to_string(), field_arg(1, DataType::Int32)?), + ("label".to_string(), utf8_literal(Some("a"))?), + ("extra".to_string(), int_literal(Some(1))?), + ], + ); + assert!(unknown.is_err(), "unknown parameter name must fail"); + + let type_mismatch = FunctionCall::try_new( + &function, + vec![ + ("x".to_string(), field_arg(1, DataType::Utf8)?), + ("label".to_string(), utf8_literal(Some("a"))?), + ], + ); + assert!(type_mismatch.is_err(), "argument type mismatch must fail"); + + // Binding is by explicit parameter name: wrong order is accepted then normalized. + let wrong_order_input = FunctionCall::try_new( + &function, + vec![ + ("label".to_string(), utf8_literal(Some("a"))?), + ("x".to_string(), field_arg(9, DataType::Int32)?), + ], + )?; + let args = wrong_order_input.arguments(); + assert_eq!(args.len(), 2); + assert_eq!(args[0].0, "x"); + assert_eq!(args[1].0, "label"); + assert_eq!(wrong_order_input.function_id().as_str(), "fn.exact.example"); + Ok(()) +} + +#[test] +fn field_argument_stores_stable_field_id_not_column_name() -> Result<()> { + let function = sample_function()?; + let call = FunctionCall::try_new( + &function, + vec![ + ("x".to_string(), field_arg(42, DataType::Int32)?), + ("label".to_string(), utf8_literal(Some("n"))?), + ], + )?; + + let (name, arg) = &call.arguments()[0]; + assert_eq!(name, "x"); + assert_eq!(arg.field_id(), Some(42)); + assert_eq!(arg.data_type(), &DataType::Int32); + assert!(arg.literal_array().is_none()); + + let json = serde_json::to_value(&call).expect("serialize FunctionCall"); + let encoded = json.to_string(); + assert!( + !encoded.contains("column_name") + && !encoded.contains("columnName") + && !encoded.to_lowercase().contains("\"column\""), + "field argument must not persist a column name: {json}" + ); + assert!( + encoded.contains("42") && (encoded.contains("field_id") || encoded.contains("fieldId")), + "field argument must persist stable field id: {json}" + ); + + let arguments = json + .get("arguments") + .and_then(Value::as_array) + .expect("arguments array"); + assert_eq!(arguments.len(), 2); + assert_json_object_keys_subset(&arguments[0], &["parameter", "value"]); + assert_eq!(arguments[0]["parameter"], Value::String("x".into())); + assert_eq!(arguments[0]["value"]["kind"], Value::String("field".into())); + assert_json_object_keys_subset(&arguments[1], &["parameter", "value"]); + assert_eq!(arguments[1]["parameter"], Value::String("label".into())); + assert_eq!( + arguments[1]["value"]["kind"], + Value::String("literal".into()) + ); + let literal_ipc = arguments[1]["value"]["ipc"] + .as_str() + .expect("literal value.ipc"); + base64_decode(literal_ipc).expect("literal ipc must be valid base64"); + Ok(()) +} + +#[test] +fn field_argument_rejects_negative_field_id() { + let err = FunctionArgument::try_field(-1, DataType::Int32) + .expect_err("negative field id must be rejected"); + let message = err.to_string().to_lowercase(); + assert!( + message.contains("field") || message.contains("negative"), + "unexpected error: {message}" + ); +} + +#[test] +fn typed_null_literal_round_trips_through_generated_metadata_json() -> Result<()> { + let function = sample_function()?; + let null_literal = int_literal(None)?; + assert_eq!(null_literal.data_type(), &DataType::Int32); + assert!(null_literal.is_typed_null()); + + let call = FunctionCall::try_new( + &function, + vec![ + ("x".to_string(), null_literal), + ("label".to_string(), utf8_literal(Some("ok"))?), + ], + )?; + let definition = GeneratedColumnDefinition::try_new(3, call, 2, 2)?; + let metadata_json = definition.to_metadata_json()?; + let restored = GeneratedColumnDefinition::from_metadata_json(&metadata_json, 3)?; + let decoded = &restored.function_call().arguments()[0].1; + assert!(decoded.is_typed_null()); + assert_eq!(decoded.data_type(), &DataType::Int32); + + let array = decoded + .literal_array() + .expect("typed NULL must expose a one-row array"); + assert_eq!(array.len(), 1); + assert!(array.is_null(0)); + assert_eq!(restored.to_metadata_json()?, metadata_json); + Ok(()) +} + +#[test] +fn generated_column_metadata_round_trips_exact_function_id() -> Result<()> { + assert_eq!(GENERATED_COLUMN_METADATA_KEY, "lancedb::generated_column"); + assert!(!GENERATED_COLUMN_METADATA_KEY.is_empty()); + + let function = sample_function()?; + let call = sample_call(&function)?; + let definition = GeneratedColumnDefinition::try_new(/* output_field_id */ 11, call, 3, 3)?; + + assert_eq!(definition.format_version(), 1); + assert_eq!(definition.output_field_id(), 11); + assert_eq!(definition.dependency_epoch(), 3); + assert_eq!(definition.materialized_epoch(), 3); + assert_eq!(definition.status(), GeneratedColumnStatus::Complete); + assert_eq!( + definition.function_call().function_id().as_str(), + "fn.exact.example" + ); + + let metadata_json = definition.to_metadata_json()?; + let value: Value = serde_json::from_str(&metadata_json).expect("metadata JSON"); + assert_json_object_keys_subset( + &value, + &[ + "format_version", + "output_field_id", + "function_call", + "dependency_epoch", + "materialized_epoch", + ], + ); + assert_eq!(value["format_version"], 1); + assert_eq!(value["output_field_id"], 11); + assert_eq!( + value["function_call"]["function_id"], + Value::String("fn.exact.example".to_string()) + ); + assert_forbidden_function_keys_absent(&value["function_call"]); + + let arguments = value["function_call"]["arguments"] + .as_array() + .expect("function_call.arguments"); + for argument in arguments { + assert_json_object_keys_subset(argument, &["parameter", "value"]); + let kind = argument["value"]["kind"] + .as_str() + .expect("argument value.kind"); + assert!( + kind == "field" || kind == "literal", + "unexpected argument kind `{kind}`" + ); + if kind == "literal" { + base64_decode(argument["value"]["ipc"].as_str().expect("literal ipc")) + .expect("literal ipc base64"); + } + } + + let restored = GeneratedColumnDefinition::from_metadata_json(&metadata_json, 11)?; + assert_eq!(restored.output_field_id(), definition.output_field_id()); + assert_eq!( + restored.function_call().function_id().as_str(), + definition.function_call().function_id().as_str() + ); + assert_eq!(restored.dependency_epoch(), definition.dependency_epoch()); + assert_eq!( + restored.materialized_epoch(), + definition.materialized_epoch() + ); + assert_eq!(restored.status(), definition.status()); + assert_eq!(restored.to_metadata_json()?, metadata_json); + Ok(()) +} + +#[test] +fn metadata_decode_is_fail_closed_for_unknown_field_variant_and_version() -> Result<()> { + let function = sample_function()?; + let call = sample_call(&function)?; + let definition = GeneratedColumnDefinition::try_new(5, call, 1, 1)?; + let value: Value = + serde_json::from_str(&definition.to_metadata_json()?).expect("metadata JSON"); + + let mut unknown_field = value.clone(); + unknown_field + .as_object_mut() + .unwrap() + .insert("unexpected_field".into(), Value::Bool(true)); + assert!( + GeneratedColumnDefinition::from_metadata_json(&unknown_field.to_string(), 5).is_err(), + "unknown metadata field must be rejected" + ); + + let mut unknown_version = value.clone(); + unknown_version["format_version"] = Value::from(2); + assert!( + GeneratedColumnDefinition::from_metadata_json(&unknown_version.to_string(), 5).is_err(), + "unknown format_version must be rejected" + ); + + let mut unknown_variant = value.clone(); + let argument_value = unknown_variant + .pointer_mut("/function_call/arguments/0/value") + .and_then(Value::as_object_mut) + .expect("argument value object"); + argument_value.clear(); + argument_value.insert("kind".into(), Value::String("expression".into())); + argument_value.insert("sql".into(), Value::String("x + 1".into())); + assert!( + GeneratedColumnDefinition::from_metadata_json(&unknown_variant.to_string(), 5).is_err(), + "unknown argument variant must be rejected" + ); + Ok(()) +} + +#[test] +fn literal_ipc_rejects_invalid_trailing_schema_only_zero_and_multi_row_payloads() -> Result<()> { + let multi_row = + FunctionArgument::try_literal(Arc::new(Int32Array::from(vec![1, 2])) as ArrayRef); + assert!(multi_row.is_err(), "multi-row literal must be rejected"); + + let zero_row = + FunctionArgument::try_literal(Arc::new(Int32Array::from(Vec::::new())) as ArrayRef); + assert!(zero_row.is_err(), "zero-row literal must be rejected"); + + let (metadata_json, output_field_id) = metadata_with_literal_call()?; + let mut value: Value = serde_json::from_str(&metadata_json).expect("metadata JSON"); + let original_bytes = base64_decode(&first_literal_ipc(&value)).expect("literal ipc base64"); + + set_first_literal_ipc(&mut value, base64_encode(b"not-arrow-ipc")); + assert!( + GeneratedColumnDefinition::from_metadata_json(&value.to_string(), output_field_id).is_err(), + "invalid Arrow IPC must be rejected" + ); + + let mut trailing = original_bytes; + trailing.extend_from_slice(b"extra"); + value = serde_json::from_str(&metadata_json).expect("metadata JSON"); + set_first_literal_ipc(&mut value, base64_encode(&trailing)); + assert!( + GeneratedColumnDefinition::from_metadata_json(&value.to_string(), output_field_id).is_err(), + "trailing IPC payload must be rejected" + ); + + let schema = Schema::new(vec![Field::new("value", DataType::Int32, true)]); + let schema_only = schema_to_ipc_file(&schema)?; + value = serde_json::from_str(&metadata_json).expect("metadata JSON"); + set_first_literal_ipc(&mut value, base64_encode(&schema_only)); + assert!( + GeneratedColumnDefinition::from_metadata_json(&value.to_string(), output_field_id).is_err(), + "schema-only IPC must be rejected for literal decode" + ); + + let zero_row_ipc = int32_batch_ipc(&[])?; + value = serde_json::from_str(&metadata_json).expect("metadata JSON"); + set_first_literal_ipc(&mut value, base64_encode(&zero_row_ipc)); + assert!( + GeneratedColumnDefinition::from_metadata_json(&value.to_string(), output_field_id).is_err(), + "zero-row literal IPC must be rejected" + ); + + let multi_row_ipc = int32_batch_ipc(&[Some(1), Some(2)])?; + value = serde_json::from_str(&metadata_json).expect("metadata JSON"); + set_first_literal_ipc(&mut value, base64_encode(&multi_row_ipc)); + assert!( + GeneratedColumnDefinition::from_metadata_json(&value.to_string(), output_field_id).is_err(), + "multi-row literal IPC must be rejected" + ); + + value = serde_json::from_str(&metadata_json).expect("metadata JSON"); + set_first_literal_ipc(&mut value, "!!!".into()); + assert!( + GeneratedColumnDefinition::from_metadata_json(&value.to_string(), output_field_id).is_err(), + "syntactically invalid base64 literal ipc must be rejected" + ); + + let two_batches_ipc = int32_two_one_row_batches_ipc()?; + value = serde_json::from_str(&metadata_json).expect("metadata JSON"); + set_first_literal_ipc(&mut value, base64_encode(&two_batches_ipc)); + assert!( + GeneratedColumnDefinition::from_metadata_json(&value.to_string(), output_field_id).is_err(), + "literal IPC with two one-row batches must be rejected" + ); + + let two_column_ipc = int32_two_column_one_row_ipc()?; + value = serde_json::from_str(&metadata_json).expect("metadata JSON"); + set_first_literal_ipc(&mut value, base64_encode(&two_column_ipc)); + assert!( + GeneratedColumnDefinition::from_metadata_json(&value.to_string(), output_field_id).is_err(), + "literal IPC with one row but two columns must be rejected" + ); + Ok(()) +} + +#[test] +fn generated_metadata_rejects_output_field_id_mismatch() -> Result<()> { + let function = sample_function()?; + let call = sample_call(&function)?; + let definition = GeneratedColumnDefinition::try_new(19, call, 2, 2)?; + let metadata_json = definition.to_metadata_json()?; + + let err = GeneratedColumnDefinition::from_metadata_json(&metadata_json, 20) + .expect_err("output field id mismatch must fail closed"); + let message = err.to_string().to_lowercase(); + assert!( + message.contains("field") || message.contains("mismatch"), + "unexpected error: {message}" + ); + Ok(()) +} + +#[test] +fn epochs_complete_to_incomplete_to_complete_and_overflow() -> Result<()> { + let function = sample_function()?; + + let mut complete = GeneratedColumnDefinition::try_new(1, sample_call(&function)?, 4, 4)?; + assert_eq!(complete.status(), GeneratedColumnStatus::Complete); + + let incomplete = GeneratedColumnDefinition::try_new(1, sample_call(&function)?, 5, 4)?; + assert_eq!(incomplete.status(), GeneratedColumnStatus::Incomplete); + + let invalid = GeneratedColumnDefinition::try_new(1, sample_call(&function)?, 4, 5); + assert!( + invalid.is_err(), + "materialized_epoch greater than dependency_epoch must be invalid" + ); + + // Complete -> Incomplete via checked dependency invalidation. + complete.invalidate()?; + assert_eq!(complete.dependency_epoch(), 5); + assert_eq!(complete.materialized_epoch(), 4); + assert_eq!(complete.status(), GeneratedColumnStatus::Incomplete); + + // Incomplete -> Complete via materialization at the current dependency epoch. + complete.mark_materialized(); + assert_eq!(complete.dependency_epoch(), 5); + assert_eq!(complete.materialized_epoch(), 5); + assert_eq!(complete.status(), GeneratedColumnStatus::Complete); + + let mut at_max = + GeneratedColumnDefinition::try_new(1, sample_call(&function)?, u64::MAX, u64::MAX)?; + assert!( + at_max.invalidate().is_err(), + "dependency_epoch overflow must fail" + ); + + // Typed NULL materialization state is distinct from Incomplete projection. + let null_call = FunctionCall::try_new( + &function, + vec![ + ("x".to_string(), int_literal(None)?), + ("label".to_string(), utf8_literal(None)?), + ], + )?; + let complete_with_nulls = GeneratedColumnDefinition::try_new(1, null_call, 9, 9)?; + assert_eq!( + complete_with_nulls.status(), + GeneratedColumnStatus::Complete + ); + assert!( + complete_with_nulls.function_call().arguments()[0] + .1 + .is_typed_null() + ); + assert_ne!( + complete_with_nulls.status(), + GeneratedColumnStatus::Incomplete + ); + Ok(()) +} + +#[test] +fn arrow_type_and_literal_encoding_is_deterministic_through_serde_json() -> Result<()> { + let data_type = DataType::Timestamp(arrow_schema::TimeUnit::Nanosecond, Some("UTC".into())); + let signature_a = FunctionSignature::try_new( + vec![FunctionParameter::new("ts", data_type.clone())], + FunctionOutput::new(data_type.clone(), false), + )?; + let signature_b = FunctionSignature::try_new( + vec![FunctionParameter::new("ts", data_type.clone())], + FunctionOutput::new(data_type, false), + )?; + let function_a = Function::new(FunctionId::try_new("fn.deterministic")?, signature_a); + let function_b = Function::new(FunctionId::try_new("fn.deterministic")?, signature_b); + let json_a = serde_json::to_value(&function_a).expect("serialize function_a"); + let json_b = serde_json::to_value(&function_b).expect("serialize function_b"); + assert_eq!( + json_a, json_b, + "Function Arrow IPC wire encoding must be deterministic" + ); + assert_eq!( + json_a["signature"]["parameters"][0]["data_type_ipc"], + json_b["signature"]["parameters"][0]["data_type_ipc"] + ); + assert_eq!( + json_a["signature"]["output"]["data_type_ipc"], + json_b["signature"]["output"]["data_type_ipc"] + ); + + let utf8_function = Function::new( + FunctionId::try_new("fn.literal.deterministic")?, + FunctionSignature::try_new( + vec![FunctionParameter::new("label", DataType::Utf8)], + FunctionOutput::new(DataType::Utf8, true), + )?, + ); + let call_a = FunctionCall::try_new( + &utf8_function, + vec![("label".to_string(), utf8_literal(Some("deterministic"))?)], + )?; + let call_b = FunctionCall::try_new( + &utf8_function, + vec![("label".to_string(), utf8_literal(Some("deterministic"))?)], + )?; + let def_a = GeneratedColumnDefinition::try_new(1, call_a, 1, 1)?; + let def_b = GeneratedColumnDefinition::try_new(1, call_b, 1, 1)?; + assert_eq!( + def_a.to_metadata_json()?, + def_b.to_metadata_json()?, + "literal IPC encoding through metadata JSON must be deterministic" + ); + + let metadata: Value = serde_json::from_str(&def_a.to_metadata_json()?).expect("metadata JSON"); + let literal_ipc = metadata["function_call"]["arguments"][0]["value"]["ipc"] + .as_str() + .expect("literal ipc"); + base64_decode(literal_ipc).expect("literal ipc base64"); + + // Contract: no untyped JSON / Lance JsonDataType substitute for Arrow types. + let encoded = json_a.to_string(); + assert!( + !encoded.contains("JsonDataType") && !encoded.contains("json_type"), + "signature types must not use Lance JsonDataType: {json_a}" + ); + Ok(()) +}