diff --git a/packages/sie_server/src/sie_server/api/options.py b/packages/sie_server/src/sie_server/api/options.py index 9667b3c4f..153ffff33 100644 --- a/packages/sie_server/src/sie_server/api/options.py +++ b/packages/sie_server/src/sie_server/api/options.py @@ -53,7 +53,9 @@ def resolve_runtime_options_with_profile( ) from e overflow_policy = merged.get("overflow_policy") - if overflow_policy is not None and overflow_policy not in VALID_OVERFLOW_POLICIES: + if overflow_policy is not None and ( + not isinstance(overflow_policy, str) or overflow_policy not in VALID_OVERFLOW_POLICIES + ): span.set_attribute("error", "invalid_overflow_policy") raise HTTPException( status_code=status.HTTP_400_BAD_REQUEST, diff --git a/packages/sie_server/tests/api/test_option.py b/packages/sie_server/tests/api/test_option.py new file mode 100644 index 000000000..489a03f82 --- /dev/null +++ b/packages/sie_server/tests/api/test_option.py @@ -0,0 +1,34 @@ +from unittest.mock import MagicMock + +import pytest +from fastapi import HTTPException +from sie_server.api.options import resolve_runtime_options_with_profile +from sie_server.config.model import EmbeddingDim, EncodeTask, ModelConfig, ProfileConfig, Tasks + + +def _make_config() -> ModelConfig: + return ModelConfig( + sie_id="test", + hf_id="org/test", + tasks=Tasks(encode=EncodeTask(dense=EmbeddingDim(dim=8))), + profiles={ + "default": ProfileConfig( + adapter_path="sie_server.adapters.base:ModelAdapter", + max_batch_tokens=8, + ) + }, + ) + + +def test_invalid_overflow_policy_type_returns_400() -> None: + span = MagicMock() + + with pytest.raises(HTTPException) as exc_info: + resolve_runtime_options_with_profile( + _make_config(), + {"overflow_policy": ["truncate_text"]}, + span, + ) + + assert exc_info.value.status_code == 400 + assert exc_info.value.detail["code"] == "INVALID_INPUT"