From 0093bc8179e6161c71264e092c698877ce9cbab3 Mon Sep 17 00:00:00 2001 From: Xuanwo Date: Wed, 12 Aug 2026 03:28:05 +0800 Subject: [PATCH] feat: add Python UDF declarations --- python/python/lancedb/__init__.py | 2 + python/python/lancedb/_udf.py | 124 +++++++++ python/python/tests/test_first_class_udf.py | 291 ++++++++++++++++++++ 3 files changed, 417 insertions(+) create mode 100644 python/python/lancedb/_udf.py create mode 100644 python/python/tests/test_first_class_udf.py diff --git a/python/python/lancedb/__init__.py b/python/python/lancedb/__init__.py index 9f8e1b649..31c0bc377 100644 --- a/python/python/lancedb/__init__.py +++ b/python/python/lancedb/__init__.py @@ -24,6 +24,7 @@ from .schema import blob, vector, BlobType from .job import AsyncJob, Job from .table import AsyncTable, Table from .types import BaseTokenizerType +from ._udf import udf from ._lancedb import Session from .namespace import ( connect_namespace, @@ -523,5 +524,6 @@ __all__ = [ "RemoteDBConnection", "Session", "Table", + "udf", "__version__", ] diff --git a/python/python/lancedb/_udf.py b/python/python/lancedb/_udf.py new file mode 100644 index 000000000..f6b8bcdca --- /dev/null +++ b/python/python/lancedb/_udf.py @@ -0,0 +1,124 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright The LanceDB Authors + +"""Local authoring declaration surface for first-class UDFs. + +This module only snapshots declaration metadata onto a Python function. It does +not package source, mint durable identity, or register anything with a database. +""" + +from __future__ import annotations + +import inspect +from collections.abc import Callable, Mapping, Sequence +from dataclasses import dataclass +from typing import ParamSpec, TypeVar + +import pyarrow as pa + +__all__ = ["udf"] + +_P = ParamSpec("_P") +_R = TypeVar("_R") + +_CONFIG_ATTR = "__lancedb_udf_config__" + + +@dataclass(frozen=True, slots=True) +class _UdfConfig: + """Private frozen snapshot of a ``@udf`` declaration.""" + + inputs: tuple[tuple[str, pa.DataType], ...] + output: pa.DataType + output_nullable: bool + python: str + packages: tuple[str, ...] + + +def _validate_inputs( + inputs: object, +) -> tuple[tuple[str, pa.DataType], ...]: + if not isinstance(inputs, Mapping): + raise TypeError("udf inputs must be a Mapping of name to pyarrow DataType") + snapshot: list[tuple[str, pa.DataType]] = [] + for key, value in inputs.items(): + if not isinstance(key, str): + raise TypeError("udf input names must be strings") + if key == "": + raise ValueError("udf input names must be non-empty") + if not isinstance(value, pa.DataType): + raise TypeError("udf input types must be pyarrow DataType values") + snapshot.append((key, value)) + return tuple(snapshot) + + +def _validate_packages(packages: object) -> tuple[str, ...]: + if isinstance(packages, (str, bytes, bytearray)): + raise TypeError("udf packages must be a sequence of strings, not a string") + if not isinstance(packages, Sequence): + raise TypeError("udf packages must be a sequence of strings") + snapshot: list[str] = [] + seen: set[str] = set() + for package in packages: + if not isinstance(package, str): + raise TypeError("udf packages must contain only strings") + if package == "": + raise ValueError("udf packages must be non-empty strings") + if package in seen: + raise ValueError(f"duplicate udf package: {package}") + seen.add(package) + snapshot.append(package) + return tuple(snapshot) + + +def udf( + *, + inputs: Mapping[str, pa.DataType], + output: pa.DataType, + python: str, + packages: Sequence[str] = (), + output_nullable: bool = True, +) -> Callable[[Callable[_P, _R]], Callable[_P, _R]]: + """Declare a local UDF without packaging or registration. + + Applying the returned decorator attaches a private frozen config snapshot + and returns the exact same function object. + """ + input_snapshot = _validate_inputs(inputs) + if not isinstance(output, pa.DataType): + raise TypeError("udf output must be a pyarrow DataType") + if not isinstance(python, str): + raise TypeError("udf python must be a string") + if python == "": + raise ValueError("udf python must be a non-empty string") + package_snapshot = _validate_packages(packages) + if not isinstance(output_nullable, bool): + raise TypeError("udf output_nullable must be a bool") + + config = _UdfConfig( + inputs=input_snapshot, + output=output, + output_nullable=output_nullable, + python=python, + packages=package_snapshot, + ) + + def decorator(fn: Callable[_P, _R]) -> Callable[_P, _R]: + if not inspect.isfunction(fn): + raise TypeError("udf can only decorate a Python function") + if hasattr(fn, _CONFIG_ATTR): + raise ValueError("function is already decorated with @udf") + setattr(fn, _CONFIG_ATTR, config) + return fn + + return decorator + + +def _get_udf_config(fn: object) -> _UdfConfig: + """Return the private declaration snapshot for a ``@udf``-decorated function.""" + config = getattr(fn, _CONFIG_ATTR, None) + if config is None: + raise TypeError("function is not decorated with @udf") + if not isinstance(config, _UdfConfig): + raise TypeError("function is not decorated with @udf") + return config diff --git a/python/python/tests/test_first_class_udf.py b/python/python/tests/test_first_class_udf.py new file mode 100644 index 000000000..dcb0026ad --- /dev/null +++ b/python/python/tests/test_first_class_udf.py @@ -0,0 +1,291 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright The LanceDB Authors + +"""RED contract tests for the local @udf declaration surface.""" + +from __future__ import annotations + +import importlib +import inspect +import types + +import pyarrow as pa +import pytest + +import lancedb +from lancedb import Function, Job, udf +from lancedb._udf import _get_udf_config + +_REMOVED_AUTHORING_KNOBS = ( + "user_version", + "idempotency_key", + "deterministic", + "null_handling", + "null_policy", + "on_error", + "error_policy", + "FunctionVersion", + "artifact", + "digest", + "geneva", +) + + +def _decorate(fn, **overrides): + kwargs = { + "inputs": {"x": pa.int32()}, + "output": pa.int64(), + "python": "3.12", + } + kwargs.update(overrides) + return udf(**kwargs)(fn) + + +def test_udf_top_level_export_and_identity_metadata_behavior(): + assert "udf" in lancedb.__all__ + assert udf is lancedb.udf + assert isinstance(importlib.import_module("lancedb._udf"), types.ModuleType) + assert not isinstance(lancedb.udf, types.ModuleType) + + def add(x, y=1): + """Add locally.""" + return x + y + + original = add + decorated = _decorate( + add, + inputs={"x": pa.int32(), "y": pa.int32()}, + output=pa.int32(), + ) + + assert decorated is original + assert decorated.__name__ == "add" + assert decorated.__doc__ == "Add locally." + assert str(inspect.signature(decorated)) == "(x, y=1)" + assert decorated(2) == 3 + assert decorated(2, 5) == 7 + assert decorated(x=4, y=6) == 10 + + +def test_udf_config_snapshot_order_defaults_and_immutability(): + inputs = {"z": pa.string(), "a": pa.int32()} + packages = ["pkg-b==2", "pkg-a==1"] + + def combine(z, a): + return f"{z}:{a}" + + decorated = udf( + inputs=inputs, + output=pa.string(), + python="3.11", + packages=packages, + output_nullable=False, + )(combine) + + config = _get_udf_config(decorated) + assert config.inputs == (("z", pa.string()), ("a", pa.int32())) + assert isinstance(config.inputs, tuple) + assert config.output == pa.string() + assert config.output_nullable is False + assert config.python == "3.11" + assert config.packages == ("pkg-b==2", "pkg-a==1") + assert isinstance(config.packages, tuple) + + inputs["extra"] = pa.bool_() + del inputs["z"] + packages.append("pkg-c==3") + packages[0] = "mutated==0" + assert config.inputs == (("z", pa.string()), ("a", pa.int32())) + assert config.packages == ("pkg-b==2", "pkg-a==1") + + for attr in ("inputs", "output", "output_nullable", "python", "packages"): + with pytest.raises(AttributeError): + setattr(config, attr, None) + + def defaults_only(x): + return x + + defaulted = udf( + inputs={"x": pa.int32()}, + output=pa.int64(), + python="3.12", + )(defaults_only) + default_config = _get_udf_config(defaulted) + assert default_config.packages == () + assert default_config.output_nullable is True + + +def test_udf_accepts_lambda_and_closure_for_local_declaration(): + ambient = "ambient-secret-value-xyz" + + lam = udf( + inputs={"n": pa.int32()}, + output=pa.int32(), + python="3.12", + )(lambda n: n + 1) + assert lam(3) == 4 + assert _get_udf_config(lam).inputs == (("n", pa.int32()),) + + def factory(offset): + @udf( + inputs={"n": pa.int32()}, + output=pa.int32(), + python="3.12", + packages=["demo==0.1"], + ) + def closed(n): + return n + offset + len(ambient) + + return closed + + closed = factory(10) + assert closed(2) == 12 + len(ambient) + assert _get_udf_config(closed).packages == ("demo==0.1",) + + +def test_udf_declaration_defers_signature_and_implementation_packaging(): + """Declaration must not validate callable signature or embed implementation.""" + + def local_add(left, right=1): + return left + right + + decorated = udf( + inputs={"x": pa.int32(), "y": pa.int32()}, + output=pa.int32(), + python="3.12", + )(local_add) + + assert decorated is local_add + assert str(inspect.signature(decorated)) == "(left, right=1)" + assert decorated(2) == 3 + assert decorated(2, 5) == 7 + + config = _get_udf_config(decorated) + assert config.inputs == (("x", pa.int32()), ("y", pa.int32())) + for attr in ( + "source", + "module", + "callable", + "function", + "implementation", + "bundle", + "artifact", + "digest", + ): + assert not hasattr(config, attr) + + +def test_udf_lookup_double_decoration_and_non_function_target(): + def plain(x): + return x + + with pytest.raises((TypeError, ValueError)): + _get_udf_config(plain) + + decorated = _decorate(plain) + + with pytest.raises((TypeError, ValueError)): + _decorate(decorated) + + with pytest.raises(TypeError): + udf( + inputs={"x": pa.int32()}, + output=pa.int32(), + python="3.12", + )(object()) + + with pytest.raises(TypeError): + udf( + inputs={"x": pa.int32()}, + output=pa.int32(), + python="3.12", + )(42) + + +def test_udf_config_validation_errors(): + def target(x): + return x + + with pytest.raises(TypeError): + udf({"x": pa.int32()}, pa.int32(), "3.12")(target) + + with pytest.raises(TypeError): + _decorate(target, inputs=[("x", pa.int32())]) + + with pytest.raises(TypeError): + _decorate(target, inputs={1: pa.int32()}) + + with pytest.raises(ValueError): + _decorate(target, inputs={"": pa.int32()}) + + with pytest.raises(TypeError): + _decorate(target, inputs={"x": "int32"}) + + with pytest.raises(TypeError): + _decorate(target, output="int64") + + with pytest.raises(TypeError): + _decorate(target, python=3.12) + + with pytest.raises(ValueError): + _decorate(target, python="") + + with pytest.raises(TypeError): + _decorate(target, packages="pkg==1") + + with pytest.raises(ValueError): + _decorate(target, packages=["pkg==1", ""]) + + with pytest.raises(ValueError): + _decorate(target, packages=["pkg==1", "pkg==1"]) + + with pytest.raises(TypeError): + _decorate(target, packages=["pkg==1", 2]) + + with pytest.raises(TypeError): + _decorate(target, output_nullable=1) + + with pytest.raises(TypeError): + _decorate(target, output_nullable="true") + + +def test_udf_rejects_removed_overdesign_and_has_no_durable_side_effects(): + params = inspect.signature(udf).parameters + for name in _REMOVED_AUTHORING_KNOBS: + assert name not in params + + def score(x): + """score body marker unique-xyz.""" + ambient = "ambient-secret-value-xyz" + return f"{ambient}:{x}" + + decorated = _decorate( + score, + packages=["score==1.0"], + output_nullable=True, + ) + config = _get_udf_config(decorated) + text = repr(config).lower() + + assert "score body marker unique-xyz" not in text + assert "ambient-secret-value-xyz" not in text + for token in ( + "user_version", + "idempotency_key", + "deterministic", + "null_policy", + "on_error", + "functionversion", + "artifact", + "digest", + "geneva", + ): + assert token not in text + + for attr in _REMOVED_AUTHORING_KNOBS: + assert not hasattr(config, attr) + + assert not isinstance(decorated, Function) + assert not isinstance(decorated, Job) + for attr in ("id", "function_id", "job", "job_id", "registration"): + assert not hasattr(decorated, attr)