use std::collections::{BTreeSet, HashMap, HashSet};
use std::sync::Arc;
use arrow_schema::{DataType, Field as ArrowField, Fields, Schema as ArrowSchema, SchemaRef};
use datafusion_common::{ScalarValue, tree_node::TreeNode};
use datafusion_expr::Expr;
use datafusion_physical_plan::PhysicalExpr;
use lance::dataset::NewColumnTransform;
use lance_arrow::{ARROW_EXT_NAME_KEY, BLOB_V2_EXT_NAME, FieldExt};
use lance_core::datatypes::{
BLOB_V2_DESC_FIELD, BlobV2Layout, format_field_path_minimal, parse_field_path,
};
use lance_datafusion::planner::Planner;
use lance_namespace::models::{JsonArrowDataType, JsonArrowField, JsonArrowSchema};
use serde::{Deserialize, Serialize};
use serde_json::Value;
use crate::function::{FUNCTION_BLOB_V2_TYPE, FunctionApplication, FunctionBinding};
use crate::utils::resolve_arrow_field_path;
use crate::{Error, Result};
pub const COMPUTED_COLUMN_META_KEY: &str = "computed_column";
pub const KIND_META_KEY: &str = "computed_column.kind";
pub const EXPRESSION_META_KEY: &str = "computed_column.expression";
pub const INPUTS_META_KEY: &str = "computed_column.inputs";
pub const FUNCTION_BINDING_ID_META_KEY: &str = "computed_column.function.binding_id";
pub const FUNCTION_OUTPUT_ORDINAL_META_KEY: &str = "computed_column.function.output_ordinal";
pub const FUNCTION_BINDINGS_META_KEY: &str = "lancedb::function_bindings";
pub const FUNCTION_BINDINGS_VERSION: u32 = 1;
pub const SQL_KIND: &str = "sql";
pub const FUNCTION_KIND: &str = "function";
pub const WHOLE_RESULT_FIELD: &str = "$value";
#[derive(Debug, Clone, PartialEq, Eq)]
#[non_exhaustive]
pub enum ComputedColumnKind {
Sql {
expression: String,
},
Function {
binding_id: String,
output_ordinal: u32,
},
Unrecognized {
kind: String,
},
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct ComputedColumn {
pub name: String,
pub kind: ComputedColumnKind,
pub inputs: Vec<String>,
}
fn computed_column_metadata(expression: &str, inputs: &[String]) -> HashMap<String, String> {
HashMap::from([
(COMPUTED_COLUMN_META_KEY.to_string(), "true".to_string()),
(KIND_META_KEY.to_string(), SQL_KIND.to_string()),
(EXPRESSION_META_KEY.to_string(), expression.to_string()),
(
INPUTS_META_KEY.to_string(),
serde_json::to_string(inputs).unwrap_or_else(|_| "[]".to_string()),
),
])
}
pub fn function_computed_column_metadata(
binding_id: &str,
output_ordinal: u32,
inputs: &[String],
) -> HashMap<String, String> {
HashMap::from([
(COMPUTED_COLUMN_META_KEY.to_string(), "true".to_string()),
(KIND_META_KEY.to_string(), FUNCTION_KIND.to_string()),
(
FUNCTION_BINDING_ID_META_KEY.to_string(),
binding_id.to_string(),
),
(
FUNCTION_OUTPUT_ORDINAL_META_KEY.to_string(),
output_ordinal.to_string(),
),
(
INPUTS_META_KEY.to_string(),
serde_json::to_string(inputs).unwrap_or_else(|_| "[]".to_string()),
),
])
}
#[derive(Debug, Serialize, Deserialize)]
struct FunctionBindingEnvelope {
version: u32,
bindings: Vec<Value>,
}
pub fn function_bindings_metadata(bindings: &[FunctionBinding]) -> Result<String> {
let bindings = bindings
.iter()
.map(serde_json::to_value)
.collect::<std::result::Result<Vec<_>, _>>()
.map_err(|e| Error::InvalidInput {
message: format!("invalid Function binding metadata: {e}"),
})?;
serde_json::to_string(&FunctionBindingEnvelope {
version: FUNCTION_BINDINGS_VERSION,
bindings,
})
.map_err(|e| Error::InvalidInput {
message: format!("invalid Function binding metadata: {e}"),
})
}
pub fn function_bindings(schema: &ArrowSchema) -> Result<Vec<FunctionBinding>> {
let Some(envelope) = function_binding_envelope(schema)? else {
return Ok(Vec::new());
};
envelope
.bindings
.into_iter()
.map(|binding| {
serde_json::from_value(binding).map_err(|e| Error::InvalidInput {
message: format!("invalid Function binding metadata: {e}"),
})
})
.collect()
}
fn function_binding_envelope(schema: &ArrowSchema) -> Result<Option<FunctionBindingEnvelope>> {
let Some(raw) = schema.metadata().get(FUNCTION_BINDINGS_META_KEY) else {
return Ok(None);
};
let envelope: FunctionBindingEnvelope =
serde_json::from_str(raw).map_err(|e| Error::InvalidInput {
message: format!("invalid Function binding metadata: {e}"),
})?;
if envelope.version != FUNCTION_BINDINGS_VERSION {
return Err(Error::NotSupported {
message: format!(
"Function binding metadata version {} is not supported by this client",
envelope.version
),
});
}
Ok(Some(envelope))
}
pub(crate) fn ensure_supported_function_metadata(schema: &ArrowSchema) -> Result<()> {
let raw_bindings = function_binding_envelope(schema)?
.map(|envelope| envelope.bindings)
.unwrap_or_default();
for value in &raw_bindings {
ensure_known_binding_shape(value)?;
}
let bindings = raw_bindings
.into_iter()
.map(|binding| {
serde_json::from_value(binding).map_err(|e| Error::InvalidInput {
message: format!("invalid Function binding metadata: {e}"),
})
})
.collect::<Result<Vec<FunctionBinding>>>()?;
let mut binding_ids = BTreeSet::new();
for binding in &bindings {
if !binding_ids.insert(binding.binding_id().to_string()) {
return Err(Error::InvalidInput {
message: format!("duplicate Function binding '{}'", binding.binding_id()),
});
}
if binding.outputs().is_empty() {
return Err(Error::InvalidInput {
message: format!("Function binding '{}' has no outputs", binding.binding_id()),
});
}
if binding.function().name.is_empty() || binding.function().version.is_empty() {
return Err(Error::InvalidInput {
message: format!(
"Function binding '{}' has no exact version",
binding.binding_id()
),
});
}
if binding.input_schema().is_none() || binding.output_schema().is_none() {
return Err(Error::NotSupported {
message: format!(
"Function binding '{}' does not contain exact Arrow schemas",
binding.binding_id()
),
});
}
for (ordinal, output) in binding.outputs().iter().enumerate() {
if output.output_ordinal != ordinal as u32 {
return Err(Error::InvalidInput {
message: format!(
"Function binding '{}' has non-canonical output ordinals",
binding.binding_id()
),
});
}
}
ensure_binding_matches_schema(schema, binding)?;
}
let bindings_by_id = bindings
.iter()
.map(|binding| (binding.binding_id(), binding))
.collect::<HashMap<_, _>>();
for field in schema.fields() {
if field
.metadata()
.get(COMPUTED_COLUMN_META_KEY)
.map(String::as_str)
!= Some("true")
{
continue;
}
match computed_column_from_field(field) {
Some(ComputedColumn {
kind:
ComputedColumnKind::Function {
binding_id,
output_ordinal,
},
..
}) => {
let binding =
bindings_by_id
.get(binding_id.as_str())
.ok_or_else(|| Error::InvalidInput {
message: format!(
"Function output '{}' references missing binding '{}'",
field.name(),
binding_id
),
})?;
let output = binding
.outputs()
.get(output_ordinal as usize)
.ok_or_else(|| Error::InvalidInput {
message: format!(
"Function output '{}' has invalid ordinal {}",
field.name(),
output_ordinal
),
})?;
if output.output_name != field.name().as_str() {
return Err(Error::InvalidInput {
message: format!(
"Function output '{}' does not match binding destination '{}'",
field.name(),
output.output_name
),
});
}
}
Some(ComputedColumn {
kind: ComputedColumnKind::Sql { .. },
..
}) => {}
Some(ComputedColumn {
kind: ComputedColumnKind::Unrecognized { kind },
..
}) => {
return Err(Error::NotSupported {
message: format!(
"computed column '{}' uses unsupported kind '{}'",
field.name(),
kind
),
});
}
None => {
return Err(Error::InvalidInput {
message: format!(
"computed column '{}' has incomplete declaration metadata",
field.name()
),
});
}
}
}
Ok(())
}
pub(crate) fn ensure_no_function_bindings_for_mutation(
schema: &ArrowSchema,
operation: &str,
) -> Result<()> {
ensure_supported_function_metadata(schema)?;
if !function_bindings(schema)?.is_empty() {
return Err(Error::NotSupported {
message: format!(
"{operation} is not supported on a table with registered Function bindings"
),
});
}
Ok(())
}
pub fn computed_column_from_field(field: &ArrowField) -> Option<ComputedColumn> {
let metadata = field.metadata();
if metadata.get(COMPUTED_COLUMN_META_KEY).map(String::as_str) != Some("true") {
return None;
}
let kind = match metadata.get(KIND_META_KEY)?.as_str() {
SQL_KIND => ComputedColumnKind::Sql {
expression: metadata.get(EXPRESSION_META_KEY)?.clone(),
},
FUNCTION_KIND => match (
metadata.get(FUNCTION_BINDING_ID_META_KEY),
metadata
.get(FUNCTION_OUTPUT_ORDINAL_META_KEY)
.and_then(|value| value.parse::<u32>().ok()),
) {
(Some(binding_id), Some(output_ordinal)) if !binding_id.is_empty() => {
ComputedColumnKind::Function {
binding_id: binding_id.clone(),
output_ordinal,
}
}
_ => ComputedColumnKind::Unrecognized {
kind: FUNCTION_KIND.to_string(),
},
},
other => ComputedColumnKind::Unrecognized {
kind: other.to_string(),
},
};
let inputs = metadata
.get(INPUTS_META_KEY)
.and_then(|raw| serde_json::from_str::<Vec<String>>(raw).ok())
.unwrap_or_default();
Some(ComputedColumn {
name: field.name().clone(),
kind,
inputs,
})
}
pub fn computed_columns(schema: &ArrowSchema) -> Vec<ComputedColumn> {
schema
.fields()
.iter()
.filter_map(|field| computed_column_from_field(field))
.collect()
}
#[derive(Debug, Clone, Serialize)]
pub(crate) struct FunctionOutputTarget {
pub result_field: String,
pub output_name: String,
pub output_ordinal: u32,
}
#[derive(Debug, Clone, Serialize)]
pub(crate) struct FunctionInputTarget {
pub parameter: String,
pub field_path: String,
pub arrow_type: String,
pub nullable: bool,
}
#[derive(Debug, Clone, Serialize)]
pub(crate) struct FunctionDeclarationPlan {
pub application: FunctionApplication,
pub binding_metadata_version: u32,
pub input_bindings: Vec<FunctionInputTarget>,
pub input_schema: JsonArrowSchema,
pub output_schema: JsonArrowSchema,
pub outputs: Vec<FunctionOutputTarget>,
}
fn invalid_function(message: impl Into<String>) -> Error {
Error::InvalidInput {
message: message.into(),
}
}
fn reject_unknown_object_fields(value: &Value, allowed: &[&str], context: &str) -> Result<()> {
let object = value.as_object().ok_or_else(|| {
invalid_function(format!(
"invalid Function binding metadata: {context} must be an object"
))
})?;
let unknown = object
.keys()
.filter(|key| !allowed.contains(&key.as_str()))
.cloned()
.collect::<Vec<_>>();
if unknown.is_empty() {
Ok(())
} else {
Err(Error::NotSupported {
message: format!(
"Function binding metadata contains newer {context} fields: {unknown:?}"
),
})
}
}
fn ensure_known_binding_shape(value: &Value) -> Result<()> {
reject_unknown_object_fields(
value,
&[
"binding_id",
"function",
"inputs",
"outputs",
"input_schema",
"output_schema",
],
"binding",
)?;
let object = value.as_object().unwrap();
reject_unknown_object_fields(
object
.get("function")
.ok_or_else(|| invalid_function("Function binding is missing its exact version"))?,
&["name", "version"],
"version reference",
)?;
for input in object
.get("inputs")
.and_then(Value::as_array)
.ok_or_else(|| invalid_function("Function binding inputs must be an array"))?
{
reject_unknown_object_fields(
input,
&[
"parameter",
"field_id",
"field_path",
"arrow_type",
"nullable",
],
"input binding",
)?;
}
for output in object
.get("outputs")
.and_then(Value::as_array)
.ok_or_else(|| invalid_function("Function binding outputs must be an array"))?
{
reject_unknown_object_fields(
output,
&[
"result_field",
"output_name",
"output_field_id",
"output_ordinal",
"arrow_type",
"nullable",
],
"output mapping",
)?;
}
Ok(())
}
struct ResolvedFieldPath<'a> {
root: &'a ArrowField,
leaf: &'a ArrowField,
}
fn resolve_field_path<'a>(schema: &'a ArrowSchema, path: &str) -> Result<ResolvedFieldPath<'a>> {
let parts = lance_core::datatypes::parse_field_path(path).map_err(|e| {
invalid_function(format!("invalid Function input field path '{path}': {e}"))
})?;
let Some((root, children)) = parts.split_first() else {
return Err(invalid_function(
"Function input field path cannot be empty",
));
};
let root = schema
.field_with_name(root)
.map_err(|_| invalid_function(format!("unknown Function input column '{path}'")))?;
let mut leaf = root;
for child in children {
let DataType::Struct(fields) = leaf.data_type() else {
return Err(invalid_function(format!(
"Function input field path '{path}' traverses a non-struct field"
)));
};
leaf = fields
.iter()
.find(|field| field.name() == child)
.map(AsRef::as_ref)
.ok_or_else(|| invalid_function(format!("unknown Function input column '{path}'")))?;
}
Ok(ResolvedFieldPath { root, leaf })
}
fn canonical_input_arrow_type(field: &JsonArrowField) -> Result<String> {
let is_blob_v2 = field
.metadata
.as_ref()
.and_then(|metadata| metadata.get(ARROW_EXT_NAME_KEY))
.map(String::as_str)
== Some(BLOB_V2_EXT_NAME);
if is_blob_v2 {
let arrow_field = lance_namespace::schema::convert_json_arrow_field(field)
.map_err(|e| invalid_function(format!("invalid Function input field: {e}")))?;
if !has_supported_blob_v2_layout(&arrow_field) {
return Err(invalid_function(format!(
"Function input '{}' has an invalid Blob v2 storage layout",
arrow_field.name()
)));
}
return Ok(FUNCTION_BLOB_V2_TYPE.to_string());
}
if field.r#type.fields.is_none() && field.r#type.length.is_none() {
Ok(field.r#type.r#type.clone())
} else {
serde_json::to_string(field.r#type.as_ref()).map_err(|e| {
invalid_function(format!("could not encode exact Function input type: {e}"))
})
}
}
fn has_supported_blob_v2_layout(field: &ArrowField) -> bool {
field.is_blob_v2()
&& matches!(
field.data_type(),
DataType::Struct(fields) if BlobV2Layout::classify(fields).is_some()
)
}
fn split_fixed_size_list(raw: &str) -> Option<(&str, i32)> {
let inner = raw.strip_prefix("fixed_size_list<")?.strip_suffix('>')?;
let mut depth = 0_u32;
let mut separator = None;
for (index, byte) in inner.bytes().enumerate() {
match byte {
b'<' => depth += 1,
b'>' => depth = depth.checked_sub(1)?,
b',' if depth == 0 => separator = Some(index),
_ => {}
}
}
let (item, size) = inner.split_at(separator?);
let size: i32 = size[1..].trim().parse().ok()?;
(size > 0).then_some((item.trim(), size))
}
fn parse_output_arrow_type(raw: &str) -> Result<JsonArrowDataType> {
fn parse(raw: &str) -> Result<JsonArrowDataType> {
let raw = raw.trim();
if raw.starts_with('{') {
return serde_json::from_str(raw).map_err(|e| {
invalid_function(format!("invalid Function Arrow type '{raw}': {e}"))
});
}
if let Some(inner) = raw
.strip_prefix("list<")
.and_then(|value| value.strip_suffix('>'))
{
let mut data_type = JsonArrowDataType::new("list".to_string());
data_type.fields = Some(vec![JsonArrowField::new(
"item".to_string(),
false,
parse(inner)?,
)]);
return Ok(data_type);
}
if let Some(inner) = raw
.strip_prefix("large_list<")
.and_then(|value| value.strip_suffix('>'))
{
let mut data_type = JsonArrowDataType::new("large_list".to_string());
data_type.fields = Some(vec![JsonArrowField::new(
"item".to_string(),
false,
parse(inner)?,
)]);
return Ok(data_type);
}
if let Some((inner, size)) = split_fixed_size_list(raw) {
let mut data_type = JsonArrowDataType::new("fixed_size_list".to_string());
data_type.fields = Some(vec![JsonArrowField::new(
"item".to_string(),
false,
parse(inner)?,
)]);
data_type.length = Some(i64::from(size));
return Ok(data_type);
}
let normalized = match raw {
"boolean" => "bool",
"string" => "utf8",
"large_string" => "large_utf8",
"halffloat" => "float16",
"float" => "float32",
"double" => "float64",
other => other,
};
Ok(JsonArrowDataType::new(normalized.to_string()))
}
let data_type = parse(raw)?;
lance_namespace::schema::convert_json_arrow_type(&data_type)
.map_err(|e| invalid_function(format!("unsupported Function Arrow type '{raw}': {e}")))?;
Ok(data_type)
}
fn function_output_field(name: &str, nullable: bool, raw: &str) -> Result<JsonArrowField> {
if raw == FUNCTION_BLOB_V2_TYPE {
return lance_namespace::schema::arrow_schema_to_json(&ArrowSchema::new(vec![
crate::blob(name, nullable),
]))
.map_err(|e| invalid_function(format!("could not encode Blob v2 output field: {e}")))?
.fields
.into_iter()
.next()
.ok_or_else(|| invalid_function("Blob v2 output field is missing"));
}
Ok(JsonArrowField::new(
name.to_string(),
nullable,
parse_output_arrow_type(raw)?,
))
}
fn function_output_field_matches(expected: &ArrowField, actual: &ArrowField) -> bool {
expected.name() == actual.name()
&& expected.is_nullable() == actual.is_nullable()
&& if expected.is_blob_v2() {
has_supported_blob_v2_layout(expected) && has_supported_blob_v2_layout(actual)
} else {
function_output_type_matches(expected.data_type(), actual.data_type())
}
}
fn function_output_type_matches(expected: &DataType, actual: &DataType) -> bool {
if expected == actual {
return true;
}
match (expected, actual) {
(DataType::Struct(expected), DataType::Struct(actual)) => {
expected.len() == actual.len()
&& expected
.iter()
.zip(actual)
.all(|(expected, actual)| function_output_field_matches(expected, actual))
}
(DataType::List(expected), DataType::List(actual))
| (DataType::LargeList(expected), DataType::LargeList(actual)) => {
function_output_field_matches(expected, actual)
}
(
DataType::FixedSizeList(expected, expected_size),
DataType::FixedSizeList(actual, actual_size),
) => expected_size == actual_size && function_output_field_matches(expected, actual),
(DataType::Map(expected, expected_sorted), DataType::Map(actual, actual_sorted)) => {
expected_sorted == actual_sorted && function_output_field_matches(expected, actual)
}
_ => false,
}
}
fn function_output_type_has_blob(data_type: &DataType) -> bool {
match data_type {
DataType::Struct(fields) => fields
.iter()
.any(|field| field.is_blob_v2() || function_output_type_has_blob(field.data_type())),
DataType::List(field)
| DataType::LargeList(field)
| DataType::FixedSizeList(field, _)
| DataType::Map(field, _) => {
field.is_blob_v2() || function_output_type_has_blob(field.data_type())
}
_ => false,
}
}
fn ensure_binding_matches_schema(schema: &ArrowSchema, binding: &FunctionBinding) -> Result<()> {
let mut input_fields = Vec::with_capacity(binding.inputs().len());
for input in binding.inputs() {
let resolved = resolve_field_path(schema, &input.field_path)?;
let field = resolved.leaf;
if field
.metadata()
.get(COMPUTED_COLUMN_META_KEY)
.map(String::as_str)
== Some("true")
{
return Err(invalid_function(format!(
"Function input '{}' is computed",
input.field_path
)));
}
if field.is_nullable() && !input.nullable {
return Err(invalid_function(format!(
"Function input column '{}' is nullable, but parameter '{}' in binding '{}' is non-nullable",
input.field_path,
input.parameter,
binding.binding_id()
)));
}
let parameter_field = ArrowField::new(
input.parameter.clone(),
field.data_type().clone(),
input.nullable,
)
.with_metadata(field.metadata().clone());
let json = lance_namespace::schema::arrow_schema_to_json(&ArrowSchema::new(vec![
parameter_field.clone(),
]))
.map_err(|e| invalid_function(format!("invalid Function input schema: {e}")))?;
let json_field = json.fields.into_iter().next().unwrap();
if canonical_input_arrow_type(&json_field)? != input.arrow_type {
return Err(invalid_function(format!(
"Function input '{}' type no longer matches binding '{}'",
input.field_path,
binding.binding_id()
)));
}
input_fields.push(parameter_field);
}
let input_schema =
lance_namespace::schema::arrow_schema_to_json(&ArrowSchema::new(input_fields))
.map_err(|e| invalid_function(format!("invalid Function input schema: {e}")))?;
let input_schema = serde_json::to_value(input_schema).map_err(|e| {
invalid_function(format!("could not encode exact Function input schema: {e}"))
})?;
if binding.input_schema() != Some(&input_schema) {
return Err(invalid_function(format!(
"Function binding '{}' input schema does not match its inputs",
binding.binding_id()
)));
}
let expected_inputs = binding
.inputs()
.iter()
.map(|input| input.field_path.clone())
.collect::<Vec<_>>();
let mut output_fields = Vec::with_capacity(binding.outputs().len());
for output in binding.outputs() {
let field = schema.field_with_name(&output.output_name).map_err(|_| {
invalid_function(format!(
"Function binding '{}' output '{}' is missing",
binding.binding_id(),
output.output_name
))
})?;
if field.name() != &output.output_name || !field.is_nullable() || output.nullable {
return Err(invalid_function(format!(
"Function output '{}' no longer matches binding '{}'",
output.output_name,
binding.binding_id()
)));
}
let (type_matches, has_semantic_blob) = if output.arrow_type == FUNCTION_BLOB_V2_TYPE {
(has_supported_blob_v2_layout(field), true)
} else {
let expected_type = parse_output_arrow_type(&output.arrow_type)?;
let expected_type = lance_namespace::schema::convert_json_arrow_type(&expected_type)
.map_err(|e| invalid_function(format!("invalid Function output type: {e}")))?;
(
function_output_type_matches(&expected_type, field.data_type()),
function_output_type_has_blob(&expected_type),
)
};
if !type_matches {
return Err(invalid_function(format!(
"Function output '{}' type no longer matches binding '{}'",
output.output_name,
binding.binding_id()
)));
}
let metadata = field.metadata();
let declared_inputs = metadata
.get(INPUTS_META_KEY)
.and_then(|raw| serde_json::from_str::<Vec<String>>(raw).ok());
if metadata.get(COMPUTED_COLUMN_META_KEY).map(String::as_str) != Some("true")
|| metadata.get(KIND_META_KEY).map(String::as_str) != Some(FUNCTION_KIND)
|| metadata
.get(FUNCTION_BINDING_ID_META_KEY)
.map(String::as_str)
!= Some(binding.binding_id())
|| metadata
.get(FUNCTION_OUTPUT_ORDINAL_META_KEY)
.and_then(|value| value.parse::<u32>().ok())
!= Some(output.output_ordinal)
|| declared_inputs.as_deref() != Some(expected_inputs.as_slice())
{
return Err(invalid_function(format!(
"Function output '{}' declaration metadata does not match binding '{}'",
output.output_name,
binding.binding_id()
)));
}
if has_semantic_blob {
output_fields.push(function_output_field(
field.name(),
true,
&output.arrow_type,
)?);
} else {
let json = lance_namespace::schema::arrow_schema_to_json(&ArrowSchema::new(vec![
ArrowField::new(field.name().clone(), field.data_type().clone(), true),
]))
.map_err(|e| invalid_function(format!("invalid Function output schema: {e}")))?;
output_fields.push(json.fields.into_iter().next().unwrap());
}
}
let output_schema = JsonArrowSchema::new(output_fields);
let output_schema = serde_json::to_value(output_schema).map_err(|e| {
invalid_function(format!(
"could not encode exact Function output schema: {e}"
))
})?;
if binding.output_schema() != Some(&output_schema) {
return Err(invalid_function(format!(
"Function binding '{}' output schema does not match physical siblings",
binding.binding_id()
)));
}
Ok(())
}
pub(crate) fn plan_function_application(
schema: &ArrowSchema,
application: &FunctionApplication,
output_name: Option<&str>,
) -> Result<FunctionDeclarationPlan> {
ensure_supported_function_metadata(schema)?;
if application.has_unknown_fields() {
return Err(Error::NotSupported {
message: "Function application contains fields from a newer contract".into(),
});
}
if application.function().name.is_empty() || application.function().version.is_empty() {
return Err(invalid_function(
"Function application requires an exact version",
));
}
let mut parameters = BTreeSet::new();
let mut input_bindings = Vec::with_capacity(application.inputs().len());
let mut input_fields = Vec::with_capacity(application.inputs().len());
for input in application.inputs() {
if !parameters.insert(input.parameter.as_str()) {
return Err(invalid_function(format!(
"duplicate Function parameter '{}'",
input.parameter
)));
}
if input.kind != "column" {
return Err(Error::NotSupported {
message: format!(
"Function input kind '{}' is not supported for column declaration",
input.kind
),
});
}
let source = input.value.as_object().ok_or_else(|| {
invalid_function(format!(
"Function parameter '{}' has an invalid column source",
input.parameter
))
})?;
if source.len() != 1 {
return Err(Error::NotSupported {
message: format!(
"Function parameter '{}' uses a newer column source contract",
input.parameter
),
});
}
let path = source.get("path").and_then(Value::as_str).ok_or_else(|| {
invalid_function(format!(
"Function parameter '{}' requires a column path",
input.parameter
))
})?;
let resolved = resolve_field_path(schema, path)?;
if resolved
.root
.metadata()
.get(COMPUTED_COLUMN_META_KEY)
.map(String::as_str)
== Some("true")
{
return Err(invalid_function(format!(
"Function input '{path}' is computed; computed-on-computed bindings are not supported"
)));
}
let field = resolved.leaf;
let parameter_field = ArrowField::new(
input.parameter.clone(),
field.data_type().clone(),
field.is_nullable(),
)
.with_metadata(field.metadata().clone());
let input_schema = lance_namespace::schema::arrow_schema_to_json(&ArrowSchema::new(vec![
parameter_field.clone(),
]))
.map_err(|e| invalid_function(format!("invalid Function input schema: {e}")))?;
let json_field = input_schema.fields.into_iter().next().unwrap();
input_bindings.push(FunctionInputTarget {
parameter: input.parameter.clone(),
field_path: path.to_string(),
arrow_type: canonical_input_arrow_type(&json_field)?,
nullable: field.is_nullable(),
});
input_fields.push(parameter_field);
}
let input_schema =
lance_namespace::schema::arrow_schema_to_json(&ArrowSchema::new(input_fields))
.map_err(|e| invalid_function(format!("invalid Function input schema: {e}")))?;
let output = application.output();
let mut outputs = Vec::new();
let mut output_fields = Vec::new();
match output.kind.as_str() {
"scalar" => {
if !application.columns().is_empty() {
return Err(invalid_function(
"scalar Function applications cannot rename result fields",
));
}
let name = output_name.ok_or_else(|| {
invalid_function(
"a scalar Function application must be mapped to one output column",
)
})?;
if output.nullable != Some(false) {
return Err(invalid_function(
"Function logical outputs must be non-nullable during NULL assignment",
));
}
let arrow_type = output.arrow_type.as_deref().ok_or_else(|| {
invalid_function("scalar Function output is missing its Arrow type")
})?;
outputs.push(FunctionOutputTarget {
result_field: WHOLE_RESULT_FIELD.to_string(),
output_name: name.to_string(),
output_ordinal: 0,
});
output_fields.push(function_output_field(name, true, arrow_type)?);
}
"named_struct" => {
if output.fields.is_empty() {
return Err(invalid_function(
"named-struct Function output requires at least one field",
));
}
let result_names = output
.fields
.iter()
.map(|field| field.name.as_str())
.collect::<BTreeSet<_>>();
if result_names.len() != output.fields.len() {
return Err(invalid_function(
"named-struct Function result field names must be unique",
));
}
if output.fields.iter().any(|field| field.nullable) {
return Err(invalid_function(
"Function logical outputs must be non-nullable during NULL assignment",
));
}
let unknown = application
.columns()
.keys()
.filter(|name| !result_names.contains(name.as_str()))
.cloned()
.collect::<Vec<_>>();
if !unknown.is_empty() {
return Err(invalid_function(format!(
"unknown Function result fields: {unknown:?}"
)));
}
if let Some(name) = output_name {
if !application.columns().is_empty() {
return Err(invalid_function(
"a named-struct mapped to one column cannot also rename expanded fields",
));
}
let fields = output
.fields
.iter()
.map(|field| function_output_field(&field.name, false, &field.arrow_type))
.collect::<Result<Vec<_>>>()?;
let mut data_type = JsonArrowDataType::new("struct".to_string());
data_type.fields = Some(fields);
outputs.push(FunctionOutputTarget {
result_field: WHOLE_RESULT_FIELD.to_string(),
output_name: name.to_string(),
output_ordinal: 0,
});
output_fields.push(JsonArrowField::new(name.to_string(), true, data_type));
} else {
let mut destinations = BTreeSet::new();
for (ordinal, field) in output.fields.iter().enumerate() {
let name = application
.columns()
.get(&field.name)
.unwrap_or(&field.name);
if !destinations.insert(name.as_str()) {
return Err(invalid_function(
"Function output destinations must be unique",
));
}
outputs.push(FunctionOutputTarget {
result_field: field.name.clone(),
output_name: name.clone(),
output_ordinal: ordinal as u32,
});
output_fields.push(function_output_field(name, true, &field.arrow_type)?);
}
}
}
kind => {
return Err(Error::NotSupported {
message: format!(
"Function output kind '{kind}' is not supported for column declaration"
),
});
}
}
for output in &outputs {
if output.output_name.is_empty() {
return Err(invalid_function(
"Function output column name cannot be empty",
));
}
if schema.field_with_name(&output.output_name).is_ok() {
return Err(Error::ColumnAlreadyExists {
name: output.output_name.clone(),
});
}
}
Ok(FunctionDeclarationPlan {
application: application.clone(),
binding_metadata_version: FUNCTION_BINDINGS_VERSION,
input_bindings,
input_schema,
output_schema: JsonArrowSchema::new(output_fields),
outputs,
})
}
pub(crate) fn ensure_not_an_input(schema: &SchemaRef, paths: &[&str]) -> Result<()> {
for declaration in computed_columns(schema) {
let inputs = match &declaration.kind {
ComputedColumnKind::Sql { expression } => Planner::new(schema.clone())
.parse_expr(expression)
.map(|parsed| Planner::column_names_in_expr(&parsed))
.map_err(|e| Error::InvalidInput {
message: format!(
"computed column '{}' has an unevaluable expression ({e}); drop it \
before changing the schema",
declaration.name
),
})?,
_ => declaration.inputs.clone(),
};
for path in paths {
if declaration.name == *path {
continue;
}
if declaration.name == root(path) {
return Err(Error::InvalidInput {
message: format!(
"'{}' is part of computed column '{}'; drop the column and declare \
it again",
path, declaration.name
),
});
}
if inputs.iter().any(|input| root(input) == root(path)) {
return Err(Error::InvalidInput {
message: format!(
"column '{}' is read by computed column '{}'; drop that column first",
path, declaration.name
),
});
}
}
}
Ok(())
}
pub(crate) fn ensure_not_written<'a>(
schema: &ArrowSchema,
written: impl IntoIterator<Item = &'a str>,
) -> Result<()> {
let declared: Vec<String> = computed_columns(schema)
.into_iter()
.map(|declaration| declaration.name)
.collect();
for name in written {
if declared.iter().any(|declared| declared == root(name)) {
return Err(Error::InvalidInput {
message: format!(
"column '{}' is computed; its values come from refresh and cannot be \
written directly",
root(name)
),
});
}
}
Ok(())
}
pub(crate) fn ensure_batch_writes_no_computed_values(
declared: &[String],
batch: &arrow_array::RecordBatch,
) -> Result<()> {
for name in declared {
if let Some(column) = batch.column_by_name(name)
&& column.null_count() != column.len()
{
return Err(Error::InvalidInput {
message: format!(
"column '{name}' is computed; its values come from refresh and cannot \
be written directly"
),
});
}
}
Ok(())
}
pub(crate) fn ensure_no_foreign_declarations<'a>(
fields: impl IntoIterator<Item = &'a Arc<ArrowField>>,
) -> Result<()> {
for field in fields {
ensure_no_foreign_declaration(field)?;
}
Ok(())
}
fn ensure_no_foreign_declaration(field: &ArrowField) -> Result<()> {
if field.metadata().keys().any(|k| is_declaration_key(k)) {
return Err(Error::InvalidInput {
message: format!(
"field '{}' carries computed-column metadata; declare computed columns \
with add_columns().computed()",
field.name()
),
});
}
Ok(())
}
pub(crate) fn is_declaration_key(key: &str) -> bool {
key == COMPUTED_COLUMN_META_KEY || key.starts_with("computed_column.")
}
pub(crate) fn ensure_not_retyped(schema: &ArrowSchema, paths: &[&str]) -> Result<()> {
for declaration in computed_columns(schema) {
for path in paths {
if declaration.name == root(path) {
return Err(Error::InvalidInput {
message: format!(
"column '{}' is computed; drop it and declare it again to change \
its type",
declaration.name
),
});
}
}
}
Ok(())
}
pub(crate) fn root(path: &str) -> &str {
path.split('.').next().unwrap_or(path)
}
pub(crate) struct BoundExpression {
pub inputs: Vec<String>,
pub roots: Vec<String>,
pub physical: Arc<dyn PhysicalExpr>,
pub data_type: DataType,
pub blob_paths: Vec<String>,
projected_blob_field: Option<ArrowField>,
}
fn is_direct_field_projection(expr: &Expr) -> bool {
match expr {
Expr::Column(_) => true,
Expr::ScalarFunction(function)
if function.name() == "get_field" && function.args.len() == 2 =>
{
is_direct_field_projection(&function.args[0])
&& matches!(
&function.args[1],
Expr::Literal(ScalarValue::Utf8(Some(_)), _)
)
}
_ => false,
}
}
fn projected_blob_field(schema: &ArrowSchema, expr: &Expr) -> Result<Option<ArrowField>> {
if !is_direct_field_projection(expr) {
return Ok(None);
}
let paths = Planner::column_names_in_expr(expr);
let [path] = paths.as_slice() else {
return Ok(None);
};
let (_, field) = resolve_arrow_field_path(schema, path)?;
Ok(field.is_blob_v2().then_some(field))
}
fn collect_blob_paths(field: &ArrowField, parent: &[String], paths: &mut Vec<Vec<String>>) {
let mut path = parent.to_vec();
path.push(field.name().clone());
if field.is_blob_v2() {
paths.push(path);
return;
}
match field.data_type() {
DataType::Struct(children) => {
for child in children {
collect_blob_paths(child, &path, paths);
}
}
DataType::List(child)
| DataType::LargeList(child)
| DataType::FixedSizeList(child, _)
| DataType::Map(child, _) => collect_blob_paths(child, &path, paths),
_ => {}
}
}
fn schema_blob_paths(schema: &ArrowSchema) -> Vec<Vec<String>> {
let mut paths = Vec::new();
for field in schema.fields() {
collect_blob_paths(field, &[], &mut paths);
}
paths
}
fn transform_blob_field(
field: &ArrowField,
parent: &[String],
materialized: &HashSet<Vec<String>>,
) -> ArrowField {
let mut path = parent.to_vec();
path.push(field.name().clone());
if field.is_blob_v2() {
if materialized.contains(&path) {
return ArrowField::new(field.name(), DataType::LargeBinary, field.is_nullable());
}
return ArrowField::new(
field.name(),
BLOB_V2_DESC_FIELD.data_type().clone(),
field.is_nullable(),
)
.with_metadata(BLOB_V2_DESC_FIELD.metadata().clone());
}
let data_type = match field.data_type() {
DataType::Struct(children) => DataType::Struct(
children
.iter()
.map(|child| Arc::new(transform_blob_field(child, &path, materialized)))
.collect(),
),
DataType::List(child) => {
DataType::List(Arc::new(transform_blob_field(child, &path, materialized)))
}
DataType::LargeList(child) => {
DataType::LargeList(Arc::new(transform_blob_field(child, &path, materialized)))
}
DataType::FixedSizeList(child, size) => DataType::FixedSizeList(
Arc::new(transform_blob_field(child, &path, materialized)),
*size,
),
DataType::Map(child, sorted) => DataType::Map(
Arc::new(transform_blob_field(child, &path, materialized)),
*sorted,
),
_ => return field.clone(),
};
ArrowField::new(field.name(), data_type, field.is_nullable())
.with_metadata(field.metadata().clone())
}
fn blob_runtime_schema(schema: &ArrowSchema, materialized: &HashSet<Vec<String>>) -> SchemaRef {
Arc::new(ArrowSchema::new_with_metadata(
schema
.fields()
.iter()
.map(|field| Arc::new(transform_blob_field(field, &[], materialized)))
.collect::<Fields>(),
schema.metadata().clone(),
))
}
fn referenced_blob_paths(schema: &ArrowSchema, inputs: &[String]) -> Result<Vec<Vec<String>>> {
let input_paths = inputs
.iter()
.map(|input| {
parse_field_path(input).map_err(|error| Error::InvalidInput {
message: format!("invalid computed-column input path '{input}': {error}"),
})
})
.collect::<Result<Vec<_>>>()?;
Ok(schema_blob_paths(schema)
.into_iter()
.filter(|blob_path| {
input_paths.iter().any(|input_path| {
input_path.len() <= blob_path.len()
&& input_path
.iter()
.zip(blob_path)
.all(|(input, blob)| input == blob)
})
})
.collect())
}
pub(crate) fn bind(schema: SchemaRef, column: &str, expression: &str) -> Result<BoundExpression> {
let invalid = |message: String| Error::InvalidExpression {
column: column.to_string(),
message,
};
let all_blob_paths = schema_blob_paths(schema.as_ref())
.into_iter()
.collect::<HashSet<_>>();
let parsing_schema = blob_runtime_schema(schema.as_ref(), &all_blob_paths);
let planner = Planner::new(parsing_schema);
let parsed = planner
.parse_expr(expression)
.map_err(|e| invalid(e.to_string()))?;
let projected_blob_field = projected_blob_field(schema.as_ref(), &parsed)?;
let mut volatile = None;
parsed
.apply(|expr| {
use datafusion_common::tree_node::TreeNodeRecursion;
if let datafusion_expr::Expr::ScalarFunction(function) = expr
&& function.func.signature().volatility != datafusion_expr::Volatility::Immutable
{
volatile = Some(function.func.name().to_string());
return Ok(TreeNodeRecursion::Stop);
}
Ok(TreeNodeRecursion::Continue)
})
.map_err(|e| invalid(e.to_string()))?;
if let Some(function) = volatile {
return Err(invalid(format!(
"'{function}' is not deterministic; a computed column's expression must \
yield the same value every time it is evaluated"
)));
}
let mut inputs = Planner::column_names_in_expr(&parsed);
inputs.sort();
inputs.dedup();
let blob_paths = referenced_blob_paths(schema.as_ref(), &inputs)?;
let runtime_schema = blob_runtime_schema(
schema.as_ref(),
&blob_paths.iter().cloned().collect::<HashSet<_>>(),
);
let mut indices = Vec::with_capacity(inputs.len());
for input in &inputs {
let index = runtime_schema
.index_of(root(input))
.map_err(|_| invalid(format!("unknown column '{input}'")))?;
if !indices.contains(&index) {
indices.push(index);
}
}
indices.sort_unstable();
let read_schema = Arc::new(
runtime_schema
.project(&indices)
.map_err(|e| invalid(e.to_string()))?,
);
let roots = read_schema
.fields()
.iter()
.map(|field| field.name().clone())
.collect();
let runtime_planner = Planner::new(runtime_schema);
let optimized = runtime_planner
.optimize_expr(parsed)
.map_err(|e| invalid(e.to_string()))?;
let physical = Planner::new(read_schema.clone())
.create_physical_expr(&optimized)
.map_err(|e| invalid(e.to_string()))?;
let data_type = physical
.data_type(read_schema.as_ref())
.map_err(|e| invalid(e.to_string()))?;
Ok(BoundExpression {
inputs,
roots,
physical,
data_type,
blob_paths: blob_paths
.iter()
.map(|path| {
let segments = path.iter().map(String::as_str).collect::<Vec<_>>();
format_field_path_minimal(&segments)
})
.collect(),
projected_blob_field,
})
}
fn plan_declarations(schema: SchemaRef, columns: &[(String, String)]) -> Result<Vec<ArrowField>> {
if columns.is_empty() {
return Err(Error::InvalidInput {
message: "at least one computed column is required".into(),
});
}
let mut schema = schema;
let mut fields = Vec::with_capacity(columns.len());
for (name, expression) in columns {
if schema.field_with_name(name).is_ok() {
return Err(Error::ColumnAlreadyExists { name: name.clone() });
}
let bound = bind(schema.clone(), name, expression)?;
let computed_metadata = computed_column_metadata(expression, &bound.inputs);
let field = match bound.projected_blob_field {
Some(source) => {
let mut metadata = source.metadata().clone();
metadata.retain(|key, _| !is_declaration_key(key));
metadata.extend(computed_metadata);
source
.with_name(name)
.with_nullable(true)
.with_metadata(metadata)
}
None => ArrowField::new(name, bound.data_type, true).with_metadata(computed_metadata),
};
schema = Arc::new(ArrowSchema::new_with_metadata(
schema
.fields()
.iter()
.cloned()
.chain(std::iter::once(Arc::new(field.clone())))
.collect::<Fields>(),
schema.metadata().clone(),
));
fields.push(field);
}
Ok(fields)
}
pub(crate) fn plan(schema: SchemaRef, columns: &[(String, String)]) -> Result<Vec<ArrowField>> {
plan_declarations(schema, columns)
}
pub fn validate_declarations(schema: SchemaRef, columns: &[(String, String)]) -> Result<()> {
ensure_no_function_bindings_for_mutation(schema.as_ref(), "schema evolution")?;
plan(schema, columns).map(drop)
}
pub(crate) fn declare(
schema: SchemaRef,
columns: &[(String, String)],
) -> Result<NewColumnTransform> {
let fields = plan_declarations(schema, columns)?;
Ok(NewColumnTransform::AllNulls(Arc::new(ArrowSchema::new(
fields,
))))
}
#[cfg(test)]
pub(super) async fn add_foreign_kind(table: &crate::Table, name: &str, kind: &str) {
let field = ArrowField::new(name, DataType::Int32, true).with_metadata(HashMap::from([
(COMPUTED_COLUMN_META_KEY.to_string(), "true".to_string()),
(KIND_META_KEY.to_string(), kind.to_string()),
(INPUTS_META_KEY.to_string(), r#"["x"]"#.to_string()),
]));
super::schema_evolution::commit_add_columns(
table.as_native().unwrap(),
NewColumnTransform::AllNulls(Arc::new(ArrowSchema::new(vec![field]))),
None,
)
.await
.unwrap();
}
#[cfg(test)]
mod tests {
#[test]
fn test_validate_declarations_matches_schema_admission_barriers() {
let schema = Arc::new(ArrowSchema::new_with_metadata(
vec![ArrowField::new("x", DataType::Int32, true)],
HashMap::from([(
FUNCTION_BINDINGS_META_KEY.to_string(),
"not valid binding metadata".to_string(),
)]),
));
let declarations = vec![("a".to_string(), "x + 1".to_string())];
assert!(super::validate_declarations(schema, &declarations).is_err());
}
#[test]
fn output_arrow_type_grammar_matches_the_shared_golden() {
let golden: serde_json::Value = serde_json::from_str(include_str!(
"../../tests/fixtures/first_class_functions/v1/arrow_types.json"
))
.unwrap();
let valid = golden["valid"].as_array().unwrap().iter();
for case in valid.chain(golden["server_only"].as_array().unwrap()) {
let raw = case["arrow_type"].as_str().unwrap();
let parsed = super::parse_output_arrow_type(raw)
.unwrap_or_else(|error| panic!("{raw}: {error}"));
assert_eq!(
serde_json::to_value(&parsed).unwrap(),
case["json"],
"{raw}"
);
}
for raw in golden["invalid"].as_array().unwrap() {
let raw = raw.as_str().unwrap();
assert!(
super::parse_output_arrow_type(raw).is_err(),
"{raw:?} should be rejected"
);
}
}
use arrow_array::record_batch;
use arrow_schema::{DataType, TimeUnit};
use futures::TryStreamExt;
use lance::dataset::ColumnAlteration;
use super::*;
use crate::connect;
use crate::query::{ExecutableQuery, QueryBase, Select};
use crate::{Error, Table};
async fn table_with_ints(name: &str) -> Table {
let conn = connect("memory://").execute().await.unwrap();
let batch = record_batch!(("x", Int32, [1, 2, 3])).unwrap();
conn.create_table(name, batch).execute().await.unwrap()
}
async fn add_computed(table: &Table, columns: &[(String, String)]) -> Result<u64> {
let mut builder = table.add_columns();
for (name, expression) in columns {
builder = builder.computed(name, expression);
}
Ok(builder.execute().await?.version)
}
async fn declared(table: &Table) -> Vec<ComputedColumn> {
computed_columns(table.schema().await.unwrap().as_ref())
}
#[tokio::test]
async fn test_declare_infers_type_and_inputs() {
let table = table_with_ints("declare_infers").await;
let initial = table.version().await.unwrap();
let version = add_computed(&table, &[("doubled".into(), "x * 2".into())])
.await
.unwrap();
assert!(version > initial);
let schema = table.schema().await.unwrap();
let field = schema.field_with_name("doubled").unwrap();
assert_eq!(field.data_type(), &DataType::Int32);
assert!(field.is_nullable());
assert_eq!(
declared(&table).await,
vec![ComputedColumn {
name: "doubled".into(),
kind: ComputedColumnKind::Sql {
expression: "x * 2".into()
},
inputs: vec!["x".into()],
}]
);
}
#[test]
fn test_direct_blob_projection_inherits_semantics() {
let schema = Arc::new(ArrowSchema::new(vec![crate::blob("image", false)]));
let fields = plan(
schema,
&[
("first".to_string(), "image".to_string()),
("second".to_string(), "first".to_string()),
],
)
.unwrap();
for field in &fields {
assert!(field.is_blob_v2());
assert!(field.is_nullable());
}
assert_eq!(
fields[1]
.metadata()
.get(EXPRESSION_META_KEY)
.map(String::as_str),
Some("first")
);
}
#[test]
fn test_blob_expression_transformation_does_not_inherit_semantics() {
let schema = Arc::new(ArrowSchema::new(vec![crate::blob("image", true)]));
let fields = plan(
schema,
&[("payload".to_string(), "coalesce(image, image)".to_string())],
)
.unwrap();
assert!(!fields[0].is_blob_v2());
assert_eq!(fields[0].data_type(), &DataType::LargeBinary);
}
#[tokio::test]
async fn test_all_nulls_preserves_field_metadata() {
let table = table_with_ints("metadata_survives").await;
add_computed(&table, &[("doubled".into(), "x * 2".into())])
.await
.unwrap();
let schema = table.schema().await.unwrap();
let metadata = schema.field_with_name("doubled").unwrap().metadata();
assert_eq!(
metadata.get(COMPUTED_COLUMN_META_KEY).map(String::as_str),
Some("true")
);
assert_eq!(metadata.get(KIND_META_KEY).map(String::as_str), Some("sql"));
assert_eq!(
metadata.get(EXPRESSION_META_KEY).map(String::as_str),
Some("x * 2")
);
assert_eq!(
metadata.get(INPUTS_META_KEY).map(String::as_str),
Some(r#"["x"]"#)
);
}
#[tokio::test]
async fn test_declared_column_is_all_null() {
let table = table_with_ints("declare_is_null").await;
add_computed(&table, &[("doubled".into(), "x * 2".into())])
.await
.unwrap();
let batches = table
.query()
.select(Select::columns(&["doubled"]))
.execute()
.await
.unwrap()
.try_collect::<Vec<_>>()
.await
.unwrap();
let total: usize = batches.iter().map(|b| b.num_rows()).sum();
assert_eq!(total, 3);
for batch in &batches {
assert_eq!(batch["doubled"].null_count(), batch.num_rows());
}
}
#[tokio::test]
async fn test_unknown_column_fails_at_declare_time() {
let table = table_with_ints("unknown_input").await;
let err = add_computed(&table, &[("bad".into(), "missing + 1".into())])
.await
.unwrap_err();
assert!(matches!(err, Error::InvalidExpression { column, .. } if column == "bad"));
let schema = table.schema().await.unwrap();
assert!(schema.field_with_name("bad").is_err());
}
#[tokio::test]
async fn test_unparsable_expression_fails_at_declare_time() {
let table = table_with_ints("bad_syntax").await;
let err = add_computed(&table, &[("bad".into(), "x *".into())])
.await
.unwrap_err();
assert!(matches!(err, Error::InvalidExpression { column, .. } if column == "bad"));
assert!(
table
.schema()
.await
.unwrap()
.field_with_name("bad")
.is_err()
);
}
#[tokio::test]
async fn test_unregistered_function_is_rejected_for_now() {
let table = table_with_ints("udf_not_yet").await;
let err = add_computed(&table, &[("vec".into(), "embed(x)".into())])
.await
.unwrap_err();
assert!(matches!(err, Error::InvalidExpression { column, .. } if column == "vec"));
assert!(
table
.schema()
.await
.unwrap()
.field_with_name("vec")
.is_err()
);
}
#[tokio::test]
async fn test_existing_column_name_is_rejected() {
let table = table_with_ints("name_taken").await;
let err = add_computed(&table, &[("x".into(), "x * 2".into())])
.await
.unwrap_err();
assert!(matches!(err, Error::ColumnAlreadyExists { name } if name == "x"));
assert!(declared(&table).await.is_empty());
}
#[tokio::test]
async fn test_constant_expression_needs_no_inputs() {
let table = table_with_ints("constant").await;
add_computed(&table, &[("answer".into(), "42".into())])
.await
.unwrap();
let declared = declared(&table).await;
assert_eq!(declared.len(), 1);
assert!(declared[0].inputs.is_empty());
}
#[tokio::test]
async fn test_multiple_columns_in_one_commit() {
let table = table_with_ints("multi").await;
let initial = table.version().await.unwrap();
add_computed(
&table,
&[
("plus".into(), "x + 1".into()),
("squared".into(), "x * x".into()),
],
)
.await
.unwrap();
assert_eq!(table.version().await.unwrap(), initial + 1);
let declared = declared(&table).await;
assert_eq!(declared.len(), 2);
assert_eq!(declared[0].name, "plus");
assert_eq!(declared[1].name, "squared");
}
#[tokio::test]
async fn test_duplicate_declaration_in_one_call_is_rejected() {
let table = table_with_ints("dupe").await;
let err = add_computed(
&table,
&[
("dup".into(), "x + 1".into()),
("dup".into(), "x + 2".into()),
],
)
.await
.unwrap_err();
assert!(matches!(err, Error::ColumnAlreadyExists { name } if name == "dup"));
assert!(declared(&table).await.is_empty());
}
#[tokio::test]
async fn test_a_declaration_may_read_one_declared_before_it() {
let table = table_with_ints("chain").await;
let before = table.version().await.unwrap();
add_computed(
&table,
&[("a".into(), "x + 1".into()), ("b".into(), "a * 2".into())],
)
.await
.unwrap();
assert_eq!(table.version().await.unwrap(), before + 1);
let declared = declared(&table).await;
assert_eq!(declared[1].name, "b");
assert_eq!(declared[1].inputs, vec!["a".to_string()]);
let err = add_computed(
&table,
&[("c".into(), "d + 1".into()), ("d".into(), "x + 1".into())],
)
.await
.unwrap_err();
assert!(matches!(err, Error::InvalidExpression { column, .. } if column == "c"));
assert!(
validate_declarations(
table.schema().await.unwrap(),
&[("e".into(), "random()".into())]
)
.is_err()
);
}
#[tokio::test]
async fn test_ordinary_columns_are_not_reported_as_computed() {
let table = table_with_ints("plain").await;
assert!(declared(&table).await.is_empty());
table
.add_columns()
.transform(NewColumnTransform::SqlExpressions(vec![(
"eager".into(),
"x * 2".into(),
)]))
.execute()
.await
.unwrap();
assert!(declared(&table).await.is_empty());
}
#[tokio::test]
async fn test_builtin_function_inference() {
let conn = connect("memory://").execute().await.unwrap();
let batch = record_batch!(("name", Utf8, ["ada", "grace"]), ("n", Int32, [-1, 2])).unwrap();
let table = conn
.create_table("builtins", batch)
.execute()
.await
.unwrap();
add_computed(
&table,
&[
("shout".into(), "upper(name)".into()),
("width".into(), "length(name)".into()),
("magnitude".into(), "abs(n)".into()),
],
)
.await
.unwrap();
let schema = table.schema().await.unwrap();
assert_eq!(
schema.field_with_name("shout").unwrap().data_type(),
&DataType::Utf8
);
assert_eq!(
schema.field_with_name("magnitude").unwrap().data_type(),
&DataType::Int32
);
assert!(
schema
.field_with_name("width")
.unwrap()
.data_type()
.is_integer()
);
let declared = declared(&table).await;
assert_eq!(declared.len(), 3);
assert_eq!(declared[0].inputs, vec!["name".to_string()]);
assert_eq!(declared[2].inputs, vec!["n".to_string()]);
}
#[tokio::test]
async fn test_unrecognized_kind_is_reported_rather_than_hidden() {
let table = table_with_ints("foreign_kind").await;
super::add_foreign_kind(&table, "embedding", "udf").await;
assert_eq!(
declared(&table).await,
vec![ComputedColumn {
name: "embedding".into(),
kind: ComputedColumnKind::Unrecognized { kind: "udf".into() },
inputs: vec!["x".into()],
}]
);
let err = add_computed(&table, &[("embedding".into(), "x * 2".into())])
.await
.unwrap_err();
assert!(matches!(err, Error::NotSupported { .. }));
}
#[test]
fn test_flag_without_a_kind_is_not_a_declaration() {
let field =
ArrowField::new("half", DataType::Int32, true).with_metadata(HashMap::from([(
COMPUTED_COLUMN_META_KEY.to_string(),
"true".to_string(),
)]));
assert_eq!(computed_column_from_field(&field), None);
}
#[test]
fn test_sql_kind_without_an_expression_is_not_a_declaration() {
let field = ArrowField::new("half", DataType::Int32, true).with_metadata(HashMap::from([
(COMPUTED_COLUMN_META_KEY.to_string(), "true".to_string()),
(KIND_META_KEY.to_string(), SQL_KIND.to_string()),
]));
assert_eq!(computed_column_from_field(&field), None);
}
#[tokio::test]
async fn test_inputs_are_deduplicated_and_sorted() {
let conn = connect("memory://").execute().await.unwrap();
let batch = record_batch!(("b", Int32, [1, 2]), ("a", Int32, [3, 4])).unwrap();
let table = conn.create_table("dedupe", batch).execute().await.unwrap();
add_computed(&table, &[("total".into(), "b + a + b".into())])
.await
.unwrap();
assert_eq!(
declared(&table).await[0].inputs,
vec!["a".to_string(), "b".to_string()]
);
}
#[tokio::test]
async fn test_dropping_an_input_is_refused() {
let table = table_with_ints("drop_input").await;
add_computed(&table, &[("doubled".into(), "x * 2".into())])
.await
.unwrap();
let err = table.drop_columns(&["x"]).await.unwrap_err();
assert!(
matches!(&err, Error::InvalidInput { message } if message.contains("doubled")),
"{err:?}"
);
}
#[tokio::test]
async fn test_renaming_an_input_is_refused() {
let table = table_with_ints("rename_input").await;
add_computed(&table, &[("doubled".into(), "x * 2".into())])
.await
.unwrap();
let err = table
.alter_columns(&[ColumnAlteration::new("x".into()).rename("y".into())])
.await
.unwrap_err();
assert!(
matches!(&err, Error::InvalidInput { message } if message.contains("doubled")),
"{err:?}"
);
}
#[tokio::test]
async fn test_altering_an_input_nullability_is_allowed() {
let table = table_with_ints("nullable_input").await;
add_computed(&table, &[("doubled".into(), "x * 2".into())])
.await
.unwrap();
table
.alter_columns(&[ColumnAlteration::new("x".into()).set_nullable(true)])
.await
.unwrap();
}
#[tokio::test]
async fn test_a_volatile_expression_is_refused() {
let table = table_with_ints("volatile_expr").await;
let err = add_computed(&table, &[("maybe".into(), "random() < 0.5".into())])
.await
.unwrap_err();
assert!(
matches!(&err, Error::InvalidExpression { message, .. }
if message.contains("random") && message.contains("deterministic")),
"{err:?}"
);
}
#[tokio::test]
async fn test_inputs_survive_expression_optimization() {
let table = table_with_ints("optimized_inputs").await;
add_computed(&table, &[("flag".into(), "true OR x > 0".into())])
.await
.unwrap();
assert_eq!(declared(&table).await[0].inputs, vec!["x".to_string()]);
let err = table.drop_columns(&["x"]).await.unwrap_err();
assert!(
matches!(&err, Error::InvalidInput { message } if message.contains("flag")),
"{err:?}"
);
}
#[tokio::test]
async fn test_retyping_the_computed_column_is_refused() {
use arrow_schema::DataType as ArrowDataType;
let table = table_with_ints("retype_computed").await;
add_computed(&table, &[("doubled".into(), "x * 2".into())])
.await
.unwrap();
let err = table
.alter_columns(&[ColumnAlteration::new("doubled".into()).cast_to(ArrowDataType::Int64)])
.await
.unwrap_err();
assert!(
matches!(&err, Error::InvalidInput { message } if message.contains("computed")),
"{err:?}"
);
table.refresh_column("doubled").await.unwrap();
}
#[tokio::test]
async fn test_declaration_metadata_is_immutable() {
use crate::table::FieldMetadataUpdate;
let table = table_with_ints("metadata_tamper").await;
add_computed(&table, &[("doubled".into(), "x * 2".into())])
.await
.unwrap();
let err = table
.update_field_metadata(&[
FieldMetadataUpdate::new("doubled").set(EXPRESSION_META_KEY, "x * 3")
])
.await
.unwrap_err();
assert!(matches!(err, Error::InvalidInput { .. }), "{err:?}");
let err = table
.update_field_metadata(&[FieldMetadataUpdate::new("x")
.set(COMPUTED_COLUMN_META_KEY, "true")
.set(KIND_META_KEY, SQL_KIND)
.set(EXPRESSION_META_KEY, "x")])
.await
.unwrap_err();
assert!(matches!(err, Error::InvalidInput { .. }), "{err:?}");
let err = table
.update_field_metadata(&[FieldMetadataUpdate::new("doubled")
.set("note", "hi")
.replace()])
.await
.unwrap_err();
assert!(matches!(err, Error::InvalidInput { .. }), "{err:?}");
table
.update_field_metadata(&[FieldMetadataUpdate::new("doubled").set("note", "hi")])
.await
.unwrap();
table.refresh_column("doubled").await.unwrap();
}
#[tokio::test]
async fn test_a_computed_column_cannot_be_written_directly() {
let table = table_with_ints("direct_write").await;
add_computed(&table, &[("doubled".into(), "x * 2".into())])
.await
.unwrap();
let batch = record_batch!(("x", Int32, [4]), ("doubled", Int32, [999])).unwrap();
let err = table.add(batch.clone()).execute().await.unwrap_err();
assert!(
matches!(&err, Error::InvalidInput { message } if message.contains("refresh")),
"{err:?}"
);
let err = table
.update()
.column("doubled", "999")
.execute()
.await
.unwrap_err();
assert!(matches!(err, Error::InvalidInput { .. }));
let mut merge = table.merge_insert(&["x"]);
merge
.when_matched_update_all(None)
.when_not_matched_insert_all();
let err = merge
.execute(Box::new(arrow_array::RecordBatchIterator::new(
vec![Ok(batch.clone())],
batch.schema(),
)))
.await
.unwrap_err();
assert!(matches!(err, Error::InvalidInput { .. }));
let plain = record_batch!(("x", Int32, [4])).unwrap();
table.add(plain).execute().await.unwrap();
}
#[tokio::test]
async fn test_installing_an_lsm_spec_over_computed_columns_is_refused() {
use crate::table::LsmWriteSpec;
let tmp_dir = tempfile::tempdir().unwrap();
let conn = connect(tmp_dir.path().to_str().unwrap())
.execute()
.await
.unwrap();
let schema = Arc::new(arrow_schema::Schema::new(vec![arrow_schema::Field::new(
"x",
DataType::Int32,
false,
)]));
let batch = arrow_array::RecordBatch::try_new(
schema,
vec![Arc::new(arrow_array::Int32Array::from(vec![1, 2])) as _],
)
.unwrap();
let table = conn
.create_table("lsm_after", batch)
.execute()
.await
.unwrap();
table.set_unenforced_primary_key(["x"]).await.unwrap();
add_computed(&table, &[("doubled".into(), "x * 2".into())])
.await
.unwrap();
let err = table
.set_lsm_write_spec(LsmWriteSpec::unsharded())
.await
.unwrap_err();
assert!(
matches!(&err, Error::NotSupported { message } if message.contains("computed")),
"{err:?}"
);
assert!(table.get_lsm_write_spec().await.unwrap().is_none());
}
#[tokio::test]
async fn test_forged_declaration_metadata_is_rejected() {
let table = table_with_ints("forged_metadata").await;
let field =
ArrowField::new("doubled", DataType::Int32, true).with_metadata(HashMap::from([
(COMPUTED_COLUMN_META_KEY.to_string(), "true".to_string()),
(KIND_META_KEY.to_string(), SQL_KIND.to_string()),
(EXPRESSION_META_KEY.to_string(), "x * 2".to_string()),
(INPUTS_META_KEY.to_string(), "[]".to_string()),
]));
let err = table
.add_columns()
.transform(NewColumnTransform::AllNulls(Arc::new(ArrowSchema::new(
vec![field],
))))
.execute()
.await
.unwrap_err();
assert!(
matches!(&err, Error::InvalidInput { message } if message.contains("computed()")),
"{err:?}"
);
assert!(declared(&table).await.is_empty());
}
#[tokio::test]
async fn test_sql_insert_cannot_write_a_computed_column() {
use datafusion::prelude::SessionContext;
let table = table_with_ints("sql_insert").await;
add_computed(&table, &[("doubled".into(), "x * 2".into())])
.await
.unwrap();
let ctx = SessionContext::new();
let provider =
crate::table::datafusion::BaseTableAdapter::try_new(table.base_table().clone())
.await
.unwrap();
ctx.register_table("t", Arc::new(provider)).unwrap();
let result = async {
ctx.sql("INSERT INTO t (x, doubled) VALUES (4, 999)")
.await?
.collect()
.await
}
.await;
let err = result.unwrap_err().to_string();
assert!(err.contains("refresh"), "{err}");
}
#[tokio::test]
async fn test_overwrite_cannot_inject_a_declaration() {
use crate::table::AddDataMode;
let table = table_with_ints("overwrite_inject").await;
let field =
ArrowField::new("doubled", DataType::Int32, true).with_metadata(HashMap::from([
(COMPUTED_COLUMN_META_KEY.to_string(), "true".to_string()),
(KIND_META_KEY.to_string(), SQL_KIND.to_string()),
(EXPRESSION_META_KEY.to_string(), "x * 2".to_string()),
]));
let schema = Arc::new(ArrowSchema::new(vec![
ArrowField::new("x", DataType::Int32, true),
field,
]));
let batch = arrow_array::RecordBatch::try_new(
schema,
vec![
Arc::new(arrow_array::Int32Array::from(vec![1])) as _,
Arc::new(arrow_array::Int32Array::from(vec![999])) as _,
],
)
.unwrap();
let err = table
.add(batch)
.mode(AddDataMode::Overwrite)
.execute()
.await
.unwrap_err();
assert!(
matches!(&err, Error::InvalidInput { message } if message.contains("declare")),
"{err:?}"
);
}
#[tokio::test]
async fn test_create_table_cannot_inject_a_declaration() {
let conn = connect("memory://").execute().await.unwrap();
let field =
ArrowField::new("doubled", DataType::Int32, true).with_metadata(HashMap::from([
(COMPUTED_COLUMN_META_KEY.to_string(), "true".to_string()),
(KIND_META_KEY.to_string(), SQL_KIND.to_string()),
(EXPRESSION_META_KEY.to_string(), "x * 2".to_string()),
]));
let schema = Arc::new(ArrowSchema::new(vec![
ArrowField::new("x", DataType::Int32, true),
field,
]));
let batch = arrow_array::RecordBatch::try_new(
schema.clone(),
vec![
Arc::new(arrow_array::Int32Array::from(vec![1])) as _,
Arc::new(arrow_array::Int32Array::from(vec![999])) as _,
],
)
.unwrap();
let err = conn
.create_table("forged_create", batch)
.execute()
.await
.unwrap_err();
assert!(
matches!(&err, Error::InvalidInput { message } if message.contains("computed()")),
"{err:?}"
);
}
#[tokio::test]
async fn test_sql_insert_omitting_computed_is_allowed() {
use datafusion::prelude::SessionContext;
let table = table_with_ints("sql_insert_omitted").await;
add_computed(&table, &[("doubled".into(), "x * 2".into())])
.await
.unwrap();
let ctx = SessionContext::new();
let provider =
crate::table::datafusion::BaseTableAdapter::try_new(table.base_table().clone())
.await
.unwrap();
ctx.register_table("t", Arc::new(provider)).unwrap();
ctx.sql("INSERT INTO t (x) VALUES (4)")
.await
.unwrap()
.collect()
.await
.unwrap();
table.checkout_latest().await.unwrap();
assert_eq!(table.count_rows(None).await.unwrap(), 4);
}
#[tokio::test]
async fn test_a_nested_computed_field_cannot_be_renamed() {
let table = table_with_ints("computed_struct_rename").await;
add_computed(&table, &[("payload".into(), "named_struct('a', x)".into())])
.await
.unwrap();
let err = table
.alter_columns(&[ColumnAlteration::new("payload.a".into()).rename("b".into())])
.await
.unwrap_err();
assert!(
matches!(&err, Error::InvalidInput { message } if message.contains("payload")),
"{err:?}"
);
}
#[tokio::test]
async fn test_stale_handles_cannot_mix_computed_and_lsm() {
use crate::table::LsmWriteSpec;
let tmp_dir = tempfile::tempdir().unwrap();
let uri = tmp_dir.path().to_str().unwrap();
let schema = Arc::new(ArrowSchema::new(vec![ArrowField::new(
"x",
DataType::Int32,
false,
)]));
let batch = arrow_array::RecordBatch::try_new(
schema.clone(),
vec![Arc::new(arrow_array::Int32Array::from(vec![1])) as _],
)
.unwrap();
let conn = connect(uri).execute().await.unwrap();
let table = conn.create_table("mix", batch).execute().await.unwrap();
table.set_unenforced_primary_key(["x"]).await.unwrap();
let stale = conn.open_table("mix").execute().await.unwrap();
add_computed(&table, &[("doubled".into(), "x * 2".into())])
.await
.unwrap();
let err = stale
.set_lsm_write_spec(LsmWriteSpec::unsharded())
.await
.unwrap_err();
assert!(matches!(err, Error::NotSupported { .. }), "install won");
let batch = arrow_array::RecordBatch::try_new(
schema,
vec![Arc::new(arrow_array::Int32Array::from(vec![1])) as _],
)
.unwrap();
let table = conn.create_table("mix2", batch).execute().await.unwrap();
table.set_unenforced_primary_key(["x"]).await.unwrap();
let stale = conn.open_table("mix2").execute().await.unwrap();
table
.set_lsm_write_spec(LsmWriteSpec::unsharded())
.await
.unwrap();
let err = add_computed(&stale, &[("doubled".into(), "x * 2".into())])
.await
.unwrap_err();
assert!(matches!(err, Error::NotSupported { .. }), "declare won");
}
#[tokio::test]
async fn test_unset_with_retained_lsm_rows_cannot_admit_a_declaration() {
use crate::table::LsmWriteSpec;
use arrow_array::{Int64Array, RecordBatchIterator};
let tmp_dir = tempfile::tempdir().unwrap();
let conn = connect(tmp_dir.path().to_str().unwrap())
.execute()
.await
.unwrap();
let schema = Arc::new(ArrowSchema::new(vec![
ArrowField::new("id", DataType::Int64, false),
ArrowField::new("value", DataType::Int64, false),
]));
let batch = arrow_array::RecordBatch::try_new(
schema.clone(),
vec![
Arc::new(Int64Array::from(vec![1, 2])) as _,
Arc::new(Int64Array::from(vec![10, 20])) as _,
],
)
.unwrap();
let table = conn
.create_table("t", batch.clone())
.execute()
.await
.unwrap();
table.set_unenforced_primary_key(["id"]).await.unwrap();
table
.set_lsm_write_spec(LsmWriteSpec::unsharded())
.await
.unwrap();
let mut merge = table.merge_insert(&["id"]);
merge
.when_matched_update_all(None)
.when_not_matched_insert_all()
.use_lsm(true);
merge
.execute(Box::new(RecordBatchIterator::new(vec![Ok(batch)], schema)))
.await
.unwrap();
table.unset_lsm_write_spec().await.unwrap();
let err = add_computed(&table, &[("doubled".into(), "value * 2".into())])
.await
.unwrap_err();
assert!(
matches!(&err, Error::NotSupported { message } if message.contains("LSM")),
"{err:?}"
);
}
#[tokio::test]
async fn test_dropping_the_computed_column_is_allowed() {
let table = table_with_ints("drop_computed").await;
add_computed(&table, &[("doubled".into(), "x * 2".into())])
.await
.unwrap();
table.drop_columns(&["doubled"]).await.unwrap();
assert!(declared(&table).await.is_empty());
}
fn function_input_schema() -> ArrowSchema {
ArrowSchema::new(vec![
ArrowField::new("title", DataType::Utf8, true),
ArrowField::new("body", DataType::Utf8, true),
])
}
fn named_struct_application(columns: &str) -> FunctionApplication {
FunctionApplication::from_json(&format!(
r#"{{
"function":{{"name":"text_features","version":"fv_exact"}},
"inputs":[
{{"parameter":"title","kind":"column","value":{{"path":"title"}}}},
{{"parameter":"body","kind":"column","value":{{"path":"body"}}}}
],
"output":{{"kind":"named_struct","fields":[
{{"name":"normalized_text","arrow_type":"utf8","nullable":false}},
{{"name":"token_count","arrow_type":"int64","nullable":false}}
]}},
"columns":{columns}
}}"#
))
.unwrap()
}
fn blob_application(output: &str) -> FunctionApplication {
FunctionApplication::from_json(&format!(
r#"{{
"function":{{"name":"blob_features","version":"fv_blob"}},
"inputs":[
{{"parameter":"image","kind":"column","value":{{"path":"image"}}}}
],
"output":{output}
}}"#
))
.unwrap()
}
fn binding_from_plan(plan: &FunctionDeclarationPlan) -> FunctionBinding {
let inputs = plan
.input_bindings
.iter()
.enumerate()
.map(|(index, input)| {
serde_json::json!({
"parameter": input.parameter,
"field_id": index,
"field_path": input.field_path,
"arrow_type": input.arrow_type,
"nullable": input.nullable,
})
})
.collect::<Vec<_>>();
let outputs = plan
.outputs
.iter()
.zip(&plan.output_schema.fields)
.enumerate()
.map(|(index, (output, field))| {
serde_json::json!({
"result_field": output.result_field,
"output_name": output.output_name,
"output_field_id": 100 + index,
"output_ordinal": output.output_ordinal,
"arrow_type": canonical_input_arrow_type(field).unwrap(),
"nullable": false,
})
})
.collect::<Vec<_>>();
FunctionBinding::from_json(
&serde_json::json!({
"binding_id": "fb_blob",
"function": plan.application.function(),
"inputs": inputs,
"outputs": outputs,
"input_schema": plan.input_schema,
"output_schema": plan.output_schema,
})
.to_string(),
)
.unwrap()
}
fn full_blob_field(name: &str, nullable: bool) -> ArrowField {
ArrowField::new(
name,
DataType::Struct(lance_core::datatypes::BLOB_V2_LOGICAL_FIELDS.clone()),
nullable,
)
.with_metadata(crate::blob(name, nullable).metadata().clone())
}
fn function_binding_schema(title_nullable: bool, body_nullable: bool) -> ArrowSchema {
ArrowSchema::new(vec![
ArrowField::new("title", DataType::Utf8, title_nullable),
ArrowField::new("body", DataType::Utf8, body_nullable),
ArrowField::new("search_text", DataType::Utf8, true),
ArrowField::new("search_token_count", DataType::Int64, true),
])
}
fn valid_function_binding_schema(
title_nullable: bool,
body_nullable: bool,
binding: &FunctionBinding,
) -> ArrowSchema {
let mut fields = function_binding_schema(title_nullable, body_nullable)
.fields()
.iter()
.map(|field| field.as_ref().clone())
.collect::<Vec<_>>();
let inputs = binding
.inputs()
.iter()
.map(|input| input.field_path.clone())
.collect::<Vec<_>>();
for output in binding.outputs() {
let index = fields
.iter()
.position(|field| field.name() == &output.output_name)
.unwrap();
fields[index] = fields[index]
.clone()
.with_metadata(function_computed_column_metadata(
binding.binding_id(),
output.output_ordinal,
&inputs,
));
}
ArrowSchema::new(fields)
}
#[test]
fn test_non_nullable_function_inputs_can_bind_to_nullable_parameters() {
let binding = FunctionBinding::from_json(include_str!(
"../../tests/fixtures/first_class_functions/v1/remote_function_binding.json"
))
.unwrap();
ensure_binding_matches_schema(
&valid_function_binding_schema(false, false, &binding),
&binding,
)
.unwrap();
}
#[test]
fn test_nullable_function_input_cannot_bind_to_non_nullable_parameter() {
let mut raw_binding: Value = serde_json::from_str(include_str!(
"../../tests/fixtures/first_class_functions/v1/remote_function_binding.json"
))
.unwrap();
raw_binding["inputs"][0]["nullable"] = Value::Bool(false);
raw_binding["input_schema"]["fields"][0]["nullable"] = Value::Bool(false);
let binding: FunctionBinding = serde_json::from_value(raw_binding).unwrap();
let err = ensure_binding_matches_schema(
&valid_function_binding_schema(true, false, &binding),
&binding,
)
.unwrap_err();
assert!(
matches!(&err, Error::InvalidInput { message }
if message.contains("input column 'title' is nullable")
&& message.contains("parameter 'title'")
&& message.contains("binding 'fb_01K3TEXT'")
&& message.contains("non-nullable")),
"{err:?}"
);
}
#[test]
fn test_second_binding_rejects_outputs_without_reciprocal_metadata() {
let binding = FunctionBinding::from_json(include_str!(
"../../tests/fixtures/first_class_functions/v1/remote_function_binding.json"
))
.unwrap();
let schema = ArrowSchema::new_with_metadata(
function_binding_schema(true, true).fields().to_vec(),
HashMap::from([(
FUNCTION_BINDINGS_META_KEY.to_string(),
function_bindings_metadata(std::slice::from_ref(&binding)).unwrap(),
)]),
);
let err = plan_function_application(
&schema,
&named_struct_application(
r#"{"normalized_text":"secondary_text","token_count":"secondary_token_count"}"#,
),
None,
)
.unwrap_err();
assert!(
matches!(&err, Error::InvalidInput { message }
if message.contains("declaration metadata")
&& message.contains("fb_01K3TEXT")),
"{err:?}"
);
}
#[test]
fn test_persisted_nested_input_keeps_leaf_level_validation() {
let mut raw_binding: Value = serde_json::from_str(include_str!(
"../../tests/fixtures/first_class_functions/v1/remote_function_binding.json"
))
.unwrap();
raw_binding["inputs"][0]["field_path"] = Value::String("title.value".to_string());
let binding: FunctionBinding = serde_json::from_value(raw_binding).unwrap();
let title = ArrowField::new(
"title",
DataType::Struct(vec![ArrowField::new("value", DataType::Utf8, true)].into()),
true,
)
.with_metadata(HashMap::from([
(COMPUTED_COLUMN_META_KEY.to_string(), "true".to_string()),
(KIND_META_KEY.to_string(), SQL_KIND.to_string()),
(EXPRESSION_META_KEY.to_string(), "title".to_string()),
]));
let mut fields = vec![title, ArrowField::new("body", DataType::Utf8, true)];
fields.extend(binding.outputs().iter().map(|output| {
let data_type = match output.arrow_type.as_str() {
"utf8" => DataType::Utf8,
"int64" => DataType::Int64,
other => panic!("unexpected fixture output type {other}"),
};
ArrowField::new(&output.output_name, data_type, true).with_metadata(
function_computed_column_metadata(
binding.binding_id(),
output.output_ordinal,
&["title.value".into(), "body".into()],
),
)
}));
ensure_binding_matches_schema(&ArrowSchema::new(fields), &binding).unwrap();
}
#[test]
fn test_function_binding_metadata_survives_schema_round_trip() {
let binding = FunctionBinding::from_json(include_str!(
"../../tests/fixtures/first_class_functions/v1/remote_function_binding.json"
))
.unwrap();
let raw = function_bindings_metadata(std::slice::from_ref(&binding)).unwrap();
let mut fields = vec![
ArrowField::new("title", DataType::Utf8, true),
ArrowField::new("body", DataType::Utf8, true),
];
fields.extend(
binding
.outputs()
.iter()
.map(|output| {
let data_type = match output.arrow_type.as_str() {
"utf8" => DataType::Utf8,
"int64" => DataType::Int64,
other => panic!("unexpected fixture output type {other}"),
};
let metadata = function_computed_column_metadata(
binding.binding_id(),
output.output_ordinal,
&["title".into(), "body".into()],
);
ArrowField::new(&output.output_name, data_type, true).with_metadata(metadata)
})
.collect::<Vec<_>>(),
);
let schema = ArrowSchema::new_with_metadata(
fields,
HashMap::from([(FUNCTION_BINDINGS_META_KEY.to_string(), raw)]),
);
let reopened =
ArrowSchema::new_with_metadata(schema.fields().to_vec(), schema.metadata().clone());
let bindings = function_bindings(&reopened).unwrap();
assert_eq!(bindings, vec![binding.clone()]);
assert!(bindings[0].input_schema().is_some());
assert!(bindings[0].output_schema().is_some());
assert!(matches!(
computed_column_from_field(reopened.field(3)).unwrap().kind,
ComputedColumnKind::Function {
ref binding_id,
output_ordinal: 1,
} if binding_id == "fb_01K3TEXT"
));
let dependent_application = FunctionApplication::from_json(
r#"{
"function":{"name":"dependent","version":"fv_dependent"},
"inputs":[
{"parameter":"text","kind":"column","value":{"path":"search_text"}}
],
"output":{"kind":"scalar","arrow_type":"int64","nullable":false}
}"#,
)
.unwrap();
let err = plan_function_application(&reopened, &dependent_application, Some("dependent"))
.unwrap_err();
assert!(
matches!(&err, Error::InvalidInput { message } if message.contains("computed-on-computed"))
);
let plan = plan_function_application(
&reopened,
&named_struct_application(
r#"{"normalized_text":"secondary_text","token_count":"secondary_token_count"}"#,
),
None,
)
.unwrap();
assert_eq!(
plan.outputs
.iter()
.map(|output| output.output_name.as_str())
.collect::<Vec<_>>(),
["secondary_text", "secondary_token_count"]
);
}
#[test]
fn test_newer_binding_fields_remain_readable_but_fail_closed_on_mutation() {
let raw_binding: Value = serde_json::from_str(include_str!(
"../../tests/fixtures/first_class_functions/v1/remote_function_binding.json"
))
.unwrap();
let binding: FunctionBinding = serde_json::from_value(raw_binding.clone()).unwrap();
assert_eq!(binding.binding_id(), "fb_01K3TEXT");
let schema = ArrowSchema::new_with_metadata(
Vec::<ArrowField>::new(),
HashMap::from([(
FUNCTION_BINDINGS_META_KEY.to_string(),
serde_json::json!({
"version": FUNCTION_BINDINGS_VERSION,
"bindings": [raw_binding],
})
.to_string(),
)]),
);
let err = ensure_supported_function_metadata(&schema).unwrap_err();
assert!(matches!(err, Error::NotSupported { .. }));
}
#[test]
fn test_named_struct_can_be_kept_as_one_nullable_physical_column() {
let application = named_struct_application("{}");
let plan =
plan_function_application(&function_input_schema(), &application, Some("features"))
.unwrap();
assert_eq!(plan.outputs.len(), 1);
assert_eq!(plan.outputs[0].result_field, WHOLE_RESULT_FIELD);
assert_eq!(plan.output_schema.fields.len(), 1);
assert!(plan.output_schema.fields[0].nullable);
assert_eq!(plan.output_schema.fields[0].r#type.r#type, "struct");
assert_eq!(
plan.output_schema.fields[0]
.r#type
.fields
.as_ref()
.unwrap()
.len(),
2
);
}
#[test]
fn test_blob_function_plans_semantic_input_and_scalar_output() {
let schema = ArrowSchema::new(vec![crate::blob("image", false)]);
let application =
blob_application(r#"{"kind":"scalar","arrow_type":"blob_v2","nullable":false}"#);
let plan = plan_function_application(&schema, &application, Some("thumbnail")).unwrap();
assert_eq!(plan.input_bindings[0].arrow_type, FUNCTION_BLOB_V2_TYPE);
let input_schema =
lance_namespace::schema::convert_json_arrow_schema(&plan.input_schema).unwrap();
assert!(input_schema.field(0).is_blob_v2());
let output_schema =
lance_namespace::schema::convert_json_arrow_schema(&plan.output_schema).unwrap();
assert!(output_schema.field(0).is_blob_v2());
}
#[test]
fn test_blob_scalar_binding_accepts_full_logical_layout() {
let input = crate::blob("image", false);
let application =
blob_application(r#"{"kind":"scalar","arrow_type":"blob_v2","nullable":false}"#);
let plan = plan_function_application(
&ArrowSchema::new(vec![input.clone()]),
&application,
Some("thumbnail"),
)
.unwrap();
let binding = binding_from_plan(&plan);
let mut metadata = full_blob_field("thumbnail", true).metadata().clone();
metadata.extend(function_computed_column_metadata(
binding.binding_id(),
0,
&["image".into()],
));
let output = full_blob_field("thumbnail", true).with_metadata(metadata);
ensure_binding_matches_schema(&ArrowSchema::new(vec![input, output]), &binding).unwrap();
}
#[test]
fn test_blob_binding_rejects_marker_on_invalid_storage_layout() {
let input = crate::blob("image", false);
let application =
blob_application(r#"{"kind":"scalar","arrow_type":"blob_v2","nullable":false}"#);
let plan = plan_function_application(
&ArrowSchema::new(vec![input.clone()]),
&application,
Some("thumbnail"),
)
.unwrap();
let binding = binding_from_plan(&plan);
let malformed = ArrowField::new("thumbnail", DataType::Int64, true)
.with_metadata(crate::blob("thumbnail", true).metadata().clone());
ensure_binding_matches_schema(&ArrowSchema::new(vec![input, malformed]), &binding)
.unwrap_err();
}
#[test]
fn test_blob_input_rejects_marker_on_invalid_storage_layout() {
let malformed = ArrowField::new("image", DataType::Int64, false)
.with_metadata(crate::blob("image", false).metadata().clone());
let application =
blob_application(r#"{"kind":"scalar","arrow_type":"blob_v2","nullable":false}"#);
plan_function_application(
&ArrowSchema::new(vec![malformed]),
&application,
Some("thumbnail"),
)
.unwrap_err();
}
#[test]
fn test_non_blob_input_does_not_require_json_round_trip() {
let json = lance_namespace::schema::arrow_schema_to_json(&ArrowSchema::new(vec![
ArrowField::new("event_time", DataType::Time64(TimeUnit::Microsecond), false),
]))
.unwrap();
assert_eq!(
canonical_input_arrow_type(&json.fields[0]).unwrap(),
"time64"
);
}
#[test]
fn test_blob_named_struct_plans_expanded_and_whole_outputs() {
let schema = ArrowSchema::new(vec![crate::blob("image", false)]);
let application = blob_application(
r#"{"kind":"named_struct","fields":[
{"name":"thumbnail","arrow_type":"blob_v2","nullable":false},
{"name":"width","arrow_type":"int32","nullable":false}
]}"#,
);
let expanded = plan_function_application(&schema, &application, None).unwrap();
let expanded_schema =
lance_namespace::schema::convert_json_arrow_schema(&expanded.output_schema).unwrap();
assert!(expanded_schema.field(0).is_blob_v2());
assert_eq!(expanded_schema.field(1).data_type(), &DataType::Int32);
let whole = plan_function_application(&schema, &application, Some("analysis")).unwrap();
let whole_schema =
lance_namespace::schema::convert_json_arrow_schema(&whole.output_schema).unwrap();
let DataType::Struct(fields) = whole_schema.field(0).data_type() else {
panic!("whole Function output should be a struct");
};
assert!(fields[0].is_blob_v2());
assert_eq!(fields[1].data_type(), &DataType::Int32);
}
#[test]
fn test_blob_whole_struct_binding_accepts_full_logical_layout() {
let input = crate::blob("image", false);
let application = blob_application(
r#"{"kind":"named_struct","fields":[
{"name":"thumbnail","arrow_type":"blob_v2","nullable":false},
{"name":"width","arrow_type":"int32","nullable":false}
]}"#,
);
let plan = plan_function_application(
&ArrowSchema::new(vec![input.clone()]),
&application,
Some("analysis"),
)
.unwrap();
let binding = binding_from_plan(&plan);
let output = ArrowField::new(
"analysis",
DataType::Struct(Fields::from(vec![
full_blob_field("thumbnail", false),
ArrowField::new("width", DataType::Int32, false),
])),
true,
)
.with_metadata(function_computed_column_metadata(
binding.binding_id(),
0,
&["image".into()],
));
ensure_binding_matches_schema(&ArrowSchema::new(vec![input, output]), &binding).unwrap();
}
#[test]
fn test_function_mapping_and_sibling_collisions_fail_before_request() {
let unknown = named_struct_application(r#"{"missing":"renamed"}"#);
let err = plan_function_application(&function_input_schema(), &unknown, None).unwrap_err();
assert!(matches!(&err, Error::InvalidInput { message } if message.contains("unknown")));
let duplicate =
named_struct_application(r#"{"normalized_text":"same","token_count":"same"}"#);
let err =
plan_function_application(&function_input_schema(), &duplicate, None).unwrap_err();
assert!(
matches!(&err, Error::InvalidInput { message } if message.contains("destinations"))
);
let mut fields = function_input_schema().fields().to_vec();
fields.push(Arc::new(ArrowField::new(
"token_count",
DataType::Int64,
true,
)));
let collision_schema = ArrowSchema::new(fields);
let err =
plan_function_application(&collision_schema, &named_struct_application("{}"), None)
.unwrap_err();
assert!(matches!(err, Error::ColumnAlreadyExists { name } if name == "token_count"));
}
#[test]
fn test_unknown_and_mixed_version_function_contracts_fail_closed() {
let application = FunctionApplication::from_json(
r#"{
"function":{"name":"f","version":"fv"},
"inputs":[{"parameter":"title","kind":"future_source","value":{"path":"title"}}],
"output":{"kind":"scalar","arrow_type":"int64","nullable":false}
}"#,
)
.unwrap();
let err = plan_function_application(&function_input_schema(), &application, Some("out"))
.unwrap_err();
assert!(matches!(err, Error::NotSupported { .. }));
let future_application = FunctionApplication::from_json(
r#"{
"function":{"name":"f","version":"fv"},
"inputs":[],
"output":{"kind":"scalar","arrow_type":"int64","nullable":false},
"future_declaration":{"mode":"managed"}
}"#,
)
.unwrap();
let err =
plan_function_application(&function_input_schema(), &future_application, Some("out"))
.unwrap_err();
assert!(matches!(err, Error::NotSupported { .. }));
let nested_future_application = FunctionApplication::from_json(
r#"{
"function":{"name":"f","version":"fv"},
"inputs":[],
"output":{"kind":"scalar","arrow_type":"int64","nullable":false,"assignment":"cell_flag"}
}"#,
)
.unwrap();
let err = plan_function_application(
&function_input_schema(),
&nested_future_application,
Some("out"),
)
.unwrap_err();
assert!(matches!(err, Error::NotSupported { .. }));
let mixed_schema = ArrowSchema::new_with_metadata(
function_input_schema().fields().to_vec(),
HashMap::from([(
FUNCTION_BINDINGS_META_KEY.to_string(),
r#"{"version":2,"bindings":[]}"#.to_string(),
)]),
);
let err = plan_function_application(&mixed_schema, &named_struct_application("{}"), None)
.unwrap_err();
assert!(matches!(err, Error::NotSupported { .. }));
}
#[test]
fn test_function_inputs_use_paths_and_cannot_be_computed() {
let mut schema = function_input_schema();
let plan =
plan_function_application(&schema, &named_struct_application("{}"), None).unwrap();
assert_eq!(plan.input_bindings[0].field_path, "title");
assert_eq!(plan.input_bindings[1].field_path, "body");
let title = schema
.field(0)
.as_ref()
.clone()
.with_metadata(HashMap::from([
(COMPUTED_COLUMN_META_KEY.to_string(), "true".to_string()),
(KIND_META_KEY.to_string(), SQL_KIND.to_string()),
(EXPRESSION_META_KEY.to_string(), "title".to_string()),
]));
schema = ArrowSchema::new(vec![title, schema.field(1).as_ref().clone()]);
let err =
plan_function_application(&schema, &named_struct_application("{}"), None).unwrap_err();
assert!(
matches!(&err, Error::InvalidInput { message } if message.contains("computed-on-computed"))
);
let nested_title = ArrowField::new(
"title",
DataType::Struct(vec![ArrowField::new("value", DataType::Utf8, true)].into()),
true,
)
.with_metadata(HashMap::from([
(COMPUTED_COLUMN_META_KEY.to_string(), "true".to_string()),
(KIND_META_KEY.to_string(), SQL_KIND.to_string()),
(
EXPRESSION_META_KEY.to_string(),
"struct('value')".to_string(),
),
]));
let nested_schema = ArrowSchema::new(vec![nested_title, schema.field(1).as_ref().clone()]);
let nested_application = FunctionApplication::from_json(
r#"{
"function":{"name":"text_features","version":"fv_exact"},
"inputs":[
{"parameter":"title","kind":"column","value":{"path":"title.value"}},
{"parameter":"body","kind":"column","value":{"path":"body"}}
],
"output":{"kind":"named_struct","fields":[
{"name":"normalized_text","arrow_type":"utf8","nullable":false},
{"name":"token_count","arrow_type":"int64","nullable":false}
]}
}"#,
)
.unwrap();
let err = plan_function_application(&nested_schema, &nested_application, None).unwrap_err();
assert!(
matches!(&err, Error::InvalidInput { message } if message.contains("computed-on-computed"))
);
}
}