diff --git a/python/python/lancedb/_lancedb.pyi b/python/python/lancedb/_lancedb.pyi index 387c85e80..d5c11699f 100644 --- a/python/python/lancedb/_lancedb.pyi +++ b/python/python/lancedb/_lancedb.pyi @@ -226,6 +226,24 @@ class Function: @property def output_nullable(self) -> bool: ... +class _FunctionDefinition: + """Private owner of the Rust FunctionDefinition registration input.""" + + def _to_json(self) -> str: ... + +def _new_function_definition( + *, + parameters: list[tuple[str, pa.DataType]], + output_type: pa.DataType, + output_nullable: bool, + module: str, + callable_name: str, + source: str, + python: str, + packages: list[str], + capabilities: list[tuple[str, str, Optional[str]]], +) -> _FunctionDefinition: ... + class Job: @property def id(self) -> Optional[str]: ... diff --git a/python/python/lancedb/_udf.py b/python/python/lancedb/_udf.py index 181bfe250..985cf4e77 100644 --- a/python/python/lancedb/_udf.py +++ b/python/python/lancedb/_udf.py @@ -23,6 +23,8 @@ from typing import NoReturn, ParamSpec, TypeVar import pyarrow as pa +from . import _lancedb + __all__ = ["FunctionCapability", "udf"] _P = ParamSpec("_P") @@ -206,6 +208,20 @@ def _validate_packages(packages: object) -> tuple[str, ...]: return tuple(snapshot) +def _reject_non_exact_capability() -> NoReturn: + # Exact-type only: subclasses are authoring inputs we never accept. Keep the + # message fixed so hostile markers never enter exception text. + raise TypeError( + "udf capabilities must contain only FunctionCapability values" + ) from None + + +def _require_exact_capability(capability: object) -> FunctionCapability: + if type(capability) is not FunctionCapability: + _reject_non_exact_capability() + return capability + + def _validate_capabilities(capabilities: object) -> tuple[FunctionCapability, ...]: if isinstance(capabilities, (str, bytes, bytearray)): raise TypeError( @@ -213,14 +229,7 @@ def _validate_capabilities(capabilities: object) -> tuple[FunctionCapability, .. ) if not isinstance(capabilities, Sequence): raise TypeError("udf capabilities must be a sequence of FunctionCapability") - snapshot: list[FunctionCapability] = [] - for capability in capabilities: - if not isinstance(capability, FunctionCapability): - raise TypeError( - "udf capabilities must contain only FunctionCapability values" - ) - snapshot.append(capability) - return tuple(snapshot) + return tuple(_require_exact_capability(capability) for capability in capabilities) def udf( @@ -486,3 +495,44 @@ def _package_udf(fn: object) -> _PackagedUdf: callable_name=callable_name, config=config, ) + + +def _normalize_capability_triple( + capability: FunctionCapability, +) -> tuple[str, str, str | None]: + """Normalize a local capability declaration to the native triple shape.""" + # Private config is untrusted; re-check exact type before any property access. + capability = _require_exact_capability(capability) + if capability.kind == "network": + origin = capability.origin + if origin is None: + raise ValueError("invalid network capability") from None + return ("network", origin, None) + if capability.kind == "secret": + reference = capability.reference + environment_variable = capability.environment_variable + if reference is None or environment_variable is None: + raise ValueError("invalid secret capability") from None + return ("secret", reference, environment_variable) + # Fail closed without echoing the unknown kind. + raise ValueError("unsupported capability kind") from None + + +def _build_function_definition(fn: object) -> _lancedb._FunctionDefinition: + """Package a ``@udf`` and bridge it to the private native definition.""" + packaged = _package_udf(fn) + config = packaged.config + capabilities = [ + _normalize_capability_triple(capability) for capability in config.capabilities + ] + return _lancedb._new_function_definition( + parameters=list(config.inputs), + output_type=config.output, + output_nullable=config.output_nullable, + module=packaged.module, + callable_name=packaged.callable_name, + source=packaged.source, + python=config.python, + packages=list(config.packages), + capabilities=capabilities, + ) diff --git a/python/python/tests/test_first_class_udf_capability.py b/python/python/tests/test_first_class_udf_capability.py index 168f840f2..537ad83a5 100644 --- a/python/python/tests/test_first_class_udf_capability.py +++ b/python/python/tests/test_first_class_udf_capability.py @@ -381,6 +381,43 @@ def test_udf_capabilities_ordered_immutable_config_default_and_validation(): assert "unique-bad-capability-string-xyz" not in repr(bad_mixed.value) +def test_udf_capabilities_rejects_function_capability_subclass_before_property_access(): + marker = "unique-hostile-capability-subclass-marker-xyz" + + class _HostileFunctionCapability(FunctionCapability): + @property + def kind(self) -> str: + raise RuntimeError(marker) + + @property + def origin(self) -> str | None: + raise RuntimeError(marker) + + @property + def reference(self) -> str | None: + raise RuntimeError(marker) + + @property + def environment_variable(self) -> str | None: + raise RuntimeError(marker) + + hostile = object.__new__(_HostileFunctionCapability) + assert isinstance(hostile, FunctionCapability) + assert type(hostile) is not FunctionCapability + + def target(x): + return x + + with pytest.raises(TypeError) as exc_info: + _decorate(target, capabilities=[hostile]) + assert marker not in str(exc_info.value) + assert marker not in repr(exc_info.value) + assert _SECRET_REFERENCE not in str(exc_info.value) + assert _SECRET_REFERENCE not in repr(exc_info.value) + assert exc_info.value.__cause__ is None + assert exc_info.value.__context__ is None + + def test_package_udf_preserves_capabilities_and_redacts_secret_reference(): packaged = _package_udf(packable_with_capabilities) config = packaged.config diff --git a/python/python/tests/test_first_class_udf_definition_bridge.py b/python/python/tests/test_first_class_udf_definition_bridge.py new file mode 100644 index 000000000..92ae98574 --- /dev/null +++ b/python/python/tests/test_first_class_udf_definition_bridge.py @@ -0,0 +1,506 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright The LanceDB Authors + +"""RED contract tests for the private UDF -> FunctionDefinition bridge.""" + +from __future__ import annotations + +import base64 +import io +import json +from pathlib import Path + +import pyarrow as pa +import pytest + +import lancedb +from lancedb import FunctionCapability, udf +from lancedb import _lancedb as _native +from lancedb import _udf as _udf_mod + +_SOURCE_MARKER = "bridge-source-marker-unique-xyz" +_SECRET_REFERENCE = "secret://team/bridge-redact-token-xyz" +_SECRET_ENV = "BRIDGE_API_TOKEN" +_NETWORK_ORIGIN = "https://api.bridge-example.com" +_NETWORK_ORIGIN_B = "https://other.bridge-example.com" + +_FORBIDDEN_WIRE_KEYS = ( + "id", + "function_id", + "FunctionId", + "catalog", + "catalog_name", + "version", + "function_version", + "FunctionVersion", + "lineage", + "user_version", + "idempotency_key", + "digest", + "artifact", + "artifact_digest", + "storage", + "storage_location", + "location", + "deterministic", + "null_policy", + "nullPolicy", + "timestamp", + "created_at", + "updated_at", + "worker", + "scheduler", + "attempt", + "attempt_id", + "replica", + "placement", + "job", + "job_id", + "retry_key", + "registration", +) + +_OVERDESIGN_ATTRS = ( + "id", + "function_id", + "FunctionVersion", + "function_version", + "artifact", + "digest", + "job", + "job_id", + "catalog", + "retry_key", + "user_version", + "idempotency_key", + "deterministic", + "null_policy", + "null_handling", +) + + +@udf( + inputs={"text": pa.string(), "limit": pa.int32()}, + output=pa.string(), + python="3.12", + packages=["pkg-b==2", "pkg-a==1"], + output_nullable=True, + capabilities=[ + FunctionCapability.network(_NETWORK_ORIGIN), + FunctionCapability.secret( + _SECRET_REFERENCE, + environment_variable=_SECRET_ENV, + ), + FunctionCapability.network(_NETWORK_ORIGIN_B), + ], +) +def packable_bridge_normalize(text, limit): + """bridge-source-marker-unique-xyz.""" + return text[:limit] + + +def _build_function_definition(fn: object): + return _udf_mod._build_function_definition(fn) + + +def _function_definition_type(): + return _native._FunctionDefinition + + +def _new_function_definition(**kwargs): + return _native._new_function_definition(**kwargs) + + +def _json_bytes(definition) -> bytes: + payload = definition._to_json() + if isinstance(payload, bytes): + return payload + assert isinstance(payload, str) + return payload.encode("utf-8") + + +def _decode_type_ipc(encoded: str) -> pa.DataType: + raw = base64.b64decode(encoded) + reader = pa.ipc.open_file(io.BytesIO(raw)) + assert reader.num_record_batches == 0 + assert len(reader.schema) == 1 + return reader.schema.field(0).type + + +def _assert_exact_object_keys(value: dict, expected: set[str], *, context: str) -> None: + assert isinstance(value, dict), f"{context} must be an object" + assert set(value) == expected, f"{context} keys must match exactly: {set(value)!r}" + + +def _assert_forbidden_keys_absent(value: object, *, context: str) -> None: + if isinstance(value, dict): + for key in value: + assert key not in _FORBIDDEN_WIRE_KEYS, ( + f"forbidden key {key!r} at {context}: {value!r}" + ) + if key == "name" and context in { + "definition", + "signature", + "signature.output", + "implementation", + }: + raise AssertionError( + f"catalog/function identity key `name` must be absent at {context}" + ) + child_context = f"{context}.{key}" + if key == "parameters" and context == "signature": + child_context = "signature.parameters" + _assert_forbidden_keys_absent(value[key], context=child_context) + elif isinstance(value, list): + for idx, item in enumerate(value): + item_context = ( + f"signature.parameters[{idx}]" + if context == "signature.parameters" + else f"{context}[{idx}]" + ) + if context == "signature.parameters": + assert isinstance(item, dict) + assert "name" in item + for key in item: + assert key not in _FORBIDDEN_WIRE_KEYS + assert key != "catalog_name" + _assert_forbidden_keys_absent( + {k: v for k, v in item.items() if k != "name"}, + context=item_context, + ) + else: + _assert_forbidden_keys_absent(item, context=item_context) + + +def _assert_sanitized_text(*parts: object) -> None: + combined = "\n".join(str(part) for part in parts) + lowered = combined.lower() + assert _SOURCE_MARKER.lower() not in lowered + assert _SECRET_REFERENCE.lower() not in lowered + assert str(Path(__file__).resolve()).lower() not in lowered + assert Path(__file__).resolve().as_posix().lower() not in lowered + + +def _assert_clean_validation_error(exc_info) -> None: + _assert_sanitized_text(exc_info.value, repr(exc_info.value)) + assert exc_info.value.__cause__ is None + assert exc_info.value.__context__ is None + + +def _valid_builder_kwargs(**overrides): + kwargs = { + "parameters": [("text", pa.string()), ("limit", pa.int32())], + "output_type": pa.string(), + "output_nullable": True, + "module": "bridge_mod", + "callable_name": "normalize", + "source": ( + "def normalize(text, limit):\n" + f" # {_SOURCE_MARKER}\n" + " return text[:limit]\n" + ), + "python": "3.12", + "packages": ["pkg-b==2", "pkg-a==1"], + "capabilities": [ + ("network", _NETWORK_ORIGIN, None), + ("secret", _SECRET_REFERENCE, _SECRET_ENV), + ("network", _NETWORK_ORIGIN_B, None), + ], + } + kwargs.update(overrides) + return kwargs + + +def test_build_function_definition_private_native_immutability_and_export_surface(): + assert "_build_function_definition" not in getattr(lancedb, "__all__", []) + assert "_FunctionDefinition" not in lancedb.__all__ + assert not hasattr(lancedb, "_FunctionDefinition") + assert not hasattr(lancedb, "_build_function_definition") + assert not hasattr(lancedb, "_new_function_definition") + + definition = _build_function_definition(packable_bridge_normalize) + definition_type = _function_definition_type() + assert type(definition) is definition_type + assert definition_type.__module__ == "lancedb._lancedb" + assert definition_type.__name__ == "_FunctionDefinition" + + with pytest.raises(TypeError): + definition_type() + + for attr in _OVERDESIGN_ATTRS: + assert not hasattr(definition, attr) + + for attr in ("signature", "module", "source", "capabilities"): + with pytest.raises(AttributeError): + setattr(definition, attr, None) + + +def test_build_function_definition_json_wire_ordered_contract_without_identity(): + definition = _build_function_definition(packable_bridge_normalize) + encoded_a = _json_bytes(definition) + encoded_b = _json_bytes(definition) + assert encoded_a == encoded_b + + wire = json.loads(encoded_a.decode("utf-8")) + _assert_exact_object_keys( + wire, + {"format_version", "signature", "implementation", "capabilities"}, + context="definition", + ) + assert wire["format_version"] == 1 + _assert_forbidden_keys_absent(wire, context="definition") + + signature = wire["signature"] + _assert_exact_object_keys(signature, {"parameters", "output"}, context="signature") + parameters = signature["parameters"] + assert [parameter["name"] for parameter in parameters] == ["text", "limit"] + for parameter in parameters: + _assert_exact_object_keys( + parameter, {"name", "data_type_ipc"}, context="parameter" + ) + assert isinstance(parameter["data_type_ipc"], str) + assert parameter["data_type_ipc"] + assert _decode_type_ipc(parameters[0]["data_type_ipc"]) == pa.string() + assert _decode_type_ipc(parameters[1]["data_type_ipc"]) == pa.int32() + + output = signature["output"] + _assert_exact_object_keys( + output, {"data_type_ipc", "nullable"}, context="signature.output" + ) + assert output["nullable"] is True + assert _decode_type_ipc(output["data_type_ipc"]) == pa.string() + + implementation = wire["implementation"] + _assert_exact_object_keys( + implementation, + {"kind", "module", "callable", "source", "python", "packages"}, + context="implementation", + ) + assert implementation["kind"] == "python" + assert implementation["module"] == __name__ + assert implementation["callable"] == "packable_bridge_normalize" + assert implementation["source"] == Path(__file__).read_text(encoding="utf-8") + assert _SOURCE_MARKER in implementation["source"] + assert implementation["python"] == "3.12" + assert implementation["packages"] == ["pkg-b==2", "pkg-a==1"] + + capabilities = wire["capabilities"] + assert capabilities == [ + {"kind": "network", "origin": _NETWORK_ORIGIN}, + { + "kind": "secret", + "reference": _SECRET_REFERENCE, + "environment_variable": _SECRET_ENV, + }, + {"kind": "network", "origin": _NETWORK_ORIGIN_B}, + ] + for capability in capabilities: + assert "value" not in capability + assert "plaintext" not in capability + assert "plaintext_secret" not in capability + assert "secret_value" not in capability + + +def test_native_definition_repr_includes_safe_structure_and_redacts_sensitive_text(): + definition = _build_function_definition(packable_bridge_normalize) + rendered = repr(definition) + assert "_FunctionDefinition" in rendered or "FunctionDefinition" in rendered + assert __name__ in rendered + assert "packable_bridge_normalize" in rendered + assert "3.12" in rendered + _assert_sanitized_text(rendered) + + +def test_new_function_definition_builder_preserves_normalized_wire(): + definition = _new_function_definition(**_valid_builder_kwargs()) + assert type(definition) is _function_definition_type() + + encoded_a = _json_bytes(definition) + encoded_b = _json_bytes(definition) + assert encoded_a == encoded_b + + wire = json.loads(encoded_a.decode("utf-8")) + assert wire["format_version"] == 1 + assert [parameter["name"] for parameter in wire["signature"]["parameters"]] == [ + "text", + "limit", + ] + assert _decode_type_ipc(wire["signature"]["parameters"][0]["data_type_ipc"]) == ( + pa.string() + ) + assert _decode_type_ipc(wire["signature"]["parameters"][1]["data_type_ipc"]) == ( + pa.int32() + ) + assert wire["signature"]["output"]["nullable"] is True + assert _decode_type_ipc(wire["signature"]["output"]["data_type_ipc"]) == pa.string() + + implementation = wire["implementation"] + assert implementation["kind"] == "python" + assert implementation["module"] == "bridge_mod" + assert implementation["callable"] == "normalize" + assert implementation["source"] == _valid_builder_kwargs()["source"] + assert _SOURCE_MARKER in implementation["source"] + assert implementation["python"] == "3.12" + assert implementation["packages"] == ["pkg-b==2", "pkg-a==1"] + assert wire["capabilities"] == [ + {"kind": "network", "origin": _NETWORK_ORIGIN}, + { + "kind": "secret", + "reference": _SECRET_REFERENCE, + "environment_variable": _SECRET_ENV, + }, + {"kind": "network", "origin": _NETWORK_ORIGIN_B}, + ] + _assert_forbidden_keys_absent(wire, context="definition") + + +@pytest.mark.parametrize( + ("overrides",), + [ + ({"parameters": [("text", pa.string()), ("text", pa.int32())]},), + ({"parameters": [("", pa.string())]},), + ({"module": ""},), + ({"callable_name": ""},), + ({"source": ""},), + ({"python": ""},), + ({"packages": ["pkg-a==1", ""]},), + ({"packages": ["pkg-a==1", "pkg-a==1"]},), + ({"capabilities": [("filesystem", _NETWORK_ORIGIN, None)]},), + ({"capabilities": [("network", _NETWORK_ORIGIN, _SECRET_ENV)]},), + ({"capabilities": [("secret", _SECRET_REFERENCE, None)]},), + ({"capabilities": [("secret", _SECRET_REFERENCE, "")]},), + ({"capabilities": [("network", "", None)]},), + ({"capabilities": [("secret", "", _SECRET_ENV)]},), + ], +) +def test_new_function_definition_strict_validation_rejections(overrides): + kwargs = _valid_builder_kwargs(**overrides) + with pytest.raises(ValueError) as exc_info: + _new_function_definition(**kwargs) + _assert_clean_validation_error(exc_info) + + +def test_new_function_definition_validation_does_not_echo_secret_or_source_marker(): + with pytest.raises(ValueError) as exc_info: + _new_function_definition(**_valid_builder_kwargs(module="")) + _assert_clean_validation_error(exc_info) + + with pytest.raises(ValueError) as exc_info: + _new_function_definition( + **_valid_builder_kwargs(packages=["pkg-a==1", "pkg-a==1"]) + ) + _assert_clean_validation_error(exc_info) + + with pytest.raises(ValueError) as exc_info: + _new_function_definition( + **_valid_builder_kwargs( + capabilities=[("secret", _SECRET_REFERENCE, None)], + ) + ) + _assert_clean_validation_error(exc_info) + + +@pytest.mark.parametrize( + ("overrides",), + [ + ({"parameters": [("text", "not-a-datatype")]},), + ({"parameters": [(123, pa.string())]},), + ({"output_type": "not-a-datatype"},), + ({"output_type": None},), + ({"output_nullable": "yes"},), + ({"packages": "pkg-a==1"},), + ({"capabilities": "network"},), + ({"capabilities": [("network", _NETWORK_ORIGIN)]},), + ({"capabilities": [("network", _NETWORK_ORIGIN, None, "extra")]},), + ], +) +def test_new_function_definition_wrong_pyarrow_and_shape_values_fail_closed(overrides): + kwargs = _valid_builder_kwargs(**overrides) + with pytest.raises((TypeError, ValueError)) as exc_info: + _new_function_definition(**kwargs) + _assert_clean_validation_error(exc_info) + + +class _HostileRaisingIterable: + def __iter__(self): + raise RuntimeError(f"{_SECRET_REFERENCE} {_SOURCE_MARKER}") + + +@pytest.mark.parametrize( + ("overrides",), + [ + ({"parameters": _HostileRaisingIterable()},), + ({"packages": _HostileRaisingIterable()},), + ({"capabilities": _HostileRaisingIterable()},), + ( + { + "capabilities": [ + ("network", _NETWORK_ORIGIN, None), + _HostileRaisingIterable(), + ("network", _NETWORK_ORIGIN_B, None), + ] + }, + ), + ], +) +def test_new_function_definition_hostile_iterable_iter_raises_fail_closed(overrides): + kwargs = _valid_builder_kwargs(**overrides) + with pytest.raises((TypeError, ValueError)) as exc_info: + _new_function_definition(**kwargs) + _assert_clean_validation_error(exc_info) + + +@udf( + inputs={"x": pa.int32()}, + output=pa.int64(), + python="3.12", + packages=["pkg-a==1"], + output_nullable=False, +) +def packable_bridge_capability_exact_type(x): + return x + 1 + + +def test_build_function_definition_rejects_forged_function_capability_subclass(): + marker = f"{_SECRET_REFERENCE} {_SOURCE_MARKER}" + + class _HostileFunctionCapability(FunctionCapability): + @property + def kind(self) -> str: + raise RuntimeError(marker) + + @property + def origin(self) -> str | None: + raise RuntimeError(marker) + + @property + def reference(self) -> str | None: + raise RuntimeError(marker) + + @property + def environment_variable(self) -> str | None: + raise RuntimeError(marker) + + hostile = object.__new__(_HostileFunctionCapability) + assert isinstance(hostile, FunctionCapability) + assert type(hostile) is not FunctionCapability + + config_attr = _udf_mod._CONFIG_ATTR + original = getattr(packable_bridge_capability_exact_type, config_attr) + forged = _udf_mod._UdfConfig( + inputs=original.inputs, + output=original.output, + output_nullable=original.output_nullable, + python=original.python, + packages=original.packages, + capabilities=(hostile,), + ) + setattr(packable_bridge_capability_exact_type, config_attr, forged) + try: + with pytest.raises((TypeError, ValueError)) as exc_info: + _build_function_definition(packable_bridge_capability_exact_type) + _assert_clean_validation_error(exc_info) + assert marker not in str(exc_info.value) + assert marker not in repr(exc_info.value) + finally: + setattr(packable_bridge_capability_exact_type, config_attr, original) diff --git a/python/src/function.rs b/python/src/function.rs index c0dcf72df..a201c23c3 100644 --- a/python/src/function.rs +++ b/python/src/function.rs @@ -1,8 +1,20 @@ // SPDX-License-Identifier: Apache-2.0 // SPDX-FileCopyrightText: Copyright The LanceDB Authors -use arrow::pyarrow::ToPyArrow; -use pyo3::{Bound, Py, PyAny, PyResult, Python, pyclass, pymethods, types::PyTuple}; +use arrow::datatypes::DataType; +use arrow::pyarrow::{FromPyArrow, ToPyArrow}; +use lancedb::function::{ + FunctionCapability, FunctionDefinition, FunctionOutput, FunctionParameter, FunctionSignature, + PythonFunctionDefinition, +}; +use pyo3::{ + Bound, Py, PyAny, PyResult, Python, + exceptions::{PyRuntimeError, PyTypeError, PyValueError}, + pyclass, pyfunction, pymethods, + types::{PyAnyMethods, PyBool, PyList, PyListMethods, PyTuple, PyTupleMethods}, +}; + +use crate::error::PythonErrorExt; /// Immutable first-class Function handle backed by the exact Rust value. #[pyclass(frozen, skip_from_py_object)] @@ -60,3 +72,206 @@ impl Function { format!("Function(id={:?})", self.inner.id().as_str()) } } + +/// Private, frozen owner of the exact Rust [`FunctionDefinition`]. +/// +/// Exposed to Python as `lancedb._lancedb._FunctionDefinition`. Not +/// constructible from Python. Sensitive fields are omitted from `__repr__` +/// via the Rust `Debug` redaction contract. +#[pyclass( + name = "_FunctionDefinition", + module = "lancedb._lancedb", + frozen, + skip_from_py_object +)] +pub struct PyFunctionDefinition { + inner: FunctionDefinition, +} + +impl PyFunctionDefinition { + pub(crate) fn new(inner: FunctionDefinition) -> Self { + Self { inner } + } + + /// Crate-private accessor for later registration slices. + #[allow(dead_code)] + pub(crate) fn inner(&self) -> &FunctionDefinition { + &self.inner + } +} + +#[pymethods] +impl PyFunctionDefinition { + fn _to_json(&self) -> PyResult { + // Use the existing serde wire. Never log or format payload data into errors. + serde_json::to_string(&self.inner) + .map_err(|_| PyRuntimeError::new_err("failed to serialize function definition")) + } + + fn __repr__(&self) -> String { + format!("{:?}", self.inner) + } +} + +/// Build a private [`PyFunctionDefinition`] from normalized keyword inputs. +/// +/// Arity mirrors the private Python FFI surface produced by +/// `_build_function_definition`; keep distinct keyword parameters at this boundary. +#[pyfunction(signature = ( + *, + parameters, + output_type, + output_nullable, + module, + callable_name, + source, + python, + packages, + capabilities, +))] +#[allow(clippy::too_many_arguments)] +pub fn _new_function_definition( + parameters: Bound<'_, PyAny>, + output_type: Bound<'_, PyAny>, + output_nullable: Bound<'_, PyAny>, + module: String, + callable_name: String, + source: String, + python: String, + packages: Bound<'_, PyAny>, + capabilities: Bound<'_, PyAny>, +) -> PyResult { + let parameters = parse_parameters(¶meters)?; + let output_data_type = parse_data_type(&output_type, "output_type")?; + let output_nullable = parse_exact_bool(&output_nullable, "output_nullable")?; + let packages = parse_string_list(&packages, "packages")?; + let capabilities = parse_capabilities(&capabilities)?; + + let signature = FunctionSignature::try_new( + parameters, + FunctionOutput::new(output_data_type, output_nullable), + ) + .infer_error()?; + let python_definition = + PythonFunctionDefinition::try_new(module, callable_name, source, python, packages) + .infer_error()?; + let definition = + FunctionDefinition::try_new(signature, python_definition, capabilities).infer_error()?; + Ok(PyFunctionDefinition::new(definition)) +} + +fn parse_exact_bool(value: &Bound<'_, PyAny>, field: &str) -> PyResult { + if !value.is_instance_of::() { + return Err(PyTypeError::new_err(format!("{field} must be a bool"))); + } + value.extract() +} + +fn parse_data_type(value: &Bound<'_, PyAny>, field: &str) -> PyResult { + DataType::from_pyarrow_bound(value) + .map_err(|_| PyTypeError::new_err(format!("{field} must be a pyarrow DataType"))) +} + +fn parse_parameters(parameters: &Bound<'_, PyAny>) -> PyResult> { + let list = parameters.cast_exact::().map_err(|_| { + PyTypeError::new_err("parameters must be a list of (name, data_type) pairs") + })?; + let mut out = Vec::with_capacity(list.len()); + for i in 0..list.len() { + let item = list.get_item(i)?; + let pair = item + .cast_exact::() + .map_err(|_| PyTypeError::new_err("each parameter must be a (name, data_type) pair"))?; + if pair.len() != 2 { + return Err(PyTypeError::new_err( + "each parameter must be a (name, data_type) pair", + )); + } + let name: String = pair + .get_item(0)? + .extract() + .map_err(|_| PyTypeError::new_err("parameter name must be a string"))?; + let data_type = parse_data_type(&pair.get_item(1)?, "parameter data_type")?; + out.push(FunctionParameter::new(name, data_type)); + } + Ok(out) +} + +fn parse_string_list(value: &Bound<'_, PyAny>, field: &str) -> PyResult> { + let list = value + .cast_exact::() + .map_err(|_| PyTypeError::new_err(format!("{field} must be a list of strings")))?; + let mut out = Vec::with_capacity(list.len()); + for i in 0..list.len() { + let item = list.get_item(i)?; + let package: String = item + .extract() + .map_err(|_| PyTypeError::new_err(format!("{field} must contain only strings")))?; + out.push(package); + } + Ok(out) +} + +fn parse_capabilities(capabilities: &Bound<'_, PyAny>) -> PyResult> { + let list = capabilities + .cast_exact::() + .map_err(|_| PyTypeError::new_err("capabilities must be a list of capability triples"))?; + let mut out = Vec::with_capacity(list.len()); + for i in 0..list.len() { + out.push(parse_capability_triple(&list.get_item(i)?)?); + } + Ok(out) +} + +fn parse_capability_triple(item: &Bound<'_, PyAny>) -> PyResult { + let triple = item.cast_exact::().map_err(|_| { + PyTypeError::new_err( + "each capability must be a 3-tuple of (kind, value, environment_variable)", + ) + })?; + if triple.len() != 3 { + return Err(PyTypeError::new_err( + "each capability must be a 3-tuple of (kind, value, environment_variable)", + )); + } + + let kind: String = triple + .get_item(0)? + .extract() + .map_err(|_| PyTypeError::new_err("capability kind must be a string"))?; + let primary: String = triple + .get_item(1)? + .extract() + .map_err(|_| PyTypeError::new_err("capability value must be a string"))?; + let env_obj = triple.get_item(2)?; + let environment_variable = if env_obj.is_none() { + None + } else { + Some(env_obj.extract::().map_err(|_| { + PyTypeError::new_err("capability environment_variable must be a string or None") + })?) + }; + + match kind.as_str() { + "network" => { + if environment_variable.is_some() { + // Fail closed without echoing kind, origin, or any env value. + return Err(PyValueError::new_err( + "network capability must not include an environment variable", + )); + } + FunctionCapability::try_network(primary).infer_error() + } + "secret" => { + let Some(environment_variable) = environment_variable else { + // Fail closed without echoing kind or secret reference. + return Err(PyValueError::new_err( + "secret capability requires an environment variable", + )); + }; + FunctionCapability::try_secret(primary, environment_variable).infer_error() + } + // Fail closed: never echo the supplied kind, source, or secret reference. + _ => Err(PyValueError::new_err("unsupported capability kind")), + } +} diff --git a/python/src/lib.rs b/python/src/lib.rs index 1f02455b8..c6355985a 100644 --- a/python/src/lib.rs +++ b/python/src/lib.rs @@ -47,6 +47,7 @@ pub fn _lancedb(_py: Python, m: &Bound<'_, PyModule>) -> PyResult<()> { m.add_class::()?; m.add_class::()?; m.add_class::()?; + m.add_class::()?; m.add_class::()?; m.add_class::()?; m.add_class::()?; @@ -90,6 +91,10 @@ pub fn _lancedb(_py: Python, m: &Bound<'_, PyModule>) -> PyResult<()> { m.add_function(wrap_pyfunction!(expr_col, m)?)?; m.add_function(wrap_pyfunction!(expr_lit, m)?)?; m.add_function(wrap_pyfunction!(expr_func, m)?)?; + m.add_function(wrap_pyfunction!( + crate::function::_new_function_definition, + m + )?)?; m.add("__version__", env!("CARGO_PKG_VERSION"))?; Ok(()) }