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
28 changes: 28 additions & 0 deletions mlx_lm/server.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down Expand Up @@ -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)
Expand Down Expand Up @@ -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,
Expand Down
37 changes: 37 additions & 0 deletions tests/test_server.py
Original file line number Diff line number Diff line change
Expand Up @@ -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")