From fd636907072cb8b9bc0cefe4e437a5e5374c2edf Mon Sep 17 00:00:00 2001 From: loyce-cheng <129246079+loyce-cheng@users.noreply.github.com> Date: Sat, 3 Oct 2026 03:41:55 +0800 Subject: [PATCH] fix: reject blank model names before registration --- docs-gb/reference/model-settings.md | 7 +++++ docs/reference/model-settings.md | 7 +++++ mlserver/errors.py | 9 +++++++ mlserver/registry.py | 7 ++++- tests/repository/test_load.py | 16 ++++++++--- tests/repository/test_repository.py | 27 +++++++++++++++++++ tests/test_registry.py | 41 ++++++++++++++++++++++++++++- 7 files changed, 109 insertions(+), 5 deletions(-) diff --git a/docs-gb/reference/model-settings.md b/docs-gb/reference/model-settings.md index 7d1998108..b531511b9 100644 --- a/docs-gb/reference/model-settings.md +++ b/docs-gb/reference/model-settings.md @@ -14,6 +14,13 @@ loaded models (unless they get overriden by a `model-settings.json` file). Additionally, if no `model-settings.json` file is found, MLServer will also try to load a _"default"_ model from these environment variables. +Before registering a model, MLServer checks that its name is neither empty nor +entirely whitespace. When loading a `model-settings.json` file, a missing or empty +`name` defaults to the name of the containing folder. When configuring a model +only through environment variables, set `MLSERVER_MODEL_NAME` to a non-empty name. +Models created programmatically must also have a non-empty name before they are +loaded into the registry. + ## Settings ```{eval-rst} diff --git a/docs/reference/model-settings.md b/docs/reference/model-settings.md index 7d1998108..b531511b9 100644 --- a/docs/reference/model-settings.md +++ b/docs/reference/model-settings.md @@ -14,6 +14,13 @@ loaded models (unless they get overriden by a `model-settings.json` file). Additionally, if no `model-settings.json` file is found, MLServer will also try to load a _"default"_ model from these environment variables. +Before registering a model, MLServer checks that its name is neither empty nor +entirely whitespace. When loading a `model-settings.json` file, a missing or empty +`name` defaults to the name of the containing folder. When configuring a model +only through environment variables, set `MLSERVER_MODEL_NAME` to a non-empty name. +Models created programmatically must also have a non-empty name before they are +loaded into the registry. + ## Settings ```{eval-rst} diff --git a/mlserver/errors.py b/mlserver/errors.py index 81cc65efb..fa1471b5e 100644 --- a/mlserver/errors.py +++ b/mlserver/errors.py @@ -17,6 +17,15 @@ def __init__(self, name: str, model_uri: Optional[str] = None): super().__init__(msg, status.HTTP_422_UNPROCESSABLE_ENTITY) +class InvalidModelName(MLServerError): + def __init__(self): + msg = ( + "Model name must not be empty or whitespace. " + "Set 'name' in model-settings.json or MLSERVER_MODEL_NAME." + ) + super().__init__(msg, status.HTTP_422_UNPROCESSABLE_ENTITY) + + class ModelNotFound(MLServerError): def __init__(self, name: str, version: Optional[str] = None): msg = f"Model {name} not found" diff --git a/mlserver/registry.py b/mlserver/registry.py index 262ac626c..262c1caaa 100644 --- a/mlserver/registry.py +++ b/mlserver/registry.py @@ -6,7 +6,7 @@ from .context import model_context from .model import MLModel -from .errors import ModelNotFound +from .errors import InvalidModelName, ModelNotFound from .logging import logger from .settings import ModelSettings @@ -287,6 +287,11 @@ def __init__( self._model_initialiser = model_initialiser async def load(self, model_settings: ModelSettings) -> MLModel: + # Repository loading may infer the name from the model's folder. + # Validate the resolved name before initialisation or registration. + if not model_settings.name.strip(): + raise InvalidModelName() + if model_settings.name not in self._models: self._models[model_settings.name] = SingleModelRegistry( model_settings, diff --git a/tests/repository/test_load.py b/tests/repository/test_load.py index 17f365199..f91be8adc 100644 --- a/tests/repository/test_load.py +++ b/tests/repository/test_load.py @@ -5,6 +5,7 @@ import sys from mlserver.model import MLModel +from mlserver.registry import MultiModelRegistry from mlserver.repository.repository import DEFAULT_MODEL_SETTINGS_FILENAME from mlserver.repository.load import load_model_settings from mlserver.settings import ModelSettings @@ -46,17 +47,21 @@ async def test_load_model_settings( assert model_settings._source == model_settings_path +@pytest.mark.parametrize("name", [None, ""]) async def test_name_fallback( sum_model_settings: ModelSettings, model_folder: str, # This is effectively the Pytest-provided `tmp_path` fixture + name, ): - # Overwrite `model-settings.json` file to be missing the `name` field + # Overwrite `model-settings.json` with a missing or empty `name` field model_settings_path = os.path.join(model_folder, DEFAULT_MODEL_SETTINGS_FILENAME) with open(model_settings_path, "w") as model_settings_file: d = sum_model_settings.model_dump(by_alias=True) - # Remove the `name` field from the JSON representation - del d["name"] + if name is None: + del d["name"] + else: + d["name"] = name # Overwrite the model settings in the temporary path json.dump(d, model_settings_file) @@ -67,6 +72,11 @@ async def test_name_fallback( # `model-settings.json`'s containing folder. assert model_settings.name == os.path.basename(model_folder) + registry = MultiModelRegistry() + model = await registry.load(model_settings) + assert model.ready + assert await registry.get_model(model_settings.name) is model + async def test_load_custom_module( custom_module_settings_path: str, sum_model_settings: ModelSettings diff --git a/tests/repository/test_repository.py b/tests/repository/test_repository.py index 6e8e9586a..725c21962 100644 --- a/tests/repository/test_repository.py +++ b/tests/repository/test_repository.py @@ -7,6 +7,8 @@ DEFAULT_MODEL_SETTINGS_FILENAME, ) from mlserver.settings import ModelSettings, ENV_PREFIX_MODEL_SETTINGS +from mlserver.errors import InvalidModelName +from mlserver.registry import MultiModelRegistry @pytest.fixture @@ -108,6 +110,31 @@ async def test_list_fallback( assert default_model_settings._source is None +@pytest.mark.parametrize("name", [None, "", " \t", "env-model"]) +async def test_register_environment_model_name(monkeypatch, tmp_path, name): + monkeypatch.chdir(tmp_path) + monkeypatch.delenv("MLSERVER_MODEL_NAME", raising=False) + monkeypatch.setenv("MLSERVER_MODEL_IMPLEMENTATION", "mlserver.MLModel") + if name is not None: + monkeypatch.setenv("MLSERVER_MODEL_NAME", name) + repository = SchemalessModelRepository(str(tmp_path)) + registry = MultiModelRegistry() + + settings_list = await repository.list() + + assert len(settings_list) == 1 + if name and name.strip(): + model = await registry.load(settings_list[0]) + assert model.name == name + assert await registry.get_model(name) is model + else: + with pytest.raises( + InvalidModelName, match="Model name must not be empty or whitespace" + ): + await registry.load(settings_list[0]) + assert list(await registry.get_models()) == [] + + async def test_find( model_repository: ModelRepository, sum_model_settings: ModelSettings ): diff --git a/tests/test_registry.py b/tests/test_registry.py index ade680b6c..7138f8f70 100644 --- a/tests/test_registry.py +++ b/tests/test_registry.py @@ -5,7 +5,7 @@ from typing import List, Union from mlserver.model import MLModel -from mlserver.errors import MLServerError, ModelNotFound +from mlserver.errors import InvalidModelName, MLServerError, ModelNotFound from mlserver.registry import MultiModelRegistry, SingleModelRegistry from mlserver.settings import ModelSettings, ModelParameters @@ -42,6 +42,45 @@ async def _async_val(model: MLModel, new_model: MLModel = None) -> MLModel: return model_registry +@pytest.mark.parametrize("name", [None, "", " ", "\t\r\n"]) +async def test_load_invalid_name(model_registry, name, monkeypatch, mocker): + monkeypatch.delenv("MLSERVER_MODEL_NAME", raising=False) + kwargs = {} if name is None else {"name": name} + model_settings = ModelSettings(implementation=MLModel, _env_file=None, **kwargs) + existing_models = list(await model_registry.get_models()) + initialiser = mocker.spy(model_registry, "_model_initialiser") + + with pytest.raises( + InvalidModelName, match="Model name must not be empty or whitespace" + ) as exc: + await model_registry.load(model_settings) + + assert exc.value.status_code == 422 + assert "MLSERVER_MODEL_NAME" in str(exc.value) + initialiser.assert_not_called() + for callback in ( + model_registry._on_model_load + + model_registry._on_model_reload + + model_registry._on_model_unload + ): + callback.assert_not_called() + with pytest.raises(ModelNotFound): + await model_registry.get_models(model_settings.name) + assert list(await model_registry.get_models()) == existing_models + + +@pytest.mark.parametrize("name", ["model", "model-1_v2.3", "模型", " model "]) +async def test_load_valid_name(name): + registry = MultiModelRegistry() + model_settings = ModelSettings(name=name, implementation=MLModel) + + model = await registry.load(model_settings) + + assert model.name == name + assert model.ready + assert await registry.get_model(name) is model + + @pytest.mark.parametrize( "name, version", [