Files
lancedb/python/python/tests/test_first_class_function_slice1.py
T
Xuanwo b944055b63 feat(python): class Functions with initialization and shipped helper code (#4254)
Stacked on #4253.

Remote Python Functions could not share helper code, take typed
parameters, or keep per-process state: helpers became `from <module>
import ...` lines a worker cannot resolve, every parameter was an `env=`
string, and state had to be stashed on imported modules. MMLB's OpenAI
preset shows the cost — 24 env vars, a 302-line callable, and a rate
limiter duplicated between the scalar and batch variants, one of which
never constructed it (ENT-2516).

**Class Functions.** `@udf` accepts a class. Each remote instance runs
`__init__` once, calls `__call__` for every row or batch it processes
(row or batch mode is inferred from annotations as before), and calls
`close()` once if defined. Calling the definition locally constructs the
class, so unit tests stay ordinary.

**Initialization.** The annotated `__init__` parameters become the
Function's initialization fields. A binding passes their values beside
its column inputs: `fn(text=col("body"), model="small",
dimensions=512)`. Values are constants of that binding; another column
can bind the same Function version with other values. Types are limited
to booleans, integers, floats, strings, and lists or structs of them,
which have one unambiguous JSON and SQL spelling. A parameter with a
default may be omitted; a null value then takes the default.
Initialization and input names must be disjoint because both are keyword
arguments of one call. Secrets stay on `EnvVarSecret`.

**Helper code.** `code=[package, ...]` ships top-level modules or
packages as Python source in a `python_bundle` artifact (a canonical
JSON object of path → source), imported normally on the worker. Changing
a helper changes the artifact digest and so the Function version; the
environment from `pip`/`conda` is reused. Functions and classes defined
in `__main__` (notebook cells, scripts) are packaged by source,
recursively and dependencies first. An import of a module that lives in
a local source tree, is not shipped with `code=`, and names no declared
package is now rejected at registration instead of failing on the
worker. Closures, lambdas, and nested definitions are still rejected,
with the reason.

A Function without these features packages byte-identically to before
(`python_callable`, same digest, no `initialization` on the wire). Class
sources are located through a method's code object, because
`inspect.getsource` cannot find a class defined in a notebook cell or
doctest.

The packaging contract is documented on `udf`. Sophon changes that build
and execute these artifacts are in lancedb/sophon; they were verified
end to end on a Linux local server (registration → REST, SQL, and
materialized-view bindings with different initialization → refresh →
query, one construction per instance).
2026-09-24 00:56:52 +08:00

463 lines
16 KiB
Python

# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
import json
from pathlib import Path
import pytest
from lancedb import col
import lancedb.functions as functions
from lancedb.functions import (
FunctionApplication,
FunctionBinding,
FunctionVersion,
PythonRuntimeSpec,
SecretBinding,
RefreshColumnResult,
SecretReference,
)
from lancedb.table import AsyncTable
FIXTURES = (
Path(__file__).parents[3]
/ "rust"
/ "lancedb"
/ "tests"
/ "fixtures"
/ "first_class_functions"
/ "v1"
)
def fixture(name: str) -> str:
return (FIXTURES / name).read_text()
def job_result(name: str) -> dict:
return json.loads(fixture(name))["result"]
def assert_no_secret_values(value):
"""No client value models a resolved credential, at any nesting depth."""
if isinstance(value, dict):
for key, child in value.items():
assert key not in {
"secret_value",
"secret_values",
"resolved_secret",
"resolved_secrets",
}
assert_no_secret_values(child)
elif isinstance(value, list):
for child in value:
assert_no_secret_values(child)
def test_public_function_values_are_in_api_reference():
docs = Path(__file__).parents[3] / "docs" / "src" / "python" / "python.md"
rendered = docs.read_text()
for name in functions.__all__:
assert f"::: lancedb.functions.{name}" in rendered
@pytest.mark.parametrize(
("fixture_name", "canonical_name", "model", "nested_result"),
[
(
"remote_function_job.json",
"remote_function_version.canonical.json",
FunctionVersion,
True,
),
(
"remote_function_application.json",
"remote_function_application.canonical.json",
FunctionApplication,
False,
),
(
"remote_function_binding.json",
"remote_function_binding.canonical.json",
FunctionBinding,
False,
),
(
"remote_initialized_function_application.json",
"remote_initialized_function_application.canonical.json",
FunctionApplication,
False,
),
(
"remote_initialized_function_binding.json",
"remote_initialized_function_binding.canonical.json",
FunctionBinding,
False,
),
(
"remote_refresh_job.json",
"remote_refresh_result.canonical.json",
RefreshColumnResult,
True,
),
(
"remote_refresh_result_without_published_version.json",
"remote_refresh_result_without_published_version.canonical.json",
RefreshColumnResult,
False,
),
],
)
def test_python_and_rust_share_remote_canonical_goldens(
fixture_name, canonical_name, model, nested_result
):
value = json.loads(fixture(fixture_name))
if nested_result:
value = value["result"]
decoded = model.from_json(json.dumps(value))
assert decoded.to_canonical_json() == fixture(canonical_name).strip()
def test_function_version_identity_is_immutable_and_exact():
value = job_result("remote_function_job.json")
version = FunctionVersion.from_json(json.dumps(value))
assert version.name == "embed"
assert version.version == "1"
assert version.image.manifest_digest.startswith("sha256:")
assert version.version != version.image.manifest_digest
assert list(version.secret_bindings) == [
SecretBinding(
kind="env", variable="HF_TOKEN", secret_ref=SecretReference(name="hf-prod")
)
]
with pytest.raises((TypeError, ValueError)):
version.version = "1"
with pytest.raises(TypeError, match="immutable"):
version.image.descriptor["format_version"] = "changed"
changed = dict(value)
changed["version"] = "2"
assert FunctionVersion(**changed) != version
assert FunctionVersion(**changed).image == version.image
for invalid in [version.image.manifest_digest, "0", "01", "-1", str(2**64)]:
with pytest.raises(ValueError):
FunctionVersion(**{**value, "version": invalid})
def test_function_version_binds_named_columns_as_one_immutable_application():
version = FunctionVersion.from_json(
json.dumps(job_result("remote_function_job.json"))
)
application = version(text=col("documents.body"))
assert application.function.name == version.name
assert application.function.version == version.version
assert application.output is version.signature.output
assert [
(value.parameter, value.kind, value.value["path"])
for value in application.inputs
] == [("text", "column", "documents.body")]
def test_function_version_binding_validates_names_and_direct_columns():
version = FunctionVersion.from_json(
json.dumps(job_result("remote_function_job.json"))
)
with pytest.raises(TypeError, match=r"missing inputs: \['text'\]"):
version()
with pytest.raises(TypeError, match=r"unknown arguments: \['body'\]"):
version(text=col("text"), body=col("body"))
with pytest.raises(TypeError, match="direct col"):
version(text=col("text").lower())
def test_function_version_keeps_named_struct_outputs_in_one_application():
value = job_result("remote_function_job.json")
value["name"] = "text_features"
value["version"] = "1"
value["signature"] = {
"inputs": [
{"name": "title", "arrow_type": "utf8", "nullable": True},
{"name": "body", "arrow_type": "utf8", "nullable": True},
],
"output": {
"kind": "named_struct",
"fields": [
{
"name": "normalized_text",
"arrow_type": "utf8",
"nullable": False,
},
{
"name": "token_count",
"arrow_type": "int64",
"nullable": False,
},
],
},
}
version = FunctionVersion(**value)
application = version(body=col("body"), title=col("title")).rename(
columns={
"normalized_text": "search_text",
"token_count": "search_token_count",
}
)
assert [value.parameter for value in application.inputs] == ["title", "body"]
assert [field.name for field in application.output.fields] == [
"normalized_text",
"token_count",
]
assert dict(application.columns) == {
"normalized_text": "search_text",
"token_count": "search_token_count",
}
def test_unknown_fields_and_discriminators_are_forward_decodable():
value = job_result("remote_function_job.json")
value["future_version_metadata"] = {"retention_class": "catalog"}
value["image"]["descriptor"]["future_interface"] = {"kind": "wasm"}
value["signature"]["output"]["kind"] = "future_output_shape"
version = FunctionVersion.from_json(json.dumps(value))
assert version.image.descriptor["future_interface"] == {"kind": "wasm"}
assert json.loads(version.to_canonical_json())["image"]["descriptor"][
"future_interface"
] == {"kind": "wasm"}
assert version.signature.output.kind == "future_output_shape"
def test_function_application_uses_rename_columns_only():
application = FunctionApplication.from_json(
fixture("remote_function_application.json")
)
renamed = application.rename(columns={"normalized_text": "body_normalized"})
assert application.columns["normalized_text"] == "search_text"
assert renamed.columns["normalized_text"] == "body_normalized"
assert renamed.function == application.function
assert not hasattr(application, "rename_outputs")
with pytest.raises(TypeError, match="immutable"):
renamed.columns["normalized_text"] = "changed"
with pytest.raises(TypeError, match="immutable"):
application.inputs[0].value["path"] = "changed"
with pytest.raises(ValueError, match="unknown Function result fields"):
application.rename(columns={"missing": "search_text"})
with pytest.raises(ValueError, match="destinations must be unique"):
application.rename(columns={"normalized_text": "same", "token_count": "same"})
bare_value = json.loads(fixture("remote_function_application.json"))
bare_value.pop("columns")
bare = FunctionApplication(**bare_value)
with pytest.raises(ValueError, match="destinations must be unique"):
bare.rename(columns={"normalized_text": "token_count"})
def test_binding_and_refresh_result_keep_stable_remote_fields():
binding = FunctionBinding.from_json(fixture("remote_function_binding.json"))
assert binding.function.version == "1"
assert [output.output_ordinal for output in binding.outputs] == [0, 1]
assert binding.input_schema is not None
assert binding.output_schema is not None
result = RefreshColumnResult.from_json(
json.dumps(job_result("remote_refresh_job.json"))
)
assert result.rows_filled == result.rows_assigned
assert result.version == result.published_version
result = RefreshColumnResult.from_json(
fixture("remote_refresh_result_without_published_version.json")
)
assert result.published_version is None
assert RefreshColumnResult.from_json(result.to_canonical_json()) == result
def test_initialization_travels_on_the_application_and_binding():
application = FunctionApplication.from_json(
fixture("remote_initialized_function_application.json")
)
assert application.initialization["temperature"] == 0.25
assert application.initialization["retry"]["backoff"] == (1, 2, 4)
application._ensure_declarable()
binding = FunctionBinding.from_json(
fixture("remote_initialized_function_binding.json")
)
assert binding.initialization_row() == {
"model": "text-embedding-3-small",
"dimensions": 512,
}
assert (
FunctionBinding.from_json(
fixture("remote_function_binding.json")
).initialization_row()
is None
)
signature = FunctionVersion.from_json(
json.dumps(job_result("remote_function_job.json"))
).signature
assert signature.initialization == ()
assert "initialization" not in json.loads(signature.to_canonical_json())
def test_function_literal_numeric_domain_matches_rust():
with pytest.raises(ValueError, match="floating-point Function literals"):
FunctionApplication.from_json(fixture("remote_function_application_float.json"))
value = json.loads(fixture("remote_function_application_float.json"))
value["inputs"][0]["value"] = 2**64
with pytest.raises(ValueError, match="outside the canonical JSON range"):
FunctionApplication.from_json(json.dumps(value))
def test_empty_default_maps_have_stable_canonical_bytes():
runtime = PythonRuntimeSpec(
kind="python", python_version="3.12", environment={"kind": "pip"}
)
assert runtime.to_canonical_json() == (
'{"environment":{"kind":"pip"},"kind":"python","python_version":"3.12"}'
)
value = json.loads(fixture("remote_function_application.json"))
value.pop("columns")
application = FunctionApplication.from_json(json.dumps(value))
assert "columns" not in json.loads(application.to_canonical_json())
@pytest.mark.parametrize("field", ["rows_assigned", "source_version"])
def test_refresh_result_rejects_non_u64_values(field):
value = job_result("remote_refresh_job.json")
value[field] = -1
with pytest.raises(ValueError):
RefreshColumnResult.from_json(json.dumps(value))
value[field] = "1"
with pytest.raises(ValueError):
RefreshColumnResult.from_json(json.dumps(value))
def test_canonical_client_values_carry_bindings_and_no_credentials():
"""A binding names a Secret; the credential behind it has no client field."""
version = FunctionVersion.from_json(
json.dumps(job_result("remote_function_job.json"))
)
canonical = json.loads(version.to_canonical_json())
assert canonical["secret_bindings"] == [
{"kind": "env", "variable": "HF_TOKEN", "secret_ref": {"name": "hf-prod"}}
]
assert_no_secret_values(canonical)
def test_a_version_without_bindings_omits_the_field_in_both_directions():
"""A Function that binds nothing carries no ``secret_bindings`` key.
Absent decodes as an empty list, and an empty list serializes back to
absent.
"""
value = job_result("remote_function_job.json")
del value["secret_bindings"]
version = FunctionVersion.from_json(json.dumps(value))
assert list(version.secret_bindings) == []
assert "secret_bindings" not in json.loads(version.to_canonical_json())
class _FunctionDeclarationInner:
def __init__(self):
self.calls = []
async def add_function_columns(self, application_json, output_name):
self.calls.append((json.loads(application_json), output_name))
return "declared"
def known_application() -> FunctionApplication:
value = json.loads(fixture("remote_function_application.json"))
value.pop("future_application")
return FunctionApplication(**value)
@pytest.mark.asyncio
async def test_add_columns_routes_struct_as_one_and_multi_output_binding_atomically():
inner = _FunctionDeclarationInner()
table = AsyncTable(inner)
application = known_application()
result = await table.add_columns(
{"features": application._copy(update={"columns": {}})}
)
assert result == "declared"
assert inner.calls[-1][1] == "features"
bare = application._copy(update={"columns": {}}).rename(
columns={"normalized_text": "search_text"}
)
result = await table.add_columns(bare)
assert result == "declared"
assert inner.calls[-1][1] is None
assert inner.calls[-1][0]["columns"] == {"normalized_text": "search_text"}
@pytest.mark.asyncio
async def test_add_columns_rejects_multiple_bindings_and_unknown_newer_application():
inner = _FunctionDeclarationInner()
table = AsyncTable(inner)
application = known_application()
with pytest.raises(ValueError, match="exactly one Function binding"):
await table.add_columns({"a": application, "b": application})
future = json.loads(fixture("remote_function_application.json"))
application = FunctionApplication(**future)
with pytest.raises(ValueError, match="newer contract"):
await table.add_columns(application)
future.pop("future_application")
future["output"]["assignment"] = "cell_flag"
application = FunctionApplication(**future)
assert "assignment" not in json.loads(application.to_canonical_json())["output"]
with pytest.raises(ValueError, match="output.assignment"):
await table.add_columns(application)
assert inner.calls == []
def test_rename_requires_named_struct_and_keeps_partial_mapping_immutable():
scalar = FunctionApplication.from_json(
json.dumps(
{
"function": {
"name": "embed",
"version": "1",
"object_id": "fixture",
"location": "memory:///fixture",
"manifest_digest": (
"sha256:"
"7e22f815b6648e14f093a3979a8e5a2082fa773ebe1ec84b135cae7e84d6f8e6"
),
},
"inputs": [],
"output": {
"kind": "scalar",
"arrow_type": "list<float32>",
"nullable": False,
},
}
)
)
with pytest.raises(ValueError, match="named-struct"):
scalar.rename(columns={"value": "embedding"})
application = known_application()._copy(update={"columns": {}})
renamed = application.rename(columns={"normalized_text": "search_text"})
assert dict(application.columns) == {}
assert dict(renamed.columns) == {"normalized_text": "search_text"}