Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
7 changes: 7 additions & 0 deletions docs-gb/reference/model-settings.md
Original file line number Diff line number Diff line change
Expand Up @@ -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}
Expand Down
7 changes: 7 additions & 0 deletions docs/reference/model-settings.md
Original file line number Diff line number Diff line change
Expand Up @@ -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}
Expand Down
9 changes: 9 additions & 0 deletions mlserver/errors.py
Original file line number Diff line number Diff line change
Expand Up @@ -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"
Expand Down
7 changes: 6 additions & 1 deletion mlserver/registry.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down Expand Up @@ -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,
Expand Down
16 changes: 13 additions & 3 deletions tests/repository/test_load.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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)
Expand All @@ -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
Expand Down
27 changes: 27 additions & 0 deletions tests/repository/test_repository.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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
):
Expand Down
41 changes: 40 additions & 1 deletion tests/test_registry.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down Expand Up @@ -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",
[
Expand Down