Compare commits

...
14 changed files with 913 additions and 46 deletions
+38 -7
View File
@@ -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)
+106 -6
View File
@@ -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),
+10 -3
View File
@@ -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")
+309 -4
View File
@@ -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;
+100 -15
View File
@@ -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();
+33 -2
View File
@@ -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 @@
{"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}}}
@@ -39,5 +39,8 @@
"env": { "env": {
"MODE": "test" "MODE": "test"
} }
} },
"required_secrets": [
"API_TOKEN"
]
} }
@@ -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"}