mirror of
https://github.com/lancedb/lancedb.git
synced 2026-08-31 02:18:27 +00:00
feat(functions): support UDF secret values
This commit is contained in:
@@ -16,6 +16,7 @@ from typing import (
|
||||
Iterable,
|
||||
List,
|
||||
Literal,
|
||||
Mapping,
|
||||
Optional,
|
||||
Union,
|
||||
)
|
||||
@@ -687,17 +688,35 @@ class DBConnection(EnforceOverrides):
|
||||
"""
|
||||
raise NotImplementedError("serialize is not supported for this connection type")
|
||||
|
||||
def create_function(self, definition: UdfDefinition) -> FunctionVersion:
|
||||
def create_function(
|
||||
self,
|
||||
definition: UdfDefinition,
|
||||
*,
|
||||
secrets: Optional[Mapping[str, str]] = None,
|
||||
) -> FunctionVersion:
|
||||
"""Register a scalar Python UDF and wait for its immutable version.
|
||||
|
||||
``secrets`` must contain exactly the names declared by
|
||||
``@udf(secrets=[...])``. Values are sent in the create request and
|
||||
stored server-side in the private execution artifact; returned
|
||||
Function and Job metadata contain only the declared names.
|
||||
This is the blocking counterpart of :meth:`create_function_async`.
|
||||
Local connections raise ``NotImplementedError``.
|
||||
"""
|
||||
return self.create_function_async(definition).wait()
|
||||
return self.create_function_async(definition, secrets=secrets).wait()
|
||||
|
||||
def create_function_async(self, definition: UdfDefinition) -> Job[FunctionVersion]:
|
||||
def create_function_async(
|
||||
self,
|
||||
definition: UdfDefinition,
|
||||
*,
|
||||
secrets: Optional[Mapping[str, str]] = None,
|
||||
) -> Job[FunctionVersion]:
|
||||
"""Register a scalar Python UDF through the remote Function catalog.
|
||||
|
||||
``secrets`` must contain exactly the names declared by
|
||||
``@udf(secrets=[...])``. Values are sent in the create request and
|
||||
stored server-side in the private execution artifact; returned
|
||||
Function and Job metadata contain only the declared names.
|
||||
Submission returns a typed job. The immutable Function version becomes
|
||||
available only when :meth:`Job.wait` succeeds. Local connections raise
|
||||
``NotImplementedError``.
|
||||
@@ -1405,8 +1424,13 @@ class LanceDBConnection(DBConnection):
|
||||
return Job(self._conn.job(job_id))
|
||||
|
||||
@override
|
||||
def create_function_async(self, definition: UdfDefinition) -> Job[FunctionVersion]:
|
||||
job = LOOP.run(self._conn.create_function_async(definition))
|
||||
def create_function_async(
|
||||
self,
|
||||
definition: UdfDefinition,
|
||||
*,
|
||||
secrets: Optional[Mapping[str, str]] = None,
|
||||
) -> Job[FunctionVersion]:
|
||||
job = LOOP.run(self._conn.create_function_async(definition, secrets=secrets))
|
||||
return Job(job)
|
||||
|
||||
@override
|
||||
@@ -2225,17 +2249,24 @@ class AsyncConnection(object):
|
||||
return AsyncJob(self._inner.job(job_id))
|
||||
|
||||
async def create_function_async(
|
||||
self, definition: UdfDefinition
|
||||
self,
|
||||
definition: UdfDefinition,
|
||||
*,
|
||||
secrets: Optional[Mapping[str, str]] = None,
|
||||
) -> AsyncJob[FunctionVersion]:
|
||||
"""Register a scalar Python UDF through the remote Function catalog.
|
||||
|
||||
``secrets`` must contain exactly the names declared by
|
||||
``@udf(secrets=[...])``. Values are sent in the create request and
|
||||
stored server-side in the private execution artifact; returned
|
||||
Function and Job metadata contain only the declared names.
|
||||
The returned typed job resolves to the immutable Function version.
|
||||
Local connections raise ``NotImplementedError``.
|
||||
"""
|
||||
if not isinstance(definition, UdfDefinition):
|
||||
raise TypeError("create_function_async requires a @udf definition")
|
||||
inner = await self._inner.create_function_async(
|
||||
definition.registration_request.to_canonical_json()
|
||||
definition._submission_json(secrets)
|
||||
)
|
||||
return _typed_job(inner, FunctionVersion.from_json)
|
||||
|
||||
|
||||
@@ -4,7 +4,7 @@
|
||||
"""Canonical Function values exchanged with LanceDB Enterprise services.
|
||||
|
||||
These immutable models contain client/wire state only. Catalog persistence,
|
||||
environment bake, and execution are owned by Sophon.
|
||||
environment bake, secret resolution, and execution are owned by Sophon.
|
||||
``RefreshColumnResult`` is also the backend-neutral result of a local
|
||||
expression-backed refresh job.
|
||||
"""
|
||||
@@ -229,7 +229,7 @@ class PythonEnvironmentSpec(_RemoteValue):
|
||||
|
||||
|
||||
class PythonRuntimeSpec(_RemoteValue):
|
||||
"""Remote runtime definition with environment values.
|
||||
"""Remote runtime definition with non-secret environment values.
|
||||
|
||||
V1 supports ``kind="python"``. Newer runtime kinds remain readable, while
|
||||
their unknown payload fields are intentionally not retained by the client.
|
||||
@@ -268,6 +268,7 @@ class FunctionVersion(_RemoteValue):
|
||||
runtime: PythonRuntimeSpec
|
||||
runtime_digest: str
|
||||
environment_digest: str
|
||||
required_secrets: tuple[str, ...] = ()
|
||||
created_at: str
|
||||
|
||||
def __call__(self, **inputs: Any) -> FunctionApplication:
|
||||
@@ -329,12 +330,17 @@ class FunctionVersion(_RemoteValue):
|
||||
|
||||
|
||||
class FunctionRegistrationRequest(_RemoteValue):
|
||||
"""Stable remote registration envelope produced by :func:`udf`."""
|
||||
"""Stable remote registration envelope produced by :func:`udf`.
|
||||
|
||||
Only secret names are represented. Secret values are supplied separately
|
||||
when the definition is submitted and are not part of this durable value.
|
||||
"""
|
||||
|
||||
name: str
|
||||
artifact: FunctionArtifactRequest
|
||||
signature: FunctionSignature
|
||||
runtime: PythonRuntimeSpec
|
||||
required_secrets: tuple[str, ...] = ()
|
||||
|
||||
|
||||
class FunctionVersionRef(_OpenRemoteValue):
|
||||
@@ -479,6 +485,18 @@ class RefreshColumnResult(_RemoteValue):
|
||||
|
||||
|
||||
_FUNCTION_NAME = re.compile(r"^[A-Za-z_][A-Za-z0-9_.-]*$")
|
||||
_SECRET_NAME = re.compile(r"^[A-Za-z_][A-Za-z0-9_]*$")
|
||||
|
||||
|
||||
def _validate_secret_value(name: str, value: Any) -> str:
|
||||
"""Validate one secret value before building the create request."""
|
||||
if not isinstance(value, str):
|
||||
raise TypeError(f"Function secret {name!r} value must be a string")
|
||||
if not value:
|
||||
raise ValueError(f"Function secret {name!r} value must be non-empty")
|
||||
if "\0" in value:
|
||||
raise ValueError(f"Function secret {name!r} value must not contain NUL")
|
||||
return value
|
||||
|
||||
|
||||
_GRAMMAR_PRIMITIVES = (
|
||||
@@ -909,6 +927,7 @@ class UdfDefinition:
|
||||
output_schema: Optional[pa.DataType | pa.Field | pa.Schema],
|
||||
pip: tuple[str, ...],
|
||||
env: Mapping[str, str],
|
||||
secrets: tuple[str, ...],
|
||||
python_version: Optional[str],
|
||||
conda: tuple[str, ...] = (),
|
||||
conda_channels: tuple[str, ...] = (),
|
||||
@@ -935,6 +954,17 @@ class UdfDefinition:
|
||||
for key, value in environment.items()
|
||||
):
|
||||
raise TypeError("Function env keys and values must be strings")
|
||||
required_secrets = tuple(sorted(set(secrets)))
|
||||
invalid_secrets = [
|
||||
secret for secret in required_secrets if not _SECRET_NAME.fullmatch(secret)
|
||||
]
|
||||
if invalid_secrets:
|
||||
raise ValueError(f"invalid Function secret names: {invalid_secrets!r}")
|
||||
overlap = set(environment) & set(required_secrets)
|
||||
if overlap:
|
||||
raise ValueError(
|
||||
f"Function env and secret names must be disjoint: {sorted(overlap)!r}"
|
||||
)
|
||||
signature = _infer_signature(function, input_schema, output_schema)
|
||||
source = _package_source(function)
|
||||
digest = f"sha256:{hashlib.sha256(source).hexdigest()}"
|
||||
@@ -963,14 +993,57 @@ class UdfDefinition:
|
||||
),
|
||||
signature=signature,
|
||||
runtime=runtime,
|
||||
required_secrets=required_secrets,
|
||||
)
|
||||
functools.update_wrapper(self, function)
|
||||
|
||||
@property
|
||||
def registration_request(self) -> FunctionRegistrationRequest:
|
||||
"""The immutable request sent by ``create_function_async``."""
|
||||
"""The immutable, value-free client model for a Function submission."""
|
||||
return self._request
|
||||
|
||||
def _submission_json(self, secrets: Optional[Mapping[str, str]]) -> str:
|
||||
"""Build one registration submission without retaining values on self."""
|
||||
if secrets is None:
|
||||
secret_values: Mapping[str, str] = {}
|
||||
elif not isinstance(secrets, Mapping):
|
||||
raise TypeError("Function secrets must be a mapping of names to strings")
|
||||
else:
|
||||
secret_values = secrets
|
||||
|
||||
if any(not isinstance(name, str) for name in secret_values):
|
||||
raise TypeError("Function secret names must be strings")
|
||||
expected = set(self._request.required_secrets)
|
||||
provided = set(secret_values)
|
||||
if provided != expected:
|
||||
missing = sorted(expected - provided)
|
||||
unexpected = sorted(provided - expected)
|
||||
details = []
|
||||
if missing:
|
||||
details.append(f"missing: {missing!r}")
|
||||
if unexpected:
|
||||
details.append(f"unexpected: {unexpected!r}")
|
||||
raise ValueError(
|
||||
"Function secret values must exactly match the declared secrets ("
|
||||
+ "; ".join(details)
|
||||
+ ")"
|
||||
)
|
||||
|
||||
canonical_values = {}
|
||||
for name in sorted(secret_values):
|
||||
canonical_values[name] = _validate_secret_value(name, secret_values[name])
|
||||
|
||||
submission = self._request._known_dict()
|
||||
if canonical_values:
|
||||
submission["secret_values"] = canonical_values
|
||||
return json.dumps(
|
||||
submission,
|
||||
ensure_ascii=False,
|
||||
allow_nan=False,
|
||||
sort_keys=True,
|
||||
separators=(",", ":"),
|
||||
)
|
||||
|
||||
def __call__(self, *args, **kwargs):
|
||||
return self._function(*args, **kwargs)
|
||||
|
||||
@@ -988,6 +1061,7 @@ def udf(
|
||||
output_schema: Optional[pa.DataType | pa.Field | pa.Schema] = None,
|
||||
pip: tuple[str, ...] | list[str] = (),
|
||||
env: Optional[Mapping[str, str]] = None,
|
||||
secrets: tuple[str, ...] | list[str] = (),
|
||||
python_version: Optional[str] = None,
|
||||
conda: tuple[str, ...] | list[str] = (),
|
||||
conda_channels: tuple[str, ...] | list[str] = (),
|
||||
@@ -1002,6 +1076,7 @@ def udf(
|
||||
output_schema: Optional[pa.DataType | pa.Field | pa.Schema] = None,
|
||||
pip: tuple[str, ...] | list[str] = (),
|
||||
env: Optional[Mapping[str, str]] = None,
|
||||
secrets: tuple[str, ...] | list[str] = (),
|
||||
python_version: Optional[str] = None,
|
||||
conda: tuple[str, ...] | list[str] = (),
|
||||
conda_channels: tuple[str, ...] | list[str] = (),
|
||||
@@ -1032,7 +1107,10 @@ def udf(
|
||||
conda_channels : sequence of str, optional
|
||||
Conda channels in priority order; requires ``conda``.
|
||||
env : mapping of str to str, optional
|
||||
Environment variables included in the Function definition.
|
||||
Non-secret environment variables. Use ``secrets`` for credentials.
|
||||
secrets : sequence of str, optional
|
||||
Names of secrets required by the callable. Supply their values separately
|
||||
to ``create_function`` or ``create_function_async``.
|
||||
python_version : str, optional
|
||||
Remote Python major/minor version. Defaults to the client version.
|
||||
|
||||
@@ -1054,11 +1132,15 @@ def udf(
|
||||
Examples
|
||||
--------
|
||||
>>> from lancedb import udf
|
||||
>>> @udf(pip=["numpy==2.2.0"])
|
||||
>>> @udf(pip=["numpy==2.2.0"], secrets=["MODEL_TOKEN"])
|
||||
... def score(value: float) -> float:
|
||||
... return value * 2
|
||||
>>> score(1.5)
|
||||
3.0
|
||||
>>> db.create_function( # doctest: +SKIP
|
||||
... score, secrets={"MODEL_TOKEN": "user-secret-value"}
|
||||
... )
|
||||
|
||||
"""
|
||||
|
||||
def decorate(target: Callable[..., Any]) -> UdfDefinition:
|
||||
@@ -1069,6 +1151,7 @@ def udf(
|
||||
output_schema=output_schema,
|
||||
pip=tuple(pip),
|
||||
env={} if env is None else env,
|
||||
secrets=tuple(secrets),
|
||||
python_version=python_version,
|
||||
conda=tuple(conda),
|
||||
conda_channels=tuple(conda_channels),
|
||||
|
||||
@@ -7,7 +7,7 @@ import json
|
||||
import logging
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
import sys
|
||||
from typing import TYPE_CHECKING, Any, Dict, Iterable, List, Optional, Union
|
||||
from typing import TYPE_CHECKING, Any, Dict, Iterable, List, Mapping, Optional, Union
|
||||
from urllib.parse import urlparse
|
||||
import warnings
|
||||
|
||||
@@ -742,8 +742,15 @@ class RemoteDBConnection(DBConnection):
|
||||
return Job(self._conn.job(job_id))
|
||||
|
||||
@override
|
||||
def create_function_async(self, definition: UdfDefinition) -> Job[FunctionVersion]:
|
||||
return Job(LOOP.run(self._conn.create_function_async(definition)))
|
||||
def create_function_async(
|
||||
self,
|
||||
definition: UdfDefinition,
|
||||
*,
|
||||
secrets: Optional[Mapping[str, str]] = None,
|
||||
) -> Job[FunctionVersion]:
|
||||
return Job(
|
||||
LOOP.run(self._conn.create_function_async(definition, secrets=secrets))
|
||||
)
|
||||
|
||||
@override
|
||||
def get_function(self, name: str, *, version: str) -> FunctionVersion:
|
||||
|
||||
@@ -37,6 +37,21 @@ def job_result(name: str) -> dict:
|
||||
return json.loads(fixture(name))["result"]
|
||||
|
||||
|
||||
def assert_no_secret_values(value):
|
||||
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()
|
||||
@@ -94,6 +109,7 @@ def test_function_version_identity_is_immutable_and_exact():
|
||||
version = FunctionVersion.from_json(json.dumps(value))
|
||||
assert version.name == "embed"
|
||||
assert version.version == "fv_01K3EXACT"
|
||||
assert version.required_secrets == ("HF_TOKEN",)
|
||||
|
||||
with pytest.raises((TypeError, ValueError)):
|
||||
version.version = "fv_changed"
|
||||
@@ -276,6 +292,15 @@ def test_refresh_result_rejects_non_u64_values(field):
|
||||
RefreshColumnResult.from_json(json.dumps(value))
|
||||
|
||||
|
||||
def test_canonical_client_values_contain_secret_names_only():
|
||||
version = FunctionVersion.from_json(
|
||||
json.dumps(job_result("remote_function_job.json"))
|
||||
)
|
||||
canonical = json.loads(version.to_canonical_json())
|
||||
assert canonical["required_secrets"] == ["HF_TOKEN"]
|
||||
assert_no_secret_values(canonical)
|
||||
|
||||
|
||||
class _FunctionDeclarationInner:
|
||||
def __init__(self):
|
||||
self.calls = []
|
||||
|
||||
@@ -19,7 +19,7 @@ import pyarrow as pa
|
||||
import pytest
|
||||
|
||||
import lancedb
|
||||
from lancedb.functions import UdfDefinition, udf
|
||||
from lancedb.functions import FunctionRegistrationRequest, UdfDefinition, udf
|
||||
|
||||
THRESHOLD = 20
|
||||
_CACHE = None
|
||||
@@ -39,12 +39,28 @@ FIXTURES = (
|
||||
@udf(
|
||||
pip=["numpy>=2"],
|
||||
env={"MODE": "test"},
|
||||
secrets=["API_TOKEN"],
|
||||
python_version="3.12",
|
||||
)
|
||||
def normalize_score(value: float) -> float:
|
||||
return value / 100.0
|
||||
|
||||
|
||||
def _assert_no_secret_values(value):
|
||||
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_scalar_udf_matches_shared_registration_golden_and_remains_callable():
|
||||
assert isinstance(normalize_score, UdfDefinition)
|
||||
assert normalize_score(25.0) == 0.25
|
||||
@@ -59,6 +75,8 @@ def test_scalar_udf_matches_shared_registration_golden_and_remains_callable():
|
||||
"kind": "scalar_to_arrow_batch",
|
||||
"version": 1,
|
||||
}
|
||||
assert request["required_secrets"] == ["API_TOKEN"]
|
||||
_assert_no_secret_values(request)
|
||||
|
||||
|
||||
def _run_packaged(definition, *args):
|
||||
@@ -372,6 +390,7 @@ def test_udf_recursion_versus_a_rebound_module_name(tmp_path):
|
||||
output_schema=None,
|
||||
pip=(),
|
||||
env={},
|
||||
secrets=(),
|
||||
python_version=None,
|
||||
)
|
||||
with pytest.raises(ValueError, match="binds that name to another value"):
|
||||
@@ -526,13 +545,54 @@ def test_annotation_and_explicit_schema_validation_fail_closed():
|
||||
return value
|
||||
|
||||
|
||||
def test_secret_names_are_canonical_and_disjoint_from_environment():
|
||||
@udf(secrets=["Z_TOKEN", "A_TOKEN", "Z_TOKEN"])
|
||||
def canonical_secrets(value: int) -> int:
|
||||
return value
|
||||
|
||||
assert canonical_secrets.registration_request.required_secrets == (
|
||||
"A_TOKEN",
|
||||
"Z_TOKEN",
|
||||
)
|
||||
|
||||
with pytest.raises(ValueError, match="must be disjoint"):
|
||||
|
||||
@udf(env={"TOKEN": "plaintext"}, secrets=["TOKEN"])
|
||||
def overlapping(value: int) -> int:
|
||||
return value
|
||||
|
||||
|
||||
def test_declared_secret_api_still_requires_explicit_create_values():
|
||||
@udf(secrets=["API_TOKEN"])
|
||||
def declared_secret(value: int) -> int:
|
||||
return value
|
||||
|
||||
with pytest.raises(ValueError, match="missing"):
|
||||
declared_secret._submission_json(None)
|
||||
submission = json.loads(
|
||||
declared_secret._submission_json({"API_TOKEN": "explicit-secret"})
|
||||
)
|
||||
assert submission["required_secrets"] == ["API_TOKEN"]
|
||||
assert submission["secret_values"] == {"API_TOKEN": "explicit-secret"}
|
||||
|
||||
|
||||
def test_no_secrets_preserve_canonical_registration_shape():
|
||||
@udf
|
||||
def no_secrets(value: int) -> int:
|
||||
return value
|
||||
|
||||
canonical = json.loads(no_secrets.registration_request.to_canonical_json())
|
||||
assert "required_secrets" not in canonical
|
||||
assert json.loads(no_secrets._submission_json(None)) == canonical
|
||||
|
||||
|
||||
def test_local_function_catalog_operations_are_not_supported(tmp_path):
|
||||
db = lancedb.connect(tmp_path)
|
||||
message = "Function catalog operations are not supported by this database"
|
||||
with pytest.raises(NotImplementedError, match=message):
|
||||
db.create_function(normalize_score)
|
||||
db.create_function(normalize_score, secrets={"API_TOKEN": "value"})
|
||||
with pytest.raises(NotImplementedError, match=message):
|
||||
db.create_function_async(normalize_score)
|
||||
db.create_function_async(normalize_score, secrets={"API_TOKEN": "value"})
|
||||
with pytest.raises(NotImplementedError, match=message):
|
||||
db.get_function("normalize_score", version="fv_exact")
|
||||
|
||||
@@ -562,6 +622,7 @@ def _mock_remote_function_catalog():
|
||||
"runtime": body["runtime"],
|
||||
"runtime_digest": "sha256:runtime",
|
||||
"environment_digest": "sha256:environment",
|
||||
"required_secrets": body.get("required_secrets", []),
|
||||
"created_at": "2026-08-21T00:00:00Z",
|
||||
}
|
||||
response = {"job_id": "job-register"}
|
||||
@@ -608,7 +669,9 @@ def test_remote_registration_job_and_exact_version_reopen_round_trip():
|
||||
host_override=host,
|
||||
client_config={"retry_config": {"retries": 0}},
|
||||
)
|
||||
registration = db.create_function_async(normalize_score)
|
||||
registration = db.create_function_async(
|
||||
normalize_score, secrets={"API_TOKEN": "secret-value"}
|
||||
)
|
||||
assert registration.id == "job-register"
|
||||
created = registration.wait()
|
||||
reopened = db.get_function("normalize_score", version=created.version)
|
||||
@@ -617,9 +680,18 @@ def test_remote_registration_job_and_exact_version_reopen_round_trip():
|
||||
assert reopened.name == "normalize_score"
|
||||
assert reopened.version == "fv_exact"
|
||||
create_request = state["requests"][0][1]
|
||||
assert create_request == json.loads(
|
||||
expected = json.loads(normalize_score.registration_request.to_canonical_json())
|
||||
expected["secret_values"] = {"API_TOKEN": "secret-value"}
|
||||
assert create_request == expected
|
||||
durable_request = FunctionRegistrationRequest.from_json(json.dumps(create_request))
|
||||
assert not hasattr(durable_request, "secret_values")
|
||||
assert "secret_values" not in json.loads(durable_request.to_canonical_json())
|
||||
assert "secret_values" not in json.loads(
|
||||
normalize_score.registration_request.to_canonical_json()
|
||||
)
|
||||
assert "secret-value" not in repr(normalize_score)
|
||||
assert "secret-value" not in repr(normalize_score.registration_request)
|
||||
assert not hasattr(created, "secret_values")
|
||||
|
||||
|
||||
def test_blocking_remote_registration_returns_function_version():
|
||||
@@ -630,7 +702,9 @@ def test_blocking_remote_registration_returns_function_version():
|
||||
host_override=host,
|
||||
client_config={"retry_config": {"retries": 0}},
|
||||
)
|
||||
created = db.create_function(normalize_score)
|
||||
created = db.create_function(
|
||||
normalize_score, secrets={"API_TOKEN": "blocking-secret"}
|
||||
)
|
||||
|
||||
assert created.name == "normalize_score"
|
||||
assert created.version == "fv_exact"
|
||||
@@ -638,3 +712,49 @@ def test_blocking_remote_registration_returns_function_version():
|
||||
"/v1/functions/create",
|
||||
"/v1/jobs/describe",
|
||||
]
|
||||
assert state["requests"][0][1]["secret_values"] == {"API_TOKEN": "blocking-secret"}
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("secret_values", "error_type", "message"),
|
||||
[
|
||||
(None, ValueError, "missing"),
|
||||
({}, ValueError, "missing"),
|
||||
({"OTHER": "value"}, ValueError, "missing.*unexpected"),
|
||||
({"API_TOKEN": ""}, ValueError, "non-empty"),
|
||||
({"API_TOKEN": "bad\0value"}, ValueError, "NUL"),
|
||||
({"API_TOKEN": 123}, TypeError, "must be a string"),
|
||||
([("API_TOKEN", "value")], TypeError, "must be a mapping"),
|
||||
],
|
||||
)
|
||||
def test_secret_values_are_validated_before_remote_request(
|
||||
secret_values, error_type, message
|
||||
):
|
||||
with _mock_remote_function_catalog() as (host, state):
|
||||
db = lancedb.connect(
|
||||
"db://dev",
|
||||
api_key="fake",
|
||||
host_override=host,
|
||||
client_config={"retry_config": {"retries": 0}},
|
||||
)
|
||||
with pytest.raises(error_type, match=message):
|
||||
db.create_function_async(normalize_score, secrets=secret_values)
|
||||
assert state["requests"] == []
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_remote_registration_submits_secret_values_only_once():
|
||||
with _mock_remote_function_catalog() as (host, state):
|
||||
db = await lancedb.connect_async(
|
||||
"db://dev",
|
||||
api_key="fake",
|
||||
host_override=host,
|
||||
client_config={"retry_config": {"retries": 0}},
|
||||
)
|
||||
registration = await db.create_function_async(
|
||||
normalize_score, secrets={"API_TOKEN": "async-secret"}
|
||||
)
|
||||
created = await registration.wait()
|
||||
|
||||
assert state["requests"][0][1]["secret_values"] == {"API_TOKEN": "async-secret"}
|
||||
assert not hasattr(created, "secret_values")
|
||||
|
||||
Reference in New Issue
Block a user