mirror of
https://github.com/lancedb/lancedb.git
synced 2026-09-11 15:52:17 +00:00
Compare commits
6
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
ef752c2e3f | ||
|
|
0f6a355593 | ||
|
|
a6ec35502a | ||
|
|
a0bb1f7597 | ||
|
|
2562e117b2 | ||
|
|
134a265ee2 |
@@ -16,6 +16,7 @@ from typing import (
|
|||||||
Iterable,
|
Iterable,
|
||||||
List,
|
List,
|
||||||
Literal,
|
Literal,
|
||||||
|
Mapping,
|
||||||
Optional,
|
Optional,
|
||||||
Union,
|
Union,
|
||||||
)
|
)
|
||||||
@@ -687,17 +688,35 @@ class DBConnection(EnforceOverrides):
|
|||||||
"""
|
"""
|
||||||
raise NotImplementedError("serialize is not supported for this connection type")
|
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.
|
"""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`.
|
This is the blocking counterpart of :meth:`create_function_async`.
|
||||||
Local connections raise ``NotImplementedError``.
|
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.
|
"""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
|
Submission returns a typed job. The immutable Function version becomes
|
||||||
available only when :meth:`Job.wait` succeeds. Local connections raise
|
available only when :meth:`Job.wait` succeeds. Local connections raise
|
||||||
``NotImplementedError``.
|
``NotImplementedError``.
|
||||||
@@ -1405,8 +1424,13 @@ class LanceDBConnection(DBConnection):
|
|||||||
return Job(self._conn.job(job_id))
|
return Job(self._conn.job(job_id))
|
||||||
|
|
||||||
@override
|
@override
|
||||||
def create_function_async(self, definition: UdfDefinition) -> Job[FunctionVersion]:
|
def create_function_async(
|
||||||
job = LOOP.run(self._conn.create_function_async(definition))
|
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)
|
return Job(job)
|
||||||
|
|
||||||
@override
|
@override
|
||||||
@@ -2225,17 +2249,24 @@ class AsyncConnection(object):
|
|||||||
return AsyncJob(self._inner.job(job_id))
|
return AsyncJob(self._inner.job(job_id))
|
||||||
|
|
||||||
async def create_function_async(
|
async def create_function_async(
|
||||||
self, definition: UdfDefinition
|
self,
|
||||||
|
definition: UdfDefinition,
|
||||||
|
*,
|
||||||
|
secrets: Optional[Mapping[str, str]] = None,
|
||||||
) -> AsyncJob[FunctionVersion]:
|
) -> AsyncJob[FunctionVersion]:
|
||||||
"""Register a scalar Python UDF through the remote Function catalog.
|
"""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.
|
The returned typed job resolves to the immutable Function version.
|
||||||
Local connections raise ``NotImplementedError``.
|
Local connections raise ``NotImplementedError``.
|
||||||
"""
|
"""
|
||||||
if not isinstance(definition, UdfDefinition):
|
if not isinstance(definition, UdfDefinition):
|
||||||
raise TypeError("create_function_async requires a @udf definition")
|
raise TypeError("create_function_async requires a @udf definition")
|
||||||
inner = await self._inner.create_function_async(
|
inner = await self._inner.create_function_async(
|
||||||
definition.registration_request.to_canonical_json()
|
definition._submission_json(secrets)
|
||||||
)
|
)
|
||||||
return _typed_job(inner, FunctionVersion.from_json)
|
return _typed_job(inner, FunctionVersion.from_json)
|
||||||
|
|
||||||
|
|||||||
@@ -4,7 +4,7 @@
|
|||||||
"""Canonical Function values exchanged with LanceDB Enterprise services.
|
"""Canonical Function values exchanged with LanceDB Enterprise services.
|
||||||
|
|
||||||
These immutable models contain client/wire state only. Catalog persistence,
|
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
|
``RefreshColumnResult`` is also the backend-neutral result of a local
|
||||||
expression-backed refresh job.
|
expression-backed refresh job.
|
||||||
"""
|
"""
|
||||||
@@ -229,7 +229,7 @@ class PythonEnvironmentSpec(_RemoteValue):
|
|||||||
|
|
||||||
|
|
||||||
class PythonRuntimeSpec(_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
|
V1 supports ``kind="python"``. Newer runtime kinds remain readable, while
|
||||||
their unknown payload fields are intentionally not retained by the client.
|
their unknown payload fields are intentionally not retained by the client.
|
||||||
@@ -268,6 +268,7 @@ class FunctionVersion(_RemoteValue):
|
|||||||
runtime: PythonRuntimeSpec
|
runtime: PythonRuntimeSpec
|
||||||
runtime_digest: str
|
runtime_digest: str
|
||||||
environment_digest: str
|
environment_digest: str
|
||||||
|
required_secrets: tuple[str, ...] = ()
|
||||||
created_at: str
|
created_at: str
|
||||||
|
|
||||||
def __call__(self, **inputs: Any) -> FunctionApplication:
|
def __call__(self, **inputs: Any) -> FunctionApplication:
|
||||||
@@ -329,12 +330,17 @@ class FunctionVersion(_RemoteValue):
|
|||||||
|
|
||||||
|
|
||||||
class FunctionRegistrationRequest(_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
|
name: str
|
||||||
artifact: FunctionArtifactRequest
|
artifact: FunctionArtifactRequest
|
||||||
signature: FunctionSignature
|
signature: FunctionSignature
|
||||||
runtime: PythonRuntimeSpec
|
runtime: PythonRuntimeSpec
|
||||||
|
required_secrets: tuple[str, ...] = ()
|
||||||
|
|
||||||
|
|
||||||
class FunctionVersionRef(_OpenRemoteValue):
|
class FunctionVersionRef(_OpenRemoteValue):
|
||||||
@@ -479,6 +485,27 @@ class RefreshColumnResult(_RemoteValue):
|
|||||||
|
|
||||||
|
|
||||||
_FUNCTION_NAME = re.compile(r"^[A-Za-z_][A-Za-z0-9_.-]*$")
|
_FUNCTION_NAME = re.compile(r"^[A-Za-z_][A-Za-z0-9_.-]*$")
|
||||||
|
_SECRET_NAME = re.compile(r"^[A-Za-z_][A-Za-z0-9_]*$")
|
||||||
|
# Keep this byte limit aligned with Sophon's MAX_FUNCTION_SECRET_VALUE_BYTES.
|
||||||
|
_MAX_FUNCTION_SECRET_VALUE_BYTES = 64 * 1024
|
||||||
|
_MAX_FUNCTION_SECRET_VALUES_BYTES = 512 * 1024
|
||||||
|
|
||||||
|
|
||||||
|
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")
|
||||||
|
value_bytes = len(value.encode("utf-8"))
|
||||||
|
if value_bytes > _MAX_FUNCTION_SECRET_VALUE_BYTES:
|
||||||
|
raise ValueError(
|
||||||
|
f"Function secret {name!r} value exceeds the "
|
||||||
|
f"{_MAX_FUNCTION_SECRET_VALUE_BYTES}-byte limit"
|
||||||
|
)
|
||||||
|
return value
|
||||||
|
|
||||||
|
|
||||||
_GRAMMAR_PRIMITIVES = (
|
_GRAMMAR_PRIMITIVES = (
|
||||||
@@ -909,6 +936,7 @@ class UdfDefinition:
|
|||||||
output_schema: Optional[pa.DataType | pa.Field | pa.Schema],
|
output_schema: Optional[pa.DataType | pa.Field | pa.Schema],
|
||||||
pip: tuple[str, ...],
|
pip: tuple[str, ...],
|
||||||
env: Mapping[str, str],
|
env: Mapping[str, str],
|
||||||
|
secrets: tuple[str, ...],
|
||||||
python_version: Optional[str],
|
python_version: Optional[str],
|
||||||
conda: tuple[str, ...] = (),
|
conda: tuple[str, ...] = (),
|
||||||
conda_channels: tuple[str, ...] = (),
|
conda_channels: tuple[str, ...] = (),
|
||||||
@@ -935,6 +963,17 @@ class UdfDefinition:
|
|||||||
for key, value in environment.items()
|
for key, value in environment.items()
|
||||||
):
|
):
|
||||||
raise TypeError("Function env keys and values must be strings")
|
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)
|
signature = _infer_signature(function, input_schema, output_schema)
|
||||||
source = _package_source(function)
|
source = _package_source(function)
|
||||||
digest = f"sha256:{hashlib.sha256(source).hexdigest()}"
|
digest = f"sha256:{hashlib.sha256(source).hexdigest()}"
|
||||||
@@ -963,14 +1002,65 @@ class UdfDefinition:
|
|||||||
),
|
),
|
||||||
signature=signature,
|
signature=signature,
|
||||||
runtime=runtime,
|
runtime=runtime,
|
||||||
|
required_secrets=required_secrets,
|
||||||
)
|
)
|
||||||
functools.update_wrapper(self, function)
|
functools.update_wrapper(self, function)
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def registration_request(self) -> FunctionRegistrationRequest:
|
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
|
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 = {}
|
||||||
|
total_bytes = 0
|
||||||
|
for name in sorted(secret_values):
|
||||||
|
value = _validate_secret_value(name, secret_values[name])
|
||||||
|
total_bytes += len(value.encode("utf-8"))
|
||||||
|
if total_bytes > _MAX_FUNCTION_SECRET_VALUES_BYTES:
|
||||||
|
raise ValueError(
|
||||||
|
"Function secret values exceed the "
|
||||||
|
f"{_MAX_FUNCTION_SECRET_VALUES_BYTES}-byte request limit"
|
||||||
|
)
|
||||||
|
canonical_values[name] = value
|
||||||
|
|
||||||
|
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):
|
def __call__(self, *args, **kwargs):
|
||||||
return self._function(*args, **kwargs)
|
return self._function(*args, **kwargs)
|
||||||
|
|
||||||
@@ -988,6 +1078,7 @@ def udf(
|
|||||||
output_schema: Optional[pa.DataType | pa.Field | pa.Schema] = None,
|
output_schema: Optional[pa.DataType | pa.Field | pa.Schema] = None,
|
||||||
pip: tuple[str, ...] | list[str] = (),
|
pip: tuple[str, ...] | list[str] = (),
|
||||||
env: Optional[Mapping[str, str]] = None,
|
env: Optional[Mapping[str, str]] = None,
|
||||||
|
secrets: tuple[str, ...] | list[str] = (),
|
||||||
python_version: Optional[str] = None,
|
python_version: Optional[str] = None,
|
||||||
conda: tuple[str, ...] | list[str] = (),
|
conda: tuple[str, ...] | list[str] = (),
|
||||||
conda_channels: tuple[str, ...] | list[str] = (),
|
conda_channels: tuple[str, ...] | list[str] = (),
|
||||||
@@ -1002,6 +1093,7 @@ def udf(
|
|||||||
output_schema: Optional[pa.DataType | pa.Field | pa.Schema] = None,
|
output_schema: Optional[pa.DataType | pa.Field | pa.Schema] = None,
|
||||||
pip: tuple[str, ...] | list[str] = (),
|
pip: tuple[str, ...] | list[str] = (),
|
||||||
env: Optional[Mapping[str, str]] = None,
|
env: Optional[Mapping[str, str]] = None,
|
||||||
|
secrets: tuple[str, ...] | list[str] = (),
|
||||||
python_version: Optional[str] = None,
|
python_version: Optional[str] = None,
|
||||||
conda: tuple[str, ...] | list[str] = (),
|
conda: tuple[str, ...] | list[str] = (),
|
||||||
conda_channels: tuple[str, ...] | list[str] = (),
|
conda_channels: tuple[str, ...] | list[str] = (),
|
||||||
@@ -1032,7 +1124,10 @@ def udf(
|
|||||||
conda_channels : sequence of str, optional
|
conda_channels : sequence of str, optional
|
||||||
Conda channels in priority order; requires ``conda``.
|
Conda channels in priority order; requires ``conda``.
|
||||||
env : mapping of str to str, optional
|
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
|
python_version : str, optional
|
||||||
Remote Python major/minor version. Defaults to the client version.
|
Remote Python major/minor version. Defaults to the client version.
|
||||||
|
|
||||||
@@ -1054,11 +1149,15 @@ def udf(
|
|||||||
Examples
|
Examples
|
||||||
--------
|
--------
|
||||||
>>> from lancedb import udf
|
>>> from lancedb import udf
|
||||||
>>> @udf(pip=["numpy==2.2.0"])
|
>>> @udf(pip=["numpy==2.2.0"], secrets=["MODEL_TOKEN"])
|
||||||
... def score(value: float) -> float:
|
... def score(value: float) -> float:
|
||||||
... return value * 2
|
... return value * 2
|
||||||
>>> score(1.5)
|
>>> score(1.5)
|
||||||
3.0
|
3.0
|
||||||
|
>>> db.create_function( # doctest: +SKIP
|
||||||
|
... score, secrets={"MODEL_TOKEN": "user-secret-value"}
|
||||||
|
... )
|
||||||
|
|
||||||
"""
|
"""
|
||||||
|
|
||||||
def decorate(target: Callable[..., Any]) -> UdfDefinition:
|
def decorate(target: Callable[..., Any]) -> UdfDefinition:
|
||||||
@@ -1069,6 +1168,7 @@ def udf(
|
|||||||
output_schema=output_schema,
|
output_schema=output_schema,
|
||||||
pip=tuple(pip),
|
pip=tuple(pip),
|
||||||
env={} if env is None else env,
|
env={} if env is None else env,
|
||||||
|
secrets=tuple(secrets),
|
||||||
python_version=python_version,
|
python_version=python_version,
|
||||||
conda=tuple(conda),
|
conda=tuple(conda),
|
||||||
conda_channels=tuple(conda_channels),
|
conda_channels=tuple(conda_channels),
|
||||||
|
|||||||
@@ -7,7 +7,7 @@ import json
|
|||||||
import logging
|
import logging
|
||||||
from concurrent.futures import ThreadPoolExecutor
|
from concurrent.futures import ThreadPoolExecutor
|
||||||
import sys
|
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
|
from urllib.parse import urlparse
|
||||||
import warnings
|
import warnings
|
||||||
|
|
||||||
@@ -742,8 +742,15 @@ class RemoteDBConnection(DBConnection):
|
|||||||
return Job(self._conn.job(job_id))
|
return Job(self._conn.job(job_id))
|
||||||
|
|
||||||
@override
|
@override
|
||||||
def create_function_async(self, definition: UdfDefinition) -> Job[FunctionVersion]:
|
def create_function_async(
|
||||||
return Job(LOOP.run(self._conn.create_function_async(definition)))
|
self,
|
||||||
|
definition: UdfDefinition,
|
||||||
|
*,
|
||||||
|
secrets: Optional[Mapping[str, str]] = None,
|
||||||
|
) -> Job[FunctionVersion]:
|
||||||
|
return Job(
|
||||||
|
LOOP.run(self._conn.create_function_async(definition, secrets=secrets))
|
||||||
|
)
|
||||||
|
|
||||||
@override
|
@override
|
||||||
def get_function(self, name: str, *, version: str) -> FunctionVersion:
|
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"]
|
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():
|
def test_public_function_values_are_in_api_reference():
|
||||||
docs = Path(__file__).parents[3] / "docs" / "src" / "python" / "python.md"
|
docs = Path(__file__).parents[3] / "docs" / "src" / "python" / "python.md"
|
||||||
rendered = docs.read_text()
|
rendered = docs.read_text()
|
||||||
@@ -94,6 +109,7 @@ def test_function_version_identity_is_immutable_and_exact():
|
|||||||
version = FunctionVersion.from_json(json.dumps(value))
|
version = FunctionVersion.from_json(json.dumps(value))
|
||||||
assert version.name == "embed"
|
assert version.name == "embed"
|
||||||
assert version.version == "fv_01K3EXACT"
|
assert version.version == "fv_01K3EXACT"
|
||||||
|
assert version.required_secrets == ("HF_TOKEN",)
|
||||||
|
|
||||||
with pytest.raises((TypeError, ValueError)):
|
with pytest.raises((TypeError, ValueError)):
|
||||||
version.version = "fv_changed"
|
version.version = "fv_changed"
|
||||||
@@ -276,6 +292,15 @@ def test_refresh_result_rejects_non_u64_values(field):
|
|||||||
RefreshColumnResult.from_json(json.dumps(value))
|
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:
|
class _FunctionDeclarationInner:
|
||||||
def __init__(self):
|
def __init__(self):
|
||||||
self.calls = []
|
self.calls = []
|
||||||
|
|||||||
@@ -19,7 +19,13 @@ import pyarrow as pa
|
|||||||
import pytest
|
import pytest
|
||||||
|
|
||||||
import lancedb
|
import lancedb
|
||||||
from lancedb.functions import UdfDefinition, udf
|
from lancedb.functions import (
|
||||||
|
_MAX_FUNCTION_SECRET_VALUE_BYTES,
|
||||||
|
_MAX_FUNCTION_SECRET_VALUES_BYTES,
|
||||||
|
FunctionRegistrationRequest,
|
||||||
|
UdfDefinition,
|
||||||
|
udf,
|
||||||
|
)
|
||||||
|
|
||||||
THRESHOLD = 20
|
THRESHOLD = 20
|
||||||
_CACHE = None
|
_CACHE = None
|
||||||
@@ -39,12 +45,28 @@ FIXTURES = (
|
|||||||
@udf(
|
@udf(
|
||||||
pip=["numpy>=2"],
|
pip=["numpy>=2"],
|
||||||
env={"MODE": "test"},
|
env={"MODE": "test"},
|
||||||
|
secrets=["API_TOKEN"],
|
||||||
python_version="3.12",
|
python_version="3.12",
|
||||||
)
|
)
|
||||||
def normalize_score(value: float) -> float:
|
def normalize_score(value: float) -> float:
|
||||||
return value / 100.0
|
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():
|
def test_scalar_udf_matches_shared_registration_golden_and_remains_callable():
|
||||||
assert isinstance(normalize_score, UdfDefinition)
|
assert isinstance(normalize_score, UdfDefinition)
|
||||||
assert normalize_score(25.0) == 0.25
|
assert normalize_score(25.0) == 0.25
|
||||||
@@ -59,6 +81,8 @@ def test_scalar_udf_matches_shared_registration_golden_and_remains_callable():
|
|||||||
"kind": "scalar_to_arrow_batch",
|
"kind": "scalar_to_arrow_batch",
|
||||||
"version": 1,
|
"version": 1,
|
||||||
}
|
}
|
||||||
|
assert request["required_secrets"] == ["API_TOKEN"]
|
||||||
|
_assert_no_secret_values(request)
|
||||||
|
|
||||||
|
|
||||||
def _run_packaged(definition, *args):
|
def _run_packaged(definition, *args):
|
||||||
@@ -372,6 +396,7 @@ def test_udf_recursion_versus_a_rebound_module_name(tmp_path):
|
|||||||
output_schema=None,
|
output_schema=None,
|
||||||
pip=(),
|
pip=(),
|
||||||
env={},
|
env={},
|
||||||
|
secrets=(),
|
||||||
python_version=None,
|
python_version=None,
|
||||||
)
|
)
|
||||||
with pytest.raises(ValueError, match="binds that name to another value"):
|
with pytest.raises(ValueError, match="binds that name to another value"):
|
||||||
@@ -526,13 +551,54 @@ def test_annotation_and_explicit_schema_validation_fail_closed():
|
|||||||
return value
|
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):
|
def test_local_function_catalog_operations_are_not_supported(tmp_path):
|
||||||
db = lancedb.connect(tmp_path)
|
db = lancedb.connect(tmp_path)
|
||||||
message = "Function catalog operations are not supported by this database"
|
message = "Function catalog operations are not supported by this database"
|
||||||
with pytest.raises(NotImplementedError, match=message):
|
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):
|
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):
|
with pytest.raises(NotImplementedError, match=message):
|
||||||
db.get_function("normalize_score", version="fv_exact")
|
db.get_function("normalize_score", version="fv_exact")
|
||||||
|
|
||||||
@@ -562,6 +628,7 @@ def _mock_remote_function_catalog():
|
|||||||
"runtime": body["runtime"],
|
"runtime": body["runtime"],
|
||||||
"runtime_digest": "sha256:runtime",
|
"runtime_digest": "sha256:runtime",
|
||||||
"environment_digest": "sha256:environment",
|
"environment_digest": "sha256:environment",
|
||||||
|
"required_secrets": body.get("required_secrets", []),
|
||||||
"created_at": "2026-08-21T00:00:00Z",
|
"created_at": "2026-08-21T00:00:00Z",
|
||||||
}
|
}
|
||||||
response = {"job_id": "job-register"}
|
response = {"job_id": "job-register"}
|
||||||
@@ -608,7 +675,9 @@ def test_remote_registration_job_and_exact_version_reopen_round_trip():
|
|||||||
host_override=host,
|
host_override=host,
|
||||||
client_config={"retry_config": {"retries": 0}},
|
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"
|
assert registration.id == "job-register"
|
||||||
created = registration.wait()
|
created = registration.wait()
|
||||||
reopened = db.get_function("normalize_score", version=created.version)
|
reopened = db.get_function("normalize_score", version=created.version)
|
||||||
@@ -617,9 +686,18 @@ def test_remote_registration_job_and_exact_version_reopen_round_trip():
|
|||||||
assert reopened.name == "normalize_score"
|
assert reopened.name == "normalize_score"
|
||||||
assert reopened.version == "fv_exact"
|
assert reopened.version == "fv_exact"
|
||||||
create_request = state["requests"][0][1]
|
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()
|
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():
|
def test_blocking_remote_registration_returns_function_version():
|
||||||
@@ -630,7 +708,9 @@ def test_blocking_remote_registration_returns_function_version():
|
|||||||
host_override=host,
|
host_override=host,
|
||||||
client_config={"retry_config": {"retries": 0}},
|
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.name == "normalize_score"
|
||||||
assert created.version == "fv_exact"
|
assert created.version == "fv_exact"
|
||||||
@@ -638,3 +718,121 @@ def test_blocking_remote_registration_returns_function_version():
|
|||||||
"/v1/functions/create",
|
"/v1/functions/create",
|
||||||
"/v1/jobs/describe",
|
"/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.parametrize(
|
||||||
|
"value",
|
||||||
|
[
|
||||||
|
"x" * _MAX_FUNCTION_SECRET_VALUE_BYTES,
|
||||||
|
"é" * (_MAX_FUNCTION_SECRET_VALUE_BYTES // len("é".encode("utf-8"))),
|
||||||
|
],
|
||||||
|
ids=["ascii", "multibyte"],
|
||||||
|
)
|
||||||
|
def test_secret_value_accepts_exact_utf8_byte_limit(value):
|
||||||
|
submission = json.loads(normalize_score._submission_json({"API_TOKEN": value}))
|
||||||
|
assert submission["secret_values"]["API_TOKEN"] == value
|
||||||
|
assert len(value.encode("utf-8")) == _MAX_FUNCTION_SECRET_VALUE_BYTES
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize(
|
||||||
|
"value",
|
||||||
|
[
|
||||||
|
"x" * (_MAX_FUNCTION_SECRET_VALUE_BYTES + 1),
|
||||||
|
"é" * (_MAX_FUNCTION_SECRET_VALUE_BYTES // len("é".encode("utf-8")) + 1),
|
||||||
|
],
|
||||||
|
ids=["ascii", "multibyte"],
|
||||||
|
)
|
||||||
|
def test_secret_value_rejects_over_utf8_byte_limit_before_json_construction(
|
||||||
|
monkeypatch, value
|
||||||
|
):
|
||||||
|
def fail_if_json_construction_starts(self):
|
||||||
|
pytest.fail("oversized secret reached JSON construction")
|
||||||
|
|
||||||
|
monkeypatch.setattr(
|
||||||
|
FunctionRegistrationRequest, "_known_dict", fail_if_json_construction_starts
|
||||||
|
)
|
||||||
|
with pytest.raises(ValueError, match=r"exceeds the 65536-byte limit"):
|
||||||
|
normalize_score._submission_json({"API_TOKEN": value})
|
||||||
|
|
||||||
|
|
||||||
|
def test_secret_values_accept_exact_aggregate_utf8_byte_limit(monkeypatch):
|
||||||
|
names = tuple(f"SECRET_{index}" for index in range(8))
|
||||||
|
value = "é" * (_MAX_FUNCTION_SECRET_VALUE_BYTES // len("é".encode("utf-8")))
|
||||||
|
values = {name: value for name in names}
|
||||||
|
monkeypatch.setattr(
|
||||||
|
normalize_score,
|
||||||
|
"_request",
|
||||||
|
normalize_score._request._copy(update={"required_secrets": names}),
|
||||||
|
)
|
||||||
|
|
||||||
|
submission = json.loads(normalize_score._submission_json(values))
|
||||||
|
|
||||||
|
assert submission["secret_values"] == values
|
||||||
|
assert sum(len(item.encode("utf-8")) for item in values.values()) == (
|
||||||
|
_MAX_FUNCTION_SECRET_VALUES_BYTES
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def test_secret_values_reject_aggregate_over_limit_before_construction(monkeypatch):
|
||||||
|
names = tuple(f"SECRET_{index}" for index in range(9))
|
||||||
|
values = {name: "x" * _MAX_FUNCTION_SECRET_VALUE_BYTES for name in names}
|
||||||
|
monkeypatch.setattr(
|
||||||
|
normalize_score,
|
||||||
|
"_request",
|
||||||
|
normalize_score._request._copy(update={"required_secrets": names}),
|
||||||
|
)
|
||||||
|
|
||||||
|
def fail_if_json_construction_starts(self):
|
||||||
|
pytest.fail("oversized aggregate reached JSON construction")
|
||||||
|
|
||||||
|
monkeypatch.setattr(
|
||||||
|
FunctionRegistrationRequest, "_known_dict", fail_if_json_construction_starts
|
||||||
|
)
|
||||||
|
with pytest.raises(ValueError, match=r"exceed.*524288-byte request limit"):
|
||||||
|
normalize_score._submission_json(values)
|
||||||
|
|
||||||
|
|
||||||
|
@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")
|
||||||
|
|||||||
@@ -5,9 +5,9 @@
|
|||||||
//! backend-neutral terminal result of a computed-column refresh.
|
//! backend-neutral terminal result of a computed-column refresh.
|
||||||
//!
|
//!
|
||||||
//! This module contains client/wire values only. Catalog persistence,
|
//! This module contains client/wire values only. Catalog persistence,
|
||||||
//! environment bake, and execution are owned by Sophon.
|
//! environment bake, secret resolution, and execution are owned by Sophon.
|
||||||
|
|
||||||
use std::collections::BTreeMap;
|
use std::collections::{BTreeMap, BTreeSet};
|
||||||
|
|
||||||
use serde::de::{self, DeserializeOwned};
|
use serde::de::{self, DeserializeOwned};
|
||||||
use serde::{Deserialize, Deserializer, Serialize, Serializer};
|
use serde::{Deserialize, Deserializer, Serialize, Serializer};
|
||||||
@@ -15,6 +15,16 @@ use serde_json::Value;
|
|||||||
|
|
||||||
use crate::{Error, Result};
|
use crate::{Error, Result};
|
||||||
|
|
||||||
|
// Keep these byte limits aligned with Sophon's Function submission validation.
|
||||||
|
pub(crate) const MAX_FUNCTION_SECRET_VALUE_BYTES: usize = 64 * 1024;
|
||||||
|
const MAX_FUNCTION_SECRET_VALUES_BYTES: usize = 512 * 1024;
|
||||||
|
|
||||||
|
fn is_portable_environment_name(name: &str) -> bool {
|
||||||
|
let mut bytes = name.bytes();
|
||||||
|
matches!(bytes.next(), Some(b'A'..=b'Z' | b'a'..=b'z' | b'_'))
|
||||||
|
&& bytes.all(|byte| matches!(byte, b'A'..=b'Z' | b'a'..=b'z' | b'0'..=b'9' | b'_'))
|
||||||
|
}
|
||||||
|
|
||||||
fn invalid_json(error: impl std::fmt::Display) -> Error {
|
fn invalid_json(error: impl std::fmt::Display) -> Error {
|
||||||
Error::InvalidInput {
|
Error::InvalidInput {
|
||||||
message: format!("invalid remote Function JSON: {error}"),
|
message: format!("invalid remote Function JSON: {error}"),
|
||||||
@@ -198,6 +208,11 @@ pub struct PythonEnvironmentSpec {
|
|||||||
}
|
}
|
||||||
|
|
||||||
/// Reproducible Python runtime definition understood by Sophon.
|
/// Reproducible Python runtime definition understood by Sophon.
|
||||||
|
///
|
||||||
|
/// `env` contains non-secret values. Secret values are submission-only in the
|
||||||
|
/// client model and do not become part of this public runtime identity;
|
||||||
|
/// [`FunctionVersion::required_secrets`] contains names only. Sophon persists
|
||||||
|
/// submitted values separately in the private execution artifact.
|
||||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||||
#[non_exhaustive]
|
#[non_exhaustive]
|
||||||
pub enum PythonRuntimeSpec {
|
pub enum PythonRuntimeSpec {
|
||||||
@@ -239,7 +254,7 @@ impl PythonRuntimeSpec {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Environment variables, or `None` for an unknown kind.
|
/// Non-secret environment variables, or `None` for an unknown kind.
|
||||||
pub fn env(&self) -> Option<&BTreeMap<String, String>> {
|
pub fn env(&self) -> Option<&BTreeMap<String, String>> {
|
||||||
match self {
|
match self {
|
||||||
Self::Python { env, .. } => Some(env),
|
Self::Python { env, .. } => Some(env),
|
||||||
@@ -324,6 +339,8 @@ pub struct FunctionVersion {
|
|||||||
runtime: PythonRuntimeSpec,
|
runtime: PythonRuntimeSpec,
|
||||||
runtime_digest: String,
|
runtime_digest: String,
|
||||||
environment_digest: String,
|
environment_digest: String,
|
||||||
|
#[serde(default, skip_serializing_if = "Vec::is_empty")]
|
||||||
|
required_secrets: Vec<String>,
|
||||||
created_at: String,
|
created_at: String,
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -356,6 +373,12 @@ impl FunctionVersion {
|
|||||||
&self.environment_digest
|
&self.environment_digest
|
||||||
}
|
}
|
||||||
|
|
||||||
|
/// Required secret names. Resolved values exist only in Sophon's private
|
||||||
|
/// execution artifact and worker launch path.
|
||||||
|
pub fn required_secrets(&self) -> &[String] {
|
||||||
|
&self.required_secrets
|
||||||
|
}
|
||||||
|
|
||||||
pub fn created_at(&self) -> &str {
|
pub fn created_at(&self) -> &str {
|
||||||
&self.created_at
|
&self.created_at
|
||||||
}
|
}
|
||||||
@@ -397,12 +420,115 @@ pub struct FunctionArtifactRequest {
|
|||||||
}
|
}
|
||||||
|
|
||||||
/// Stable request envelope for remote immutable Function registration.
|
/// Stable request envelope for remote immutable Function registration.
|
||||||
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
|
///
|
||||||
|
/// Secret values are submission-only in the client model. Sophon persists them
|
||||||
|
/// in the database-scoped private execution artifact; returned
|
||||||
|
/// [`FunctionVersion`] and Job metadata contain only
|
||||||
|
/// [`Self::required_secrets`] names. Debug formatting always redacts values.
|
||||||
|
#[derive(Clone, PartialEq, Eq, Serialize, Deserialize)]
|
||||||
pub struct FunctionRegistrationRequest {
|
pub struct FunctionRegistrationRequest {
|
||||||
pub name: String,
|
pub name: String,
|
||||||
pub artifact: FunctionArtifactRequest,
|
pub artifact: FunctionArtifactRequest,
|
||||||
pub signature: FunctionSignature,
|
pub signature: FunctionSignature,
|
||||||
pub runtime: PythonRuntimeSpec,
|
pub runtime: PythonRuntimeSpec,
|
||||||
|
#[serde(default, skip_serializing_if = "Vec::is_empty")]
|
||||||
|
pub required_secrets: Vec<String>,
|
||||||
|
#[serde(default, skip_serializing_if = "BTreeMap::is_empty")]
|
||||||
|
pub secret_values: BTreeMap<String, String>,
|
||||||
|
}
|
||||||
|
|
||||||
|
impl FunctionRegistrationRequest {
|
||||||
|
pub(crate) fn validate_secret_values(&self) -> Result<()> {
|
||||||
|
let mut required = BTreeSet::new();
|
||||||
|
for name in &self.required_secrets {
|
||||||
|
if !is_portable_environment_name(name) {
|
||||||
|
return Err(Error::InvalidInput {
|
||||||
|
message: format!(
|
||||||
|
"Function secret name {name:?} must be a portable environment variable name"
|
||||||
|
),
|
||||||
|
});
|
||||||
|
}
|
||||||
|
if !required.insert(name) {
|
||||||
|
return Err(Error::InvalidInput {
|
||||||
|
message: format!("Function required_secrets contains duplicate name {name:?}"),
|
||||||
|
});
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if let PythonRuntimeSpec::Python { env, .. } = &self.runtime
|
||||||
|
&& let Some(name) = required.iter().find(|name| env.contains_key(**name))
|
||||||
|
{
|
||||||
|
return Err(Error::InvalidInput {
|
||||||
|
message: format!(
|
||||||
|
"Function runtime env and secret names must be disjoint: {name:?}"
|
||||||
|
),
|
||||||
|
});
|
||||||
|
}
|
||||||
|
|
||||||
|
let provided = self.secret_values.keys().collect::<BTreeSet<_>>();
|
||||||
|
if required != provided {
|
||||||
|
return Err(Error::InvalidInput {
|
||||||
|
message: "Function secret_values keys must exactly match required_secrets"
|
||||||
|
.to_string(),
|
||||||
|
});
|
||||||
|
}
|
||||||
|
|
||||||
|
let mut total_bytes = 0usize;
|
||||||
|
for (name, value) in &self.secret_values {
|
||||||
|
if value.is_empty() {
|
||||||
|
return Err(Error::InvalidInput {
|
||||||
|
message: format!("Function secret {name:?} value must be non-empty"),
|
||||||
|
});
|
||||||
|
}
|
||||||
|
if value.contains('\0') {
|
||||||
|
return Err(Error::InvalidInput {
|
||||||
|
message: format!("Function secret {name:?} value must not contain NUL"),
|
||||||
|
});
|
||||||
|
}
|
||||||
|
if value.len() > MAX_FUNCTION_SECRET_VALUE_BYTES {
|
||||||
|
return Err(Error::InvalidInput {
|
||||||
|
message: format!(
|
||||||
|
"Function secret {name:?} value exceeds the \
|
||||||
|
{MAX_FUNCTION_SECRET_VALUE_BYTES}-byte limit"
|
||||||
|
),
|
||||||
|
});
|
||||||
|
}
|
||||||
|
total_bytes =
|
||||||
|
total_bytes
|
||||||
|
.checked_add(value.len())
|
||||||
|
.ok_or_else(|| Error::InvalidInput {
|
||||||
|
message: "Function secret values exceed the request byte limit".to_string(),
|
||||||
|
})?;
|
||||||
|
}
|
||||||
|
if total_bytes > MAX_FUNCTION_SECRET_VALUES_BYTES {
|
||||||
|
return Err(Error::InvalidInput {
|
||||||
|
message: format!(
|
||||||
|
"Function secret values exceed the \
|
||||||
|
{MAX_FUNCTION_SECRET_VALUES_BYTES}-byte request limit"
|
||||||
|
),
|
||||||
|
});
|
||||||
|
}
|
||||||
|
Ok(())
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
impl std::fmt::Debug for FunctionRegistrationRequest {
|
||||||
|
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||||
|
let secret_values = self
|
||||||
|
.secret_values
|
||||||
|
.keys()
|
||||||
|
.map(|name| (name, "[REDACTED]"))
|
||||||
|
.collect::<BTreeMap<_, _>>();
|
||||||
|
formatter
|
||||||
|
.debug_struct("FunctionRegistrationRequest")
|
||||||
|
.field("name", &self.name)
|
||||||
|
.field("artifact", &self.artifact)
|
||||||
|
.field("signature", &self.signature)
|
||||||
|
.field("runtime", &self.runtime)
|
||||||
|
.field("required_secrets", &self.required_secrets)
|
||||||
|
.field("secret_values", &secret_values)
|
||||||
|
.finish()
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
impl_json!(FunctionRegistrationRequest);
|
impl_json!(FunctionRegistrationRequest);
|
||||||
@@ -587,6 +713,185 @@ impl RefreshColumnResult {
|
|||||||
|
|
||||||
impl_json!(RefreshColumnResult);
|
impl_json!(RefreshColumnResult);
|
||||||
|
|
||||||
|
#[cfg(test)]
|
||||||
|
mod secret_value_tests {
|
||||||
|
use super::{
|
||||||
|
FunctionRegistrationRequest, MAX_FUNCTION_SECRET_VALUE_BYTES,
|
||||||
|
MAX_FUNCTION_SECRET_VALUES_BYTES, PythonRuntimeSpec,
|
||||||
|
};
|
||||||
|
use crate::Error;
|
||||||
|
|
||||||
|
fn request() -> FunctionRegistrationRequest {
|
||||||
|
FunctionRegistrationRequest::from_json(include_str!(
|
||||||
|
"../tests/fixtures/first_class_functions/v1/remote_function_registration_request.json"
|
||||||
|
))
|
||||||
|
.unwrap()
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn validates_secret_name_and_value_invariants() {
|
||||||
|
let missing = request();
|
||||||
|
assert!(matches!(
|
||||||
|
missing.validate_secret_values(),
|
||||||
|
Err(Error::InvalidInput { message }) if message.contains("exactly match")
|
||||||
|
));
|
||||||
|
|
||||||
|
let mut empty = request();
|
||||||
|
empty
|
||||||
|
.secret_values
|
||||||
|
.insert("API_TOKEN".to_string(), String::new());
|
||||||
|
assert!(matches!(
|
||||||
|
empty.validate_secret_values(),
|
||||||
|
Err(Error::InvalidInput { message }) if message.contains("non-empty")
|
||||||
|
));
|
||||||
|
|
||||||
|
let mut nul = request();
|
||||||
|
nul.secret_values
|
||||||
|
.insert("API_TOKEN".to_string(), "before\0after".to_string());
|
||||||
|
assert!(matches!(
|
||||||
|
nul.validate_secret_values(),
|
||||||
|
Err(Error::InvalidInput { message }) if message.contains("NUL")
|
||||||
|
));
|
||||||
|
|
||||||
|
let mut unexpected = request();
|
||||||
|
unexpected
|
||||||
|
.secret_values
|
||||||
|
.insert("OTHER".to_string(), "value".to_string());
|
||||||
|
assert!(matches!(
|
||||||
|
unexpected.validate_secret_values(),
|
||||||
|
Err(Error::InvalidInput { message }) if message.contains("exactly match")
|
||||||
|
));
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn rejects_invalid_duplicate_and_overlapping_secret_declarations() {
|
||||||
|
let mut invalid_name = request();
|
||||||
|
invalid_name.required_secrets = vec!["BAD=NAME".to_string()];
|
||||||
|
invalid_name
|
||||||
|
.secret_values
|
||||||
|
.insert("BAD=NAME".to_string(), "secret".to_string());
|
||||||
|
|
||||||
|
let mut duplicate = request();
|
||||||
|
duplicate.required_secrets = vec!["API_TOKEN".to_string(), "API_TOKEN".to_string()];
|
||||||
|
duplicate
|
||||||
|
.secret_values
|
||||||
|
.insert("API_TOKEN".to_string(), "secret".to_string());
|
||||||
|
|
||||||
|
let mut overlap = request();
|
||||||
|
overlap
|
||||||
|
.secret_values
|
||||||
|
.insert("API_TOKEN".to_string(), "secret".to_string());
|
||||||
|
if let PythonRuntimeSpec::Python { env, .. } = &mut overlap.runtime {
|
||||||
|
env.insert("API_TOKEN".to_string(), "public".to_string());
|
||||||
|
}
|
||||||
|
|
||||||
|
assert!(matches!(
|
||||||
|
invalid_name.validate_secret_values(),
|
||||||
|
Err(Error::InvalidInput { message }) if message.contains("portable environment variable")
|
||||||
|
));
|
||||||
|
assert!(matches!(
|
||||||
|
duplicate.validate_secret_values(),
|
||||||
|
Err(Error::InvalidInput { message }) if message.contains("duplicate")
|
||||||
|
));
|
||||||
|
assert!(matches!(
|
||||||
|
overlap.validate_secret_values(),
|
||||||
|
Err(Error::InvalidInput { message }) if message.contains("must be disjoint")
|
||||||
|
));
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn enforces_portable_secret_name_boundaries() {
|
||||||
|
for name in ["A", "_", "A0_"] {
|
||||||
|
let mut request = request();
|
||||||
|
request.required_secrets = vec![name.to_string()];
|
||||||
|
request
|
||||||
|
.secret_values
|
||||||
|
.insert(name.to_string(), "secret".to_string());
|
||||||
|
request.validate_secret_values().unwrap();
|
||||||
|
}
|
||||||
|
|
||||||
|
for name in ["", "0TOKEN", "BAD-NAME", "TÖKEN"] {
|
||||||
|
let mut request = request();
|
||||||
|
request.required_secrets = vec![name.to_string()];
|
||||||
|
request
|
||||||
|
.secret_values
|
||||||
|
.insert(name.to_string(), "secret".to_string());
|
||||||
|
assert!(matches!(
|
||||||
|
request.validate_secret_values(),
|
||||||
|
Err(Error::InvalidInput { message })
|
||||||
|
if message.contains("portable environment variable")
|
||||||
|
));
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn accepts_exact_secret_value_utf8_byte_limit() {
|
||||||
|
for value in [
|
||||||
|
"x".repeat(MAX_FUNCTION_SECRET_VALUE_BYTES),
|
||||||
|
"é".repeat(MAX_FUNCTION_SECRET_VALUE_BYTES / "é".len()),
|
||||||
|
] {
|
||||||
|
assert_eq!(value.len(), MAX_FUNCTION_SECRET_VALUE_BYTES);
|
||||||
|
let mut request = request();
|
||||||
|
request.secret_values.insert("API_TOKEN".to_string(), value);
|
||||||
|
request.validate_secret_values().unwrap();
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn rejects_secret_value_over_utf8_byte_limit() {
|
||||||
|
for value in [
|
||||||
|
"x".repeat(MAX_FUNCTION_SECRET_VALUE_BYTES + 1),
|
||||||
|
"é".repeat(MAX_FUNCTION_SECRET_VALUE_BYTES / "é".len() + 1),
|
||||||
|
] {
|
||||||
|
assert!(value.len() > MAX_FUNCTION_SECRET_VALUE_BYTES);
|
||||||
|
let mut request = request();
|
||||||
|
request.secret_values.insert("API_TOKEN".to_string(), value);
|
||||||
|
assert!(matches!(
|
||||||
|
request.validate_secret_values(),
|
||||||
|
Err(Error::InvalidInput { message }) if message.contains("65536-byte limit")
|
||||||
|
));
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn rejects_aggregate_secret_value_bytes_over_server_limit() {
|
||||||
|
let mut request = request();
|
||||||
|
request.required_secrets = (0..9).map(|index| format!("SECRET_{index}")).collect();
|
||||||
|
request.secret_values = request
|
||||||
|
.required_secrets
|
||||||
|
.iter()
|
||||||
|
.map(|name| (name.clone(), "x".repeat(MAX_FUNCTION_SECRET_VALUE_BYTES)))
|
||||||
|
.collect();
|
||||||
|
|
||||||
|
assert!(matches!(
|
||||||
|
request.validate_secret_values(),
|
||||||
|
Err(Error::InvalidInput { message })
|
||||||
|
if message.contains(&format!("{MAX_FUNCTION_SECRET_VALUES_BYTES}-byte request limit"))
|
||||||
|
));
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn accepts_exact_aggregate_secret_value_byte_limit() {
|
||||||
|
let mut request = request();
|
||||||
|
request.required_secrets = (0..8).map(|index| format!("SECRET_{index}")).collect();
|
||||||
|
request.secret_values = request
|
||||||
|
.required_secrets
|
||||||
|
.iter()
|
||||||
|
.map(|name| (name.clone(), "x".repeat(MAX_FUNCTION_SECRET_VALUE_BYTES)))
|
||||||
|
.collect();
|
||||||
|
|
||||||
|
assert_eq!(
|
||||||
|
request
|
||||||
|
.secret_values
|
||||||
|
.values()
|
||||||
|
.map(String::len)
|
||||||
|
.sum::<usize>(),
|
||||||
|
MAX_FUNCTION_SECRET_VALUES_BYTES
|
||||||
|
);
|
||||||
|
request.validate_secret_values().unwrap();
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
#[cfg(test)]
|
#[cfg(test)]
|
||||||
mod conda_environment_tests {
|
mod conda_environment_tests {
|
||||||
use super::PythonEnvironmentSpec;
|
use super::PythonEnvironmentSpec;
|
||||||
|
|||||||
@@ -7,6 +7,7 @@ use reqwest::{
|
|||||||
Body, Request, RequestBuilder, Response,
|
Body, Request, RequestBuilder, Response,
|
||||||
header::{HeaderMap, HeaderValue},
|
header::{HeaderMap, HeaderValue},
|
||||||
};
|
};
|
||||||
|
use serde_json::Value;
|
||||||
use std::{collections::HashMap, future::Future, str::FromStr, sync::Arc, time::Duration};
|
use std::{collections::HashMap, future::Future, str::FromStr, sync::Arc, time::Duration};
|
||||||
|
|
||||||
use crate::error::{Error, Result};
|
use crate::error::{Error, Result};
|
||||||
@@ -14,6 +15,60 @@ use crate::remote::db::RemoteOptions;
|
|||||||
use crate::remote::retry::{ResolvedRetryConfig, RetryCounter};
|
use crate::remote::retry::{ResolvedRetryConfig, RetryCounter};
|
||||||
|
|
||||||
const REQUEST_ID_HEADER: HeaderName = HeaderName::from_static("x-request-id");
|
const REQUEST_ID_HEADER: HeaderName = HeaderName::from_static("x-request-id");
|
||||||
|
const REDACTED_JSON_VALUE: &str = "[REDACTED]";
|
||||||
|
const SUPPRESSED_JSON_BODY: &str = "[JSON BODY SUPPRESSED]";
|
||||||
|
|
||||||
|
fn is_sensitive_json_field(name: &str) -> bool {
|
||||||
|
name.to_ascii_lowercase().contains("secret")
|
||||||
|
}
|
||||||
|
|
||||||
|
fn redact_sensitive_json_fields(value: &mut Value) {
|
||||||
|
match value {
|
||||||
|
Value::Object(fields) => {
|
||||||
|
for (name, child) in fields {
|
||||||
|
if is_sensitive_json_field(name) {
|
||||||
|
*child = Value::String(REDACTED_JSON_VALUE.to_string());
|
||||||
|
} else {
|
||||||
|
redact_sensitive_json_fields(child);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
Value::Array(values) => values.iter_mut().for_each(redact_sensitive_json_fields),
|
||||||
|
_ => {}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
fn redacted_json_body(request: &Request) -> Option<String> {
|
||||||
|
let body = request.body()?.as_bytes()?;
|
||||||
|
let mut value = serde_json::from_slice(body).ok()?;
|
||||||
|
redact_sensitive_json_fields(&mut value);
|
||||||
|
serde_json::to_string(&value).ok()
|
||||||
|
}
|
||||||
|
|
||||||
|
fn request_log_message(request: &Request, request_id: &str) -> String {
|
||||||
|
let prefix = format!(
|
||||||
|
"Sending request_id={}: {} {}",
|
||||||
|
request_id,
|
||||||
|
request.method(),
|
||||||
|
request.url()
|
||||||
|
);
|
||||||
|
let content_type = request
|
||||||
|
.headers()
|
||||||
|
.get("content-type")
|
||||||
|
.and_then(|value| value.to_str().ok())
|
||||||
|
.and_then(|value| value.split(';').next());
|
||||||
|
if content_type.is_some_and(|value| value.eq_ignore_ascii_case("application/json")) {
|
||||||
|
// Never format the raw Request here: its Debug representation is not a
|
||||||
|
// redaction boundary and may include the original body. If the JSON body
|
||||||
|
// cannot be structurally parsed, suppress it instead of logging raw bytes.
|
||||||
|
let body = redacted_json_body(request).unwrap_or_else(|| SUPPRESSED_JSON_BODY.to_string());
|
||||||
|
format!("{prefix} with body {body}")
|
||||||
|
} else {
|
||||||
|
// Method and URL are sufficient request context. Raw Request formatting
|
||||||
|
// may expose headers or a non-JSON body, so it is never a logging fallback.
|
||||||
|
prefix
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
/// Configuration for TLS/mTLS settings.
|
/// Configuration for TLS/mTLS settings.
|
||||||
#[derive(Clone, Debug)]
|
#[derive(Clone, Debug)]
|
||||||
@@ -839,22 +894,9 @@ impl<S: HttpSend> RestfulLanceDbClient<S> {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
pub(crate) fn log_request(&self, request: &Request, request_id: &String) {
|
pub(crate) fn log_request(&self, request: &Request, request_id: &str) {
|
||||||
if log::log_enabled!(log::Level::Debug) {
|
if log::log_enabled!(log::Level::Debug) {
|
||||||
let content_type = request
|
debug!("{}", request_log_message(request, request_id));
|
||||||
.headers()
|
|
||||||
.get("content-type")
|
|
||||||
.map(|v| v.to_str().unwrap());
|
|
||||||
if content_type == Some("application/json") {
|
|
||||||
let body = request.body().as_ref().unwrap().as_bytes().unwrap();
|
|
||||||
let body = String::from_utf8_lossy(body);
|
|
||||||
debug!(
|
|
||||||
"Sending request_id={}: {:?} with body {}",
|
|
||||||
request_id, request, body
|
|
||||||
);
|
|
||||||
} else {
|
|
||||||
debug!("Sending request_id={}: {:?}", request_id, request);
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -1077,6 +1119,49 @@ mod tests {
|
|||||||
ENV_MUTEX.lock().unwrap_or_else(|e| e.into_inner())
|
ENV_MUTEX.lock().unwrap_or_else(|e| e.into_inner())
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn test_request_log_message_redacts_secrets_and_never_formats_raw_requests() {
|
||||||
|
const SECRET_SENTINEL: &str = "udf-secret-log-sentinel-7e4e";
|
||||||
|
const MALFORMED_SENTINEL: &str = "malformed-secret-log-sentinel-b652";
|
||||||
|
const NON_JSON_SENTINEL: &str = "non-json-secret-log-sentinel-7fd1";
|
||||||
|
|
||||||
|
let request = reqwest::Client::new()
|
||||||
|
.post("https://example.com/v1/functions/create")
|
||||||
|
.json(&serde_json::json!({
|
||||||
|
"name": "uses_secret",
|
||||||
|
"nested": {
|
||||||
|
"secret_values": {"OPENAI_API_KEY": SECRET_SENTINEL},
|
||||||
|
"safe": "visible-value"
|
||||||
|
}
|
||||||
|
}))
|
||||||
|
.build()
|
||||||
|
.unwrap();
|
||||||
|
let log_message = request_log_message(&request, "valid-json");
|
||||||
|
|
||||||
|
let malformed_request = reqwest::Client::new()
|
||||||
|
.post("https://example.com/v1/functions/create")
|
||||||
|
.header("content-type", "application/json; charset=utf-8")
|
||||||
|
.body(format!(r#"{{"secret_values":"{MALFORMED_SENTINEL}""#))
|
||||||
|
.build()
|
||||||
|
.unwrap();
|
||||||
|
let malformed_log_message = request_log_message(&malformed_request, "malformed-json");
|
||||||
|
|
||||||
|
let non_json_request = reqwest::Client::new()
|
||||||
|
.post("https://example.com/v1/functions/create")
|
||||||
|
.header("content-type", "text/plain")
|
||||||
|
.body(NON_JSON_SENTINEL)
|
||||||
|
.build()
|
||||||
|
.unwrap();
|
||||||
|
let non_json_log_message = request_log_message(&non_json_request, "non-json");
|
||||||
|
|
||||||
|
assert!(log_message.contains("visible-value"));
|
||||||
|
assert!(log_message.contains(REDACTED_JSON_VALUE));
|
||||||
|
assert!(!log_message.contains(SECRET_SENTINEL));
|
||||||
|
assert!(malformed_log_message.contains(SUPPRESSED_JSON_BODY));
|
||||||
|
assert!(!malformed_log_message.contains(MALFORMED_SENTINEL));
|
||||||
|
assert!(!non_json_log_message.contains(NON_JSON_SENTINEL));
|
||||||
|
}
|
||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
fn test_timeout_config_default() {
|
fn test_timeout_config_default() {
|
||||||
let config = TimeoutConfig::default();
|
let config = TimeoutConfig::default();
|
||||||
|
|||||||
@@ -554,6 +554,7 @@ impl<S: HttpSend> Database for RemoteDatabase<S> {
|
|||||||
&self,
|
&self,
|
||||||
request: FunctionRegistrationRequest,
|
request: FunctionRegistrationRequest,
|
||||||
) -> Result<Job<FunctionVersion>> {
|
) -> Result<Job<FunctionVersion>> {
|
||||||
|
request.validate_secret_values()?;
|
||||||
let req = self.client.post("/v1/functions/create").json(&request);
|
let req = self.client.post("/v1/functions/create").json(&request);
|
||||||
let (request_id, response) = self.client.send(req).await?;
|
let (request_id, response) = self.client.send(req).await?;
|
||||||
let response = self.client.check_response(&request_id, response).await?;
|
let response = self.client.check_response(&request_id, response).await?;
|
||||||
@@ -2642,7 +2643,8 @@ mod tests {
|
|||||||
);
|
);
|
||||||
const FUNCTION_JOB: &str =
|
const FUNCTION_JOB: &str =
|
||||||
include_str!("../../tests/fixtures/first_class_functions/v1/remote_function_job.json");
|
include_str!("../../tests/fixtures/first_class_functions/v1/remote_function_job.json");
|
||||||
let expected: serde_json::Value = serde_json::from_str(REQUEST).unwrap();
|
let mut expected: serde_json::Value = serde_json::from_str(REQUEST).unwrap();
|
||||||
|
expected["secret_values"] = serde_json::json!({"API_TOKEN": "secret-value"});
|
||||||
let conn = Connection::new_with_handler(move |request| match request.url().path() {
|
let conn = Connection::new_with_handler(move |request| match request.url().path() {
|
||||||
"/v1/functions/create" => {
|
"/v1/functions/create" => {
|
||||||
assert_eq!(request.method(), &reqwest::Method::POST);
|
assert_eq!(request.method(), &reqwest::Method::POST);
|
||||||
@@ -2660,7 +2662,10 @@ mod tests {
|
|||||||
.unwrap(),
|
.unwrap(),
|
||||||
path => panic!("unexpected path: {path}"),
|
path => panic!("unexpected path: {path}"),
|
||||||
});
|
});
|
||||||
let request = crate::function::FunctionRegistrationRequest::from_json(REQUEST).unwrap();
|
let mut request = crate::function::FunctionRegistrationRequest::from_json(REQUEST).unwrap();
|
||||||
|
request
|
||||||
|
.secret_values
|
||||||
|
.insert("API_TOKEN".to_string(), "secret-value".to_string());
|
||||||
let job = conn.create_function_async(request).await.unwrap();
|
let job = conn.create_function_async(request).await.unwrap();
|
||||||
assert_eq!(job.id(), Some("job-function-1"));
|
assert_eq!(job.id(), Some("job-function-1"));
|
||||||
let version = job.wait().await.unwrap();
|
let version = job.wait().await.unwrap();
|
||||||
@@ -2668,6 +2673,32 @@ mod tests {
|
|||||||
assert_eq!(version.version(), "fv_01K3EXACT");
|
assert_eq!(version.version(), "fv_01K3EXACT");
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn test_create_function_async_validates_secrets_before_serialization_and_send() {
|
||||||
|
const REQUEST: &str = include_str!(
|
||||||
|
"../../tests/fixtures/first_class_functions/v1/remote_function_registration_request.json"
|
||||||
|
);
|
||||||
|
let sends = Arc::new(AtomicUsize::new(0));
|
||||||
|
let sends_ref = sends.clone();
|
||||||
|
let conn = Connection::new_with_handler(move |_| {
|
||||||
|
sends_ref.fetch_add(1, Ordering::SeqCst);
|
||||||
|
http::Response::builder().status(500).body("").unwrap()
|
||||||
|
});
|
||||||
|
let mut request = crate::function::FunctionRegistrationRequest::from_json(REQUEST).unwrap();
|
||||||
|
request.secret_values.insert(
|
||||||
|
"API_TOKEN".to_string(),
|
||||||
|
"x".repeat(crate::function::MAX_FUNCTION_SECRET_VALUE_BYTES + 1),
|
||||||
|
);
|
||||||
|
|
||||||
|
let error = conn.create_function_async(request).await.unwrap_err();
|
||||||
|
|
||||||
|
assert!(matches!(
|
||||||
|
error,
|
||||||
|
Error::InvalidInput { message } if message.contains("65536-byte limit")
|
||||||
|
));
|
||||||
|
assert_eq!(sends.load(Ordering::SeqCst), 0);
|
||||||
|
}
|
||||||
|
|
||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
async fn test_get_function_requires_and_sends_exact_version() {
|
async fn test_get_function_requires_and_sends_exact_version() {
|
||||||
const VERSION: &str = include_str!(
|
const VERSION: &str = include_str!(
|
||||||
|
|||||||
@@ -20,6 +20,25 @@ fn job_result(name: &str) -> Value {
|
|||||||
serde_json::from_str::<Value>(&fixture(name)).expect("remote Job fixture")["result"].clone()
|
serde_json::from_str::<Value>(&fixture(name)).expect("remote Job fixture")["result"].clone()
|
||||||
}
|
}
|
||||||
|
|
||||||
|
fn assert_no_secret_values(value: &Value) {
|
||||||
|
match value {
|
||||||
|
Value::Object(values) => {
|
||||||
|
for (key, value) in values {
|
||||||
|
assert!(
|
||||||
|
!matches!(
|
||||||
|
key.as_str(),
|
||||||
|
"secret_value" | "secret_values" | "resolved_secret" | "resolved_secrets"
|
||||||
|
),
|
||||||
|
"client canonical value must not model resolved secret material"
|
||||||
|
);
|
||||||
|
assert_no_secret_values(value);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
Value::Array(values) => values.iter().for_each(assert_no_secret_values),
|
||||||
|
_ => {}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
fn function_version_job_result_matches_shared_canonical_golden() {
|
fn function_version_job_result_matches_shared_canonical_golden() {
|
||||||
let result = job_result("remote_function_job.json");
|
let result = job_result("remote_function_job.json");
|
||||||
@@ -28,6 +47,7 @@ fn function_version_job_result_matches_shared_canonical_golden() {
|
|||||||
assert_eq!(version.name(), "embed");
|
assert_eq!(version.name(), "embed");
|
||||||
assert_eq!(version.version(), "fv_01K3EXACT");
|
assert_eq!(version.version(), "fv_01K3EXACT");
|
||||||
assert_eq!(version.runtime_digest(), "sha256:runtime");
|
assert_eq!(version.runtime_digest(), "sha256:runtime");
|
||||||
|
assert_eq!(version.required_secrets(), &["HF_TOKEN"]);
|
||||||
assert_eq!(
|
assert_eq!(
|
||||||
version.to_canonical_json().expect("canonical JSON"),
|
version.to_canonical_json().expect("canonical JSON"),
|
||||||
fixture("remote_function_version.canonical.json").trim()
|
fixture("remote_function_version.canonical.json").trim()
|
||||||
@@ -142,3 +162,21 @@ fn floating_point_application_literals_are_rejected_consistently() {
|
|||||||
.contains("floating-point Function literals")
|
.contains("floating-point Function literals")
|
||||||
);
|
);
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn canonical_client_values_contain_secret_names_only() {
|
||||||
|
let result = job_result("remote_function_job.json");
|
||||||
|
let version = FunctionVersion::from_json(&result.to_string()).expect("FunctionVersion result");
|
||||||
|
let canonical: Value = serde_json::from_str(
|
||||||
|
&version
|
||||||
|
.to_canonical_json()
|
||||||
|
.expect("canonical FunctionVersion"),
|
||||||
|
)
|
||||||
|
.expect("canonical JSON");
|
||||||
|
|
||||||
|
assert_eq!(
|
||||||
|
canonical["required_secrets"],
|
||||||
|
serde_json::json!(["HF_TOKEN"])
|
||||||
|
);
|
||||||
|
assert_no_secret_values(&canonical);
|
||||||
|
}
|
||||||
|
|||||||
@@ -6,6 +6,7 @@ use std::path::PathBuf;
|
|||||||
|
|
||||||
use lancedb::Error;
|
use lancedb::Error;
|
||||||
use lancedb::function::FunctionRegistrationRequest;
|
use lancedb::function::FunctionRegistrationRequest;
|
||||||
|
use serde_json::Value;
|
||||||
|
|
||||||
fn fixture(name: &str) -> String {
|
fn fixture(name: &str) -> String {
|
||||||
let path = PathBuf::from(env!("CARGO_MANIFEST_DIR"))
|
let path = PathBuf::from(env!("CARGO_MANIFEST_DIR"))
|
||||||
@@ -14,6 +15,25 @@ fn fixture(name: &str) -> String {
|
|||||||
fs::read_to_string(path).expect("fixture must be readable")
|
fs::read_to_string(path).expect("fixture must be readable")
|
||||||
}
|
}
|
||||||
|
|
||||||
|
fn assert_no_secret_values(value: &Value) {
|
||||||
|
match value {
|
||||||
|
Value::Object(values) => {
|
||||||
|
for (key, value) in values {
|
||||||
|
assert!(
|
||||||
|
!matches!(
|
||||||
|
key.as_str(),
|
||||||
|
"secret_value" | "secret_values" | "resolved_secret" | "resolved_secrets"
|
||||||
|
),
|
||||||
|
"registration requests must not model resolved secret material"
|
||||||
|
);
|
||||||
|
assert_no_secret_values(value);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
Value::Array(values) => values.iter().for_each(assert_no_secret_values),
|
||||||
|
_ => {}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
fn registration_request_matches_shared_canonical_golden() {
|
fn registration_request_matches_shared_canonical_golden() {
|
||||||
let request = FunctionRegistrationRequest::from_json(&fixture(
|
let request = FunctionRegistrationRequest::from_json(&fixture(
|
||||||
@@ -22,10 +42,33 @@ fn registration_request_matches_shared_canonical_golden() {
|
|||||||
.expect("registration request");
|
.expect("registration request");
|
||||||
assert_eq!(request.name, "normalize_score");
|
assert_eq!(request.name, "normalize_score");
|
||||||
assert_eq!(request.artifact.adapter.kind, "scalar_to_arrow_batch");
|
assert_eq!(request.artifact.adapter.kind, "scalar_to_arrow_batch");
|
||||||
|
assert_eq!(request.required_secrets, ["API_TOKEN"]);
|
||||||
|
assert!(request.secret_values.is_empty());
|
||||||
assert_eq!(
|
assert_eq!(
|
||||||
request.to_canonical_json().expect("canonical request"),
|
request.to_canonical_json().expect("canonical request"),
|
||||||
fixture("remote_function_registration_request.canonical.json").trim()
|
fixture("remote_function_registration_request.canonical.json").trim()
|
||||||
);
|
);
|
||||||
|
|
||||||
|
let value: Value =
|
||||||
|
serde_json::from_str(&request.to_canonical_json().expect("canonical request"))
|
||||||
|
.expect("request JSON");
|
||||||
|
assert_no_secret_values(&value);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn registration_request_serializes_secret_values_but_redacts_debug_output() {
|
||||||
|
let mut value: Value =
|
||||||
|
serde_json::from_str(&fixture("remote_function_registration_request.json")).unwrap();
|
||||||
|
value["secret_values"] = serde_json::json!({"API_TOKEN": "secret-plaintext"});
|
||||||
|
let request = FunctionRegistrationRequest::from_json(&value.to_string()).unwrap();
|
||||||
|
|
||||||
|
assert_eq!(request.secret_values["API_TOKEN"], "secret-plaintext");
|
||||||
|
let canonical = request.to_canonical_json().unwrap();
|
||||||
|
assert!(canonical.contains("secret-plaintext"));
|
||||||
|
let debug = format!("{request:?}");
|
||||||
|
assert!(debug.contains("API_TOKEN"));
|
||||||
|
assert!(debug.contains("[REDACTED]"));
|
||||||
|
assert!(!debug.contains("secret-plaintext"));
|
||||||
}
|
}
|
||||||
|
|
||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
|
|||||||
@@ -24,6 +24,7 @@
|
|||||||
},
|
},
|
||||||
"runtime_digest": "sha256:runtime",
|
"runtime_digest": "sha256:runtime",
|
||||||
"environment_digest": "sha256:environment",
|
"environment_digest": "sha256:environment",
|
||||||
|
"required_secrets": ["HF_TOKEN"],
|
||||||
"created_at": "2026-08-21T00:00:00Z"
|
"created_at": "2026-08-21T00:00:00Z"
|
||||||
},
|
},
|
||||||
"future_job": {"trace_id": "trace-1"}
|
"future_job": {"trace_id": "trace-1"}
|
||||||
|
|||||||
+1
-1
@@ -1 +1 @@
|
|||||||
{"artifact":{"adapter":{"kind":"scalar_to_arrow_batch","version":1},"content":{"data":"ZnJvbSBfX2Z1dHVyZV9fIGltcG9ydCBhbm5vdGF0aW9ucwoKZGVmIG5vcm1hbGl6ZV9zY29yZSh2YWx1ZTogZmxvYXQpIC0+IGZsb2F0OgogICAgcmV0dXJuIHZhbHVlIC8gMTAwLjAK","encoding":"base64"},"digest":"sha256:760784bdcef57b802f389b97804cc0b618aae39e86044733451bcd0b13089a7f","entrypoint":"normalize_score","kind":"python_callable"},"name":"normalize_score","runtime":{"env":{"MODE":"test"},"environment":{"kind":"pip","packages":["numpy>=2"]},"kind":"python","python_version":"3.12"},"signature":{"inputs":[{"arrow_type":"float64","name":"value","nullable":false}],"output":{"arrow_type":"float64","kind":"scalar","nullable":false}}}
|
{"artifact":{"adapter":{"kind":"scalar_to_arrow_batch","version":1},"content":{"data":"ZnJvbSBfX2Z1dHVyZV9fIGltcG9ydCBhbm5vdGF0aW9ucwoKZGVmIG5vcm1hbGl6ZV9zY29yZSh2YWx1ZTogZmxvYXQpIC0+IGZsb2F0OgogICAgcmV0dXJuIHZhbHVlIC8gMTAwLjAK","encoding":"base64"},"digest":"sha256:760784bdcef57b802f389b97804cc0b618aae39e86044733451bcd0b13089a7f","entrypoint":"normalize_score","kind":"python_callable"},"name":"normalize_score","required_secrets":["API_TOKEN"],"runtime":{"env":{"MODE":"test"},"environment":{"kind":"pip","packages":["numpy>=2"]},"kind":"python","python_version":"3.12"},"signature":{"inputs":[{"arrow_type":"float64","name":"value","nullable":false}],"output":{"arrow_type":"float64","kind":"scalar","nullable":false}}}
|
||||||
|
|||||||
Vendored
+4
-1
@@ -39,5 +39,8 @@
|
|||||||
"env": {
|
"env": {
|
||||||
"MODE": "test"
|
"MODE": "test"
|
||||||
}
|
}
|
||||||
}
|
},
|
||||||
|
"required_secrets": [
|
||||||
|
"API_TOKEN"
|
||||||
|
]
|
||||||
}
|
}
|
||||||
|
|||||||
Vendored
+1
-1
@@ -1 +1 @@
|
|||||||
{"artifact":{"digest":"sha256:code","entrypoint":"embed","kind":"python_callable"},"created_at":"2026-08-21T00:00:00Z","environment_digest":"sha256:environment","name":"embed","runtime":{"env":{"TOKENIZERS_PARALLELISM":"false"},"environment":{"kind":"pip","packages":["sentence-transformers>=3"]},"kind":"python","python_version":"3.12"},"runtime_digest":"sha256:runtime","signature":{"inputs":[{"arrow_type":"utf8","name":"text","nullable":true}],"output":{"arrow_type":"list<float32>","kind":"scalar","nullable":false}},"version":"fv_01K3EXACT"}
|
{"artifact":{"digest":"sha256:code","entrypoint":"embed","kind":"python_callable"},"created_at":"2026-08-21T00:00:00Z","environment_digest":"sha256:environment","name":"embed","required_secrets":["HF_TOKEN"],"runtime":{"env":{"TOKENIZERS_PARALLELISM":"false"},"environment":{"kind":"pip","packages":["sentence-transformers>=3"]},"kind":"python","python_version":"3.12"},"runtime_digest":"sha256:runtime","signature":{"inputs":[{"arrow_type":"utf8","name":"text","nullable":true}],"output":{"arrow_type":"list<float32>","kind":"scalar","nullable":false}},"version":"fv_01K3EXACT"}
|
||||||
|
|||||||
Reference in New Issue
Block a user