feat: bridge Python UDF definitions to Rust

This commit is contained in:
Xuanwo
2026-08-12 04:34:58 +08:00
parent f8bb90405f
commit 2b10f2a7ce
6 changed files with 841 additions and 10 deletions
+18
View File
@@ -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]: ...
+58 -8
View File
@@ -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,
)
@@ -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
@@ -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)
+217 -2
View File
@@ -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<String> {
// 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<PyFunctionDefinition> {
let parameters = parse_parameters(&parameters)?;
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<bool> {
if !value.is_instance_of::<PyBool>() {
return Err(PyTypeError::new_err(format!("{field} must be a bool")));
}
value.extract()
}
fn parse_data_type(value: &Bound<'_, PyAny>, field: &str) -> PyResult<DataType> {
DataType::from_pyarrow_bound(value)
.map_err(|_| PyTypeError::new_err(format!("{field} must be a pyarrow DataType")))
}
fn parse_parameters(parameters: &Bound<'_, PyAny>) -> PyResult<Vec<FunctionParameter>> {
let list = parameters.cast_exact::<PyList>().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::<PyTuple>()
.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<Vec<String>> {
let list = value
.cast_exact::<PyList>()
.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<Vec<FunctionCapability>> {
let list = capabilities
.cast_exact::<PyList>()
.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<FunctionCapability> {
let triple = item.cast_exact::<PyTuple>().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::<String>().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")),
}
}
+5
View File
@@ -47,6 +47,7 @@ pub fn _lancedb(_py: Python, m: &Bound<'_, PyModule>) -> PyResult<()> {
m.add_class::<Session>()?;
m.add_class::<Table>()?;
m.add_class::<crate::function::Function>()?;
m.add_class::<crate::function::PyFunctionDefinition>()?;
m.add_class::<crate::job::Job>()?;
m.add_class::<crate::job::JobInfo>()?;
m.add_class::<crate::job::JobDescription>()?;
@@ -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(())
}