mirror of
https://github.com/lancedb/lancedb.git
synced 2026-09-04 04:28:44 +00:00
fix(functions): complete secret validation parity
This commit is contained in:
@@ -488,6 +488,7 @@ _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:
|
||||
@@ -1038,8 +1039,16 @@ class UdfDefinition:
|
||||
)
|
||||
|
||||
canonical_values = {}
|
||||
total_bytes = 0
|
||||
for name in sorted(secret_values):
|
||||
canonical_values[name] = _validate_secret_value(name, secret_values[name])
|
||||
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:
|
||||
|
||||
@@ -21,6 +21,7 @@ import pytest
|
||||
import lancedb
|
||||
from lancedb.functions import (
|
||||
_MAX_FUNCTION_SECRET_VALUE_BYTES,
|
||||
_MAX_FUNCTION_SECRET_VALUES_BYTES,
|
||||
FunctionRegistrationRequest,
|
||||
UdfDefinition,
|
||||
udf,
|
||||
@@ -782,6 +783,43 @@ def test_secret_value_rejects_over_utf8_byte_limit_before_json_construction(
|
||||
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):
|
||||
|
||||
Reference in New Issue
Block a user