|
|
|
@@ -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
|