diff --git a/python/python/lancedb/_functions.py b/python/python/lancedb/_functions.py index be6053b29..fc9d5a892 100644 --- a/python/python/lancedb/_functions.py +++ b/python/python/lancedb/_functions.py @@ -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) diff --git a/python/python/lancedb/_lancedb.pyi b/python/python/lancedb/_lancedb.pyi index 6dabb6848..f7d0ece73 100644 --- a/python/python/lancedb/_lancedb.pyi +++ b/python/python/lancedb/_lancedb.pyi @@ -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( diff --git a/python/python/lancedb/db.py b/python/python/lancedb/db.py index d70eac26e..d7f8854cd 100644 --- a/python/python/lancedb/db.py +++ b/python/python/lancedb/db.py @@ -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) diff --git a/python/python/lancedb/remote/db.py b/python/python/lancedb/remote/db.py index 5bbe48b60..4c157a4d6 100644 --- a/python/python/lancedb/remote/db.py +++ b/python/python/lancedb/remote/db.py @@ -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)) diff --git a/python/python/tests/test_first_class_function_lookup.py b/python/python/tests/test_first_class_function_lookup.py index ece823995..657a52a37 100644 --- a/python/python/tests/test_first_class_function_lookup.py +++ b/python/python/tests/test_first_class_function_lookup.py @@ -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") diff --git a/python/python/tests/test_first_class_function_replace.py b/python/python/tests/test_first_class_function_replace.py new file mode 100644 index 000000000..885d10ca0 --- /dev/null +++ b/python/python/tests/test_first_class_function_replace.py @@ -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 diff --git a/python/src/connection.rs b/python/src/connection.rs index 4d60d0b89..1a1b1b8cc 100644 --- a/python/src/connection.rs +++ b/python/src/connection.rs @@ -610,6 +610,30 @@ impl Connection { }) } + /// Submit a first-class Function conditional replace job. + /// + /// Accepts the observed native [`crate::function::Function`] handle and the + /// exact private [`crate::function::PyFunctionDefinition`], then builds + /// [`RegisterFunctionJobSpec`] with `expected_current_function_id = + /// Some(current.id)`. Reads only `current.inner().id().clone()`. Does not + /// JSON round-trip the definition. + pub fn _replace_function<'py>( + self_: PyRef<'py, Self>, + name: String, + current: Bound<'_, crate::function::Function>, + definition: Bound<'_, crate::function::PyFunctionDefinition>, + ) -> PyResult> { + let inner = self_.get_inner()?.clone(); + let definition = definition.get().inner().clone(); + let current_id = current.get().inner().id().clone(); + future_into_py(self_.py(), async move { + let spec = RegisterFunctionJobSpec::try_new(name, definition, Some(current_id)) + .infer_error()?; + let job = inner.register_function(spec).await.infer_error()?; + Ok(crate::job::Job::new(job)) + }) + } + /// Look up the Function currently bound to a database-scoped name. /// /// Wraps the exact Rust [`lancedb::function::Function`] once. Empty names