feat(python): add conditional function replacement

This commit is contained in:
Xuanwo
2026-08-12 07:02:38 +08:00
parent 29be3e5509
commit 1524ee0669
7 changed files with 652 additions and 1 deletions
+18
View File
@@ -37,6 +37,16 @@ class _SyncFunctions:
native_job = self._connection._submit_register_function(name, definition)
return Job(AsyncJob(native_job))
def replace(
self, name: str, current: Function, decorated_udf: Callable[..., object]
) -> Job:
"""Conditionally replace a Function; return sync [Job][lancedb.job.Job]."""
definition = _udf._build_function_definition(decorated_udf)
native_job = self._connection._submit_replace_function(
name, current, definition
)
return Job(AsyncJob(native_job))
def get(self, name: str) -> Function:
"""Return the Function currently bound to a database-scoped name."""
return self._connection._lookup_function_by_name(name)
@@ -65,6 +75,14 @@ class _AsyncFunctions:
native_job = await self._connection._register_function(name, definition)
return AsyncJob(native_job)
async def replace(
self, name: str, current: Function, decorated_udf: Callable[..., object]
) -> AsyncJob:
"""Conditionally replace a Function; return [AsyncJob][lancedb.job.AsyncJob]."""
definition = _udf._build_function_definition(decorated_udf)
native_job = await self._connection._replace_function(name, current, definition)
return AsyncJob(native_job)
async def get(self, name: str) -> Function:
"""Return the Function currently bound to a database-scoped name."""
return await self._connection._lookup_function_by_name(name)
+3
View File
@@ -156,6 +156,9 @@ class Connection(object):
async def _register_function(
self, name: str, definition: "_FunctionDefinition"
) -> Job: ...
async def _replace_function(
self, name: str, current: Function, definition: "_FunctionDefinition"
) -> Job: ...
async def _lookup_function_by_name(self, name: str) -> Function: ...
async def _lookup_function_by_id(self, function_id: str) -> Function: ...
async def create_table(
+22
View File
@@ -672,6 +672,17 @@ class DBConnection(EnforceOverrides):
"function registration is not supported for this connection type"
)
def _submit_replace_function(
self, name: str, current: "Function", definition: "_FunctionDefinition"
) -> NativeJob:
"""Submit a Function conditional replace job via the native connection.
Connection subclasses that support registration override this hook.
"""
raise NotImplementedError(
"function replace is not supported for this connection type"
)
def _lookup_function_by_name(self, name: str) -> "Function":
"""Look up a Function by database-scoped name via the native connection.
@@ -1313,6 +1324,12 @@ class LanceDBConnection(DBConnection):
) -> NativeJob:
return LOOP.run(self._conn._register_function(name, definition))
@override
def _submit_replace_function(
self, name: str, current: "Function", definition: "_FunctionDefinition"
) -> NativeJob:
return LOOP.run(self._conn._replace_function(name, current, definition))
@override
def _lookup_function_by_name(self, name: str) -> "Function":
return LOOP.run(self._conn._lookup_function_by_name(name))
@@ -2079,6 +2096,11 @@ class AsyncConnection(object):
) -> NativeJob:
return await self._inner._register_function(name, definition)
async def _replace_function(
self, name: str, current: "Function", definition: "_FunctionDefinition"
) -> NativeJob:
return await self._inner._replace_function(name, current, definition)
async def _lookup_function_by_name(self, name: str) -> "Function":
return await self._inner._lookup_function_by_name(name)
+6
View File
@@ -742,6 +742,12 @@ class RemoteDBConnection(DBConnection):
) -> "NativeJob":
return LOOP.run(self._conn._register_function(name, definition))
@override
def _submit_replace_function(
self, name: str, current: "Function", definition: "_FunctionDefinition"
) -> "NativeJob":
return LOOP.run(self._conn._replace_function(name, current, definition))
@override
def _lookup_function_by_name(self, name: str) -> "Function":
return LOOP.run(self._conn._lookup_function_by_name(name))
@@ -618,7 +618,6 @@ def test_no_direct_db_lookup_methods_and_no_deleted_keywords():
assert not hasattr(db, "get_function")
assert not hasattr(db.functions, "get_by_name")
assert not hasattr(db.functions, "list")
assert not hasattr(db.functions, "replace")
assert not hasattr(db.functions, "remove")
assert not hasattr(db.functions, "revoke")
@@ -0,0 +1,579 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
"""Contract tests for Python conditional first-class Function replacement."""
from __future__ import annotations
import contextlib
import http.server
import json
import threading
from collections.abc import AsyncIterator, Iterator
from datetime import timedelta
from typing import Any, Callable
from unittest import mock
import pyarrow as pa
import pytest
import lancedb
import lancedb._udf as _udf_mod
import lancedb.job
from lancedb import FunctionCapability, udf
from lancedb.exceptions import JobFailedError
_SOURCE_MARKER = "replace-source-marker-unique-xyz"
_SECRET_REFERENCE = "secret://team/replace-redact-token-xyz"
_SECRET_ENV = "REPLACE_API_TOKEN"
_NETWORK_ORIGIN = "https://api.replace-example.com"
_FUNCTION_NAME = "text.normalize"
_CURRENT_FUNCTION_ID = "fn.replace-current-1"
_REPLACED_FUNCTION_ID = "fn.replace-result-1"
_JOB_ID_SYNC = "job-replace-sync-1"
_JOB_ID_ASYNC = "job-replace-async-1"
_JOB_ID_CONFLICT = "job-replace-conflict-1"
_REGISTER_PATH = "/v1/functions/register"
_LOOKUP_PATH = "/v1/functions/lookup"
_DESCRIBE_PATH = "/v1/jobs/describe"
_CONFLICTING_MESSAGE_CODE = "definition_validation_failure"
# Pinned Rust-canonical schema-only type IPC (base64). Same fixtures as lookup /
# job-result tests: PyArrow FileWriter bytes are not byte-identical to Arrow Rust.
_INT32_TYPE_IPC_B64 = (
"QVJST1cxAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAP"
"////94AAAAEAAAAAAACgAMAAoACQAEAAoAAAAQAAAAAAEEAAgACAAAAAQACAAAAAQAAAABAAAAFAAAABAAFAAQ"
"AA4ADwAEAAAACAAQAAAAGAAAACAAAAAAAAECHAAAAAgADAAEAAsACAAAACAAAAAAAAABAAAAAAAAAAAAAAAA/"
"////wAAAAAUAAAAAAAAAAwAFAASAAwACAAEAAwAAABsAAAAcAAAABAAAAAAAAQACAAIAAAABAAIAAAABAAAAA"
"EAAAAUAAAAEAAUABAADgAPAAQAAAAIABAAAAAYAAAAIAAAAAAAAQIcAAAACAAMAAQACwAIAAAAIAAAAAAAAAE"
"AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAACQAAAAQVJST1cx"
)
_UTF8_TYPE_IPC_B64 = (
"QVJST1cxAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAP"
"////94AAAAEAAAAAAACgAMAAoACQAEAAoAAAAQAAAAAAEEAAgACAAAAAQACAAAAAQAAAABAAAAFAAAABAAFAAQ"
"AA4ADwAEAAAACAAQAAAAGAAAAAwAAAAAAAEFEAAAAAAAAAAEAAQABAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA/"
"////wAAAAAQAAAADAAUABIADAAIAAQADAAAAGAAAABkAAAAEAAAAAAABAAIAAgAAAAEAAgAAAAEAAAAAQAAAB"
"QAAAAQABQAEAAOAA8ABAAAAAgAEAAAABgAAAAMAAAAAAABBRAAAAAAAAAABAAEAAQAAAAAAAAAAAAAAAAAAAA"
"AAAAAAAAAAIAAAABBUlJPVzE="
)
_DELETED_REPLACE_KEYWORDS = (
"idempotency_key",
"retry_key",
"user_version",
"version",
"deterministic",
"null_policy",
"replace",
"expected_current_function_id",
"alias",
"lineage",
)
_SPEC_KEYS = {
"format_version",
"name",
"definition",
"expected_current_function_id",
}
@udf(
inputs={"text": pa.string(), "limit": pa.int32()},
output=pa.string(),
python="3.12",
packages=["pkg-a==1"],
output_nullable=True,
capabilities=[
FunctionCapability.network(_NETWORK_ORIGIN),
FunctionCapability.secret(
_SECRET_REFERENCE,
environment_variable=_SECRET_ENV,
),
],
)
def packable_replace_normalize(text, limit):
"""replace-source-marker-unique-xyz."""
return text[:limit]
def _current_function_wire() -> dict[str, Any]:
return {
"format_version": 1,
"id": _CURRENT_FUNCTION_ID,
"signature": {
"parameters": [
{"name": "text", "data_type_ipc": _UTF8_TYPE_IPC_B64},
{"name": "limit", "data_type_ipc": _INT32_TYPE_IPC_B64},
],
"output": {
"data_type_ipc": _UTF8_TYPE_IPC_B64,
"nullable": True,
},
},
}
def _definition_json(fn: object) -> dict[str, Any]:
payload = _udf_mod._build_function_definition(fn)._to_json()
if isinstance(payload, bytes):
return json.loads(payload.decode("utf-8"))
assert isinstance(payload, str)
return json.loads(payload)
def _expected_replace_spec(name: str, current_id: str, fn: object) -> dict[str, Any]:
return {
"format_version": 1,
"name": name,
"definition": _definition_json(fn),
"expected_current_function_id": current_id,
}
def _read_body(request: http.server.BaseHTTPRequestHandler) -> bytes:
content_len = int(request.headers.get("Content-Length", 0))
if content_len <= 0:
return b""
return request.rfile.read(content_len)
def _make_handler(handler: Callable[[http.server.BaseHTTPRequestHandler], None]):
class _Handler(http.server.BaseHTTPRequestHandler):
def do_GET(self):
handler(self)
def do_POST(self):
handler(self)
def log_message(self, format, *args): # noqa: A003
return
return _Handler
@contextlib.contextmanager
def _mock_remote_db(handler) -> Iterator[Any]:
server = http.server.HTTPServer(("localhost", 0), _make_handler(handler))
port = server.server_address[1]
thread = threading.Thread(target=server.serve_forever)
thread.start()
try:
db = lancedb.connect(
"db://dev",
api_key="fake",
host_override=f"http://localhost:{port}",
client_config={
"retry_config": {
"retries": 2,
"backoff_factor": 0.0,
"backoff_jitter": 0.0,
},
"timeout_config": {"connect_timeout": 1},
},
)
yield db
finally:
server.shutdown()
thread.join()
@contextlib.asynccontextmanager
async def _mock_remote_db_async(handler) -> AsyncIterator[Any]:
server = http.server.HTTPServer(("localhost", 0), _make_handler(handler))
port = server.server_address[1]
thread = threading.Thread(target=server.serve_forever)
thread.start()
try:
db = await lancedb.connect_async(
"db://dev",
api_key="fake",
host_override=f"http://localhost:{port}",
client_config={
"retry_config": {
"retries": 2,
"backoff_factor": 0.0,
"backoff_jitter": 0.0,
},
"timeout_config": {"connect_timeout": 1},
},
)
yield db
finally:
server.shutdown()
thread.join()
def _assert_exact_replace_spec(
body: dict[str, Any], expected: dict[str, Any], current_id: str
) -> None:
assert set(body) == _SPEC_KEYS
assert body == expected
assert body["format_version"] == 1
assert body["expected_current_function_id"] == current_id
assert body["expected_current_function_id"] is not None
assert _SOURCE_MARKER in json.dumps(body["definition"])
assert any(
capability.get("reference") == _SECRET_REFERENCE
for capability in body["definition"]["capabilities"]
)
def _lookup_success_handler(
counters: dict[str, int],
*,
after_lookup: (
Callable[[http.server.BaseHTTPRequestHandler, bytes], None] | None
) = None,
):
"""Serve exact lookup; optionally continue for register/describe."""
def handler(request: http.server.BaseHTTPRequestHandler) -> None:
assert request.command == "POST"
raw = _read_body(request)
if request.path == _LOOKUP_PATH:
counters["lookup"] = counters.get("lookup", 0) + 1
body = json.loads(raw.decode("utf-8"))
assert body == {"name": _FUNCTION_NAME}
request.send_response(200)
request.send_header("Content-Type", "application/json")
request.end_headers()
request.wfile.write(
json.dumps({"function": _current_function_wire()}).encode("utf-8")
)
return
if after_lookup is not None:
after_lookup(request, raw)
return
counters["register"] = counters.get("register", 0) + 1
request.send_response(500)
request.end_headers()
request.wfile.write(b"unexpected register")
return handler
def _observe_current(db) -> lancedb.Function:
current = db.functions.get(_FUNCTION_NAME)
assert type(current) is lancedb.Function
assert current.id == _CURRENT_FUNCTION_ID
assert not hasattr(current, "name")
assert not hasattr(current, "replace")
return current
async def _observe_current_async(db) -> lancedb.Function:
current = await db.functions.get(_FUNCTION_NAME)
assert type(current) is lancedb.Function
assert current.id == _CURRENT_FUNCTION_ID
assert not hasattr(current, "name")
assert not hasattr(current, "replace")
return current
def test_sync_remote_replace_exact_body_one_package_job_and_function_result():
counters: dict[str, int] = {"lookup": 0, "register": 0, "describe": 0}
expected_spec = _expected_replace_spec(
_FUNCTION_NAME, _CURRENT_FUNCTION_ID, packable_replace_normalize
)
register_attempts: list[dict[str, Any]] = []
function_result_wire = {
"kind": "function",
"format_version": 1,
"function": {
"format_version": 1,
"id": _REPLACED_FUNCTION_ID,
"signature": expected_spec["definition"]["signature"],
},
}
def after_lookup(
request: http.server.BaseHTTPRequestHandler, payload: bytes
) -> None:
if request.path == _REGISTER_PATH:
counters["register"] += 1
register_attempts.append(
{
"request_id": request.headers.get("x-request-id"),
"raw": payload,
"body": json.loads(payload.decode("utf-8")),
}
)
request.send_response(200)
request.send_header("Content-Type", "application/json")
request.end_headers()
request.wfile.write(json.dumps({"job_id": _JOB_ID_SYNC}).encode("utf-8"))
return
assert request.path == _DESCRIBE_PATH
counters["describe"] += 1
body = json.loads(payload.decode("utf-8"))
assert body["job_id"] == _JOB_ID_SYNC
request.send_response(200)
request.send_header("Content-Type", "application/json")
request.end_headers()
request.wfile.write(
json.dumps(
{
"job_id": _JOB_ID_SYNC,
"job_state": "DONE",
"job_type": "register_function",
"creation_ms": 1,
"spec": {},
"result": function_result_wire,
}
).encode("utf-8")
)
package_calls = {"n": 0}
original_package = _udf_mod._package_udf
def counting_package(fn: object):
package_calls["n"] += 1
return original_package(fn)
with _mock_remote_db(
_lookup_success_handler(counters, after_lookup=after_lookup)
) as db:
assert not hasattr(db, "replace_function")
current = _observe_current(db)
assert counters["lookup"] == 1
assert counters["register"] == 0
with mock.patch.object(_udf_mod, "_package_udf", side_effect=counting_package):
job = db.functions.replace(
_FUNCTION_NAME, current, packable_replace_normalize
)
assert type(job) is lancedb.job.Job
assert job.id == _JOB_ID_SYNC
waited = job.wait(timeout=timedelta(seconds=5))
assert package_calls["n"] == 1
assert counters["lookup"] == 1
assert counters["register"] == 1
assert counters["describe"] == 1
assert len(register_attempts) == 1
attempt = register_attempts[0]
assert isinstance(attempt["request_id"], str) and attempt["request_id"]
assert attempt["raw"]
_assert_exact_replace_spec(attempt["body"], expected_spec, current.id)
assert type(waited) is lancedb.Function
assert waited.id == _REPLACED_FUNCTION_ID
assert waited.parameters == (("text", pa.string()), ("limit", pa.int32()))
assert waited.output_type == pa.string()
assert waited.output_nullable is True
@pytest.mark.asyncio
async def test_async_remote_replace_exact_body_returns_async_job():
counters: dict[str, int] = {"lookup": 0, "register": 0}
expected_spec = _expected_replace_spec(
_FUNCTION_NAME, _CURRENT_FUNCTION_ID, packable_replace_normalize
)
seen: dict[str, Any] = {}
def after_lookup(request: http.server.BaseHTTPRequestHandler, raw: bytes) -> None:
assert request.path == _REGISTER_PATH
counters["register"] += 1
seen["raw"] = raw
seen["body"] = json.loads(raw.decode("utf-8"))
seen["request_id"] = request.headers.get("x-request-id")
request.send_response(200)
request.send_header("Content-Type", "application/json")
request.end_headers()
request.wfile.write(json.dumps({"job_id": _JOB_ID_ASYNC}).encode("utf-8"))
async with _mock_remote_db_async(
_lookup_success_handler(counters, after_lookup=after_lookup)
) as db:
assert not hasattr(db, "replace_function")
current = await _observe_current_async(db)
assert counters["lookup"] == 1
assert counters["register"] == 0
job = await db.functions.replace(
_FUNCTION_NAME, current, packable_replace_normalize
)
assert counters["lookup"] == 1
assert counters["register"] == 1
assert seen.get("raw")
assert isinstance(seen.get("request_id"), str) and seen["request_id"]
_assert_exact_replace_spec(seen["body"], expected_spec, current.id)
assert type(job) is lancedb.job.AsyncJob
assert job.id == _JOB_ID_ASYNC
def test_sync_remote_replace_failed_name_conflict_raises_job_failed_error_code():
counters: dict[str, int] = {"lookup": 0, "register": 0, "describe": 0}
def after_lookup(request: http.server.BaseHTTPRequestHandler, raw: bytes) -> None:
if request.path == _REGISTER_PATH:
counters["register"] += 1
request.send_response(200)
request.send_header("Content-Type", "application/json")
request.end_headers()
request.wfile.write(
json.dumps({"job_id": _JOB_ID_CONFLICT}).encode("utf-8")
)
return
assert request.path == _DESCRIBE_PATH
counters["describe"] += 1
body = json.loads(raw.decode("utf-8"))
assert body["job_id"] == _JOB_ID_CONFLICT
request.send_response(200)
request.send_header("Content-Type", "application/json")
request.end_headers()
request.wfile.write(
json.dumps(
{
"job_id": _JOB_ID_CONFLICT,
"job_type": "register_function",
"job_state": "FAILED",
"creation_ms": 1,
"spec": {},
"failure": {
"phase": "validate",
"message": (
f"looks like {_CONFLICTING_MESSAGE_CODE} during CAS"
),
"retryable": False,
"error_code": "name_conflict",
},
}
).encode("utf-8")
)
with _mock_remote_db(
_lookup_success_handler(counters, after_lookup=after_lookup)
) as db:
current = _observe_current(db)
assert counters["lookup"] == 1
assert counters["register"] == 0
job = db.functions.replace(_FUNCTION_NAME, current, packable_replace_normalize)
assert type(job) is lancedb.job.Job
with pytest.raises(JobFailedError) as exc_info:
job.wait(timeout=timedelta(seconds=5))
assert counters["lookup"] == 1
assert counters["register"] == 1
assert counters["describe"] == 1
err = exc_info.value
assert isinstance(err, JobFailedError)
assert err.error_code == "name_conflict"
assert err.error_code != _CONFLICTING_MESSAGE_CODE
def test_empty_name_rejects_before_register_transport():
counters: dict[str, int] = {"lookup": 0, "register": 0}
with _mock_remote_db(_lookup_success_handler(counters)) as db:
current = _observe_current(db)
assert counters["lookup"] == 1
assert counters["register"] == 0
with pytest.raises(ValueError):
db.functions.replace("", current, packable_replace_normalize)
assert counters["lookup"] == 1
assert counters["register"] == 0
@pytest.mark.parametrize(
"bad_current",
[
_CURRENT_FUNCTION_ID,
{"id": _CURRENT_FUNCTION_ID},
object(),
123,
],
)
def test_raw_id_or_arbitrary_current_rejected_without_register(bad_current):
counters: dict[str, int] = {"lookup": 0, "register": 0}
with _mock_remote_db(_lookup_success_handler(counters)) as db:
# Observe a real handle separately so the bad-current path is isolated.
_ = _observe_current(db)
assert counters["lookup"] == 1
assert counters["register"] == 0
with pytest.raises(TypeError):
db.functions.replace(
_FUNCTION_NAME, bad_current, packable_replace_normalize
)
assert counters["lookup"] == 1
assert counters["register"] == 0
def test_local_sync_replace_not_implemented_without_table_mutation(tmp_path):
counters: dict[str, int] = {"lookup": 0, "register": 0}
with _mock_remote_db(_lookup_success_handler(counters)) as remote_db:
current = _observe_current(remote_db)
assert counters["lookup"] == 1
assert counters["register"] == 0
assert type(current) is lancedb.Function
db = lancedb.connect(tmp_path)
before = db.list_tables().tables
assert before == []
assert not hasattr(db, "replace_function")
with pytest.raises(NotImplementedError):
db.functions.replace(_FUNCTION_NAME, current, packable_replace_normalize)
assert db.list_tables().tables == before
@pytest.mark.asyncio
async def test_local_async_replace_not_implemented_without_table_mutation(tmp_path):
counters: dict[str, int] = {"lookup": 0, "register": 0}
with _mock_remote_db(_lookup_success_handler(counters)) as remote_db:
current = _observe_current(remote_db)
assert counters["lookup"] == 1
assert counters["register"] == 0
assert type(current) is lancedb.Function
db = await lancedb.connect_async(tmp_path)
before = (await db.list_tables()).tables
assert before == []
assert not hasattr(db, "replace_function")
with pytest.raises(NotImplementedError):
await db.functions.replace(_FUNCTION_NAME, current, packable_replace_normalize)
assert (await db.list_tables()).tables == before
@pytest.mark.parametrize("keyword", _DELETED_REPLACE_KEYWORDS)
def test_replace_rejects_deleted_cas_retry_version_kwargs(keyword):
counters: dict[str, int] = {"lookup": 0, "register": 0}
with _mock_remote_db(_lookup_success_handler(counters)) as db:
current = _observe_current(db)
assert counters["lookup"] == 1
assert counters["register"] == 0
with pytest.raises(TypeError):
db.functions.replace(
_FUNCTION_NAME,
current,
packable_replace_normalize,
**{keyword: True},
)
assert counters["lookup"] == 1
assert counters["register"] == 0
def test_no_direct_replace_function_methods_and_function_has_no_replace():
counters: dict[str, int] = {"lookup": 0, "register": 0}
with _mock_remote_db(_lookup_success_handler(counters)) as db:
current = _observe_current(db)
assert not hasattr(db, "replace_function")
assert not hasattr(db, "register_function")
assert not hasattr(current, "replace")
assert not hasattr(current, "replace_function")
assert callable(getattr(db.functions, "replace", None))
assert counters["lookup"] == 1
assert counters["register"] == 0