From 8007c65d32b7dfc33ebfad6021533e9620c9603e Mon Sep 17 00:00:00 2001 From: Philip Cushen Date: Tue, 15 Sep 2026 12:40:35 +0100 Subject: [PATCH] server: --allowed-models to restrict which model paths a request may load On a shared machine every script that names a different model swaps or adds a resident model for everyone else. With --allowed-models the server refuses any model path not listed (the --model given at start is always allowed) before touching weights; the request gets the existing 404 error path. Default is unchanged (load on demand). Fixes #1892. --- mlx_lm/server.py | 28 ++++++++++++++++++++++++++++ tests/test_server.py | 37 +++++++++++++++++++++++++++++++++++++ 2 files changed, 65 insertions(+) diff --git a/mlx_lm/server.py b/mlx_lm/server.py index 9462d6e57..7f59f6902 100644 --- a/mlx_lm/server.py +++ b/mlx_lm/server.py @@ -302,6 +302,17 @@ def __init__(self, cli_args: argparse.Namespace): self._adapter_map["default_model"] = self.cli_args.adapter_path self._draft_model_map["default_model"] = self.cli_args.draft_model + # Optional allow-list: when set, a request may only load these model + # paths (the --model given at start is always allowed). Anything else + # is refused instead of loaded, so one caller naming a different model + # cannot swap out the resident model for everyone else. + allowed = getattr(self.cli_args, "allowed_models", None) + self.allowed_models = None + if allowed: + self.allowed_models = set(allowed) + if self.cli_args.model is not None: + self.allowed_models.add(self.cli_args.model) + # Build the tokenizer config for later use in load self._tokenizer_config = {"trust_remote_code": cli_args.trust_remote_code} if cli_args.chat_template: @@ -383,6 +394,12 @@ def load(self, model_path, adapter_path=None, draft_model_path=None): model_path = self._model_map.get(model_path, model_path) draft_model_path = self._draft_model_map.get(draft_model_path, draft_model_path) + if self.allowed_models is not None and model_path not in self.allowed_models: + raise ValueError( + f"Model '{model_path}' is not in --allowed-models; " + f"this server only serves: {sorted(self.allowed_models)}" + ) + model_key = (model_path, adapter_path, draft_model_path) if self.model_key != model_key: self._load(*model_key) @@ -1777,6 +1794,17 @@ def main(): type=str, help="The path to the MLX model weights, tokenizer, and config", ) + parser.add_argument( + "--allowed-models", + type=str, + nargs="*", + default=None, + help=( + "Restrict which model paths a request may load. The --model given " + "at start is always allowed. Requests naming any other model are " + "refused (HTTP 404 with an error) instead of loaded. Default: any." + ), + ) parser.add_argument( "--adapter-path", type=str, diff --git a/tests/test_server.py b/tests/test_server.py index 3579fae34..726fbcc0c 100644 --- a/tests/test_server.py +++ b/tests/test_server.py @@ -912,3 +912,40 @@ def test_inflight_request_does_not_hang_after_thread_crash(self): if __name__ == "__main__": unittest.main() + + +class TestAllowedModels(unittest.TestCase): + """--allowed-models: a request may only load the listed models (plus the --model + given at start); anything else is refused before any weights are touched.""" + + def _provider(self, model, allowed): + import argparse + from mlx_lm.server import ModelProvider + + args = argparse.Namespace( + model=model, + adapter_path=None, + draft_model=None, + trust_remote_code=False, + chat_template=None, + pipeline=False, + allowed_models=allowed, + ) + provider = ModelProvider(args) + provider._load = lambda *a, **k: None # never load real weights here + return provider + + def test_refuses_model_not_in_list(self): + p = self._provider("a/model", ["b/model"]) + with self.assertRaises(ValueError): + p.load("c/other") + + def test_allows_listed_and_default(self): + p = self._provider("a/model", ["b/model"]) + p.load("b/model") + p.load("a/model") + p.load("default_model") + + def test_no_list_means_any(self): + p = self._provider("a/model", None) + p.load("c/other")