From 66f2acc108986f16aed5b27bb3ebf1f2cae4c698 Mon Sep 17 00:00:00 2001 From: Sai Kiran Polisetty Date: Tue, 3 Mar 2026 22:36:16 +0530 Subject: [PATCH 01/24] Update --- python/openai/README.md | 152 +++++ .../openai/openai_frontend/engine/engine.py | 18 +- .../openai_frontend/engine/triton_engine.py | 138 ++-- .../fastapi/middleware/api_restriction.py | 5 +- .../fastapi/routers/model_management.py | 90 +++ .../frontend/fastapi_frontend.py | 4 +- python/openai/openai_frontend/main.py | 44 +- python/openai/tests/test_model_management.py | 642 ++++++++++++++++++ python/openai/tests/utils.py | 18 +- 9 files changed, 1066 insertions(+), 45 deletions(-) create mode 100644 python/openai/openai_frontend/frontend/fastapi/routers/model_management.py create mode 100644 python/openai/tests/test_model_management.py diff --git a/python/openai/README.md b/python/openai/README.md index 945b050be0..061d5c6187 100644 --- a/python/openai/README.md +++ b/python/openai/README.md @@ -542,6 +542,155 @@ available arguments and default values. For more information on the `tritonfrontend` python bindings, see the docs [here](https://github.com/triton-inference-server/server/blob/main/docs/customization_guide/tritonfrontend.md). +## Model Management + +The OpenAI-compatible frontend supports explicit model control, allowing you to +dynamically load and unload models at runtime without restarting the server. +This is particularly useful when hosting multiple large models on a shared GPU +cluster and needing to swap models on demand. + +### Model Control Mode + +Use `--model-control-mode` to specify how models are managed at startup and +runtime. The default is `none`. + +| Mode | Behavior | +|------|----------| +| `none` (default) | All models in the repository are loaded at startup. Load/unload APIs are not available. | +| `explicit` | No models are loaded at startup unless specified with `--load-model`. Load and unload are controlled via the management API. | + +> [!NOTE] +> This matches the native `tritonserver --model-control-mode` behavior. +> See [Triton Model Management](../../docs/user_guide/model_management.md) for +> more details. + +### Loading Models at Startup (Explicit Mode) + +When using `--model-control-mode=explicit`, use `--load-model` to specify which +models should be loaded at startup. It may be specified multiple times to load +multiple models. + +```bash +# Start in explicit mode with no models loaded +python3 openai_frontend/main.py \ + --model-repository /path/to/models \ + --tokenizer meta-llama/Meta-Llama-3.1-8B-Instruct \ + --model-control-mode explicit + +# Start in explicit mode and load a specific model at startup +python3 openai_frontend/main.py \ + --model-repository /path/to/models \ + --tokenizer meta-llama/Meta-Llama-3.1-8B-Instruct \ + --model-control-mode explicit \ + --load-model llama-3.1-8b-instruct + +# Load multiple models at startup +python3 openai_frontend/main.py \ + --model-repository /path/to/models \ + --tokenizer meta-llama/Meta-Llama-3.1-8B-Instruct \ + --model-control-mode explicit \ + --load-model model-a \ + --load-model model-b + +# Load ALL models in the repository at startup +python3 openai_frontend/main.py \ + --model-repository /path/to/models \ + --tokenizer meta-llama/Meta-Llama-3.1-8B-Instruct \ + --model-control-mode explicit \ + --load-model '*' +``` + +> [!IMPORTANT] +> `--load-model` requires `--model-control-mode=explicit`. Using `--load-model` +> without it is an error. `--load-model=*` must be the **only** `--load-model` +> argument; combining it with a named model is an error. + +### Dynamic Load / Unload API + +Once the server is running in `explicit` mode, models can be loaded and unloaded +at runtime via these endpoints: + +| Method | Endpoint | Description | +|--------|----------|-------------| +| `POST` | `/v1/models/{model_name}/load` | Load a model. Blocks until model is fully loaded and ready. | +| `POST` | `/v1/models/{model_name}/unload` | Unload a model. Blocks until fully unloaded. In-flight requests complete before removal. | + +Both endpoints return an error if `--model-control-mode` is not `explicit`. + +#### Load a model + +```bash +MODEL="llama-3.1-8b-instruct" +curl -s -X POST http://localhost:9000/v1/models/${MODEL}/load | jq +``` + +
+Example output + +```json +{ + "id": "llama-3.1-8b-instruct", + "object": "model", + "created": 1750000000, + "owned_by": "Triton Inference Server" +} +``` + +
+ +#### Unload a model + +```bash +MODEL="llama-3.1-8b-instruct" +curl -s -X POST http://localhost:9000/v1/models/${MODEL}/unload | jq +``` + +
+Example output + +```json +{ + "status": "success", + "model": "llama-3.1-8b-instruct" +} +``` + +
+ +#### Error cases + +| Scenario | HTTP Status | Detail | +|----------|-------------|--------| +| `--model-control-mode` is not `explicit` | 400 | `Model load/unload requires --model-control-mode=explicit` | +| Model already loaded (duplicate load) | 400 | `Model '' is already loaded` | +| Model not loaded (unload unknown) | 400 | `Unknown model: ` | +| Model not found in repository | 500 | Triton error: `failed to poll from model repository` | + +### Restricting the Management API + +The model management endpoints can be protected with authentication headers using +`--openai-restricted-api model-management`. This is separate from the inference +API restriction so operations teams can lock down management access independently. + +```bash +python3 openai_frontend/main.py \ + --model-repository /path/to/models \ + --tokenizer meta-llama/Meta-Llama-3.1-8B-Instruct \ + --model-control-mode explicit \ + --load-model llama-3.1-8b-instruct \ + --openai-restricted-api model-management admin-key admin-secret +``` + +Clients must then include the header for load/unload requests: + +```bash +curl -H "admin-key: admin-secret" \ + -X POST http://localhost:9000/v1/models/llama-3.1-8b-instruct/load +``` + +See [Limit Endpoint Access](#limit-endpoint-access) for more details on +restricting API groups. + ## Model Parallelism Support - [x] vLLM ([EngineArgs](https://github.com/triton-inference-server/vllm_backend/blob/main/README.md#using-the-vllm-backend)) @@ -797,6 +946,9 @@ Use the `--openai-restricted-api` command-line argument to configure endpoint re - **model-repository**: Model listing and information endpoints - `GET /v1/models` - `GET /v1/models/{model_name}` + - **model-management**: Dynamic model load/unload endpoints (requires `--model-control-mode=explicit`) + - `POST /v1/models/{model_name}/load` + - `POST /v1/models/{model_name}/unload` - **metrics**: Server metrics endpoint - `GET /metrics` - **health**: Health check endpoint diff --git a/python/openai/openai_frontend/engine/engine.py b/python/openai/openai_frontend/engine/engine.py index 2dfeafb1db..da68ac024b 100644 --- a/python/openai/openai_frontend/engine/engine.py +++ b/python/openai/openai_frontend/engine/engine.py @@ -1,4 +1,4 @@ -# Copyright 2024-2025, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# Copyright 2024-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. # # Redistribution and use in source and binary forms, with or without # modification, are permitted provided that the following conditions @@ -100,3 +100,19 @@ def embedding(self, request: CreateEmbeddingRequest) -> CreateEmbeddingResponse: Returns a CreateEmbeddingResponse. """ pass + + async def load_model(self, model_name: str) -> Model: + """ + Loads a model by name. Only available in EXPLICIT model control mode. + Blocks until the model is fully loaded and ready, matching standard + Triton server load behavior. + """ + pass + + async def unload_model(self, model_name: str) -> None: + """ + Unloads a model by name. Only available in EXPLICIT model control mode. + Blocks until the model is fully unloaded, matching standard Triton + server unload behavior. In-flight requests complete before unload. + """ + pass diff --git a/python/openai/openai_frontend/engine/triton_engine.py b/python/openai/openai_frontend/engine/triton_engine.py index d194fc2e1f..867dd7eb8c 100644 --- a/python/openai/openai_frontend/engine/triton_engine.py +++ b/python/openai/openai_frontend/engine/triton_engine.py @@ -27,6 +27,7 @@ from __future__ import annotations +import asyncio import base64 import json import time @@ -141,6 +142,7 @@ def __init__( # now, and won't account for dynamically loading/unloading models. self.create_time = int(time.time()) self.model_metadata = self._get_model_metadata() + self._metadata_lock = asyncio.Lock() self.tool_call_parser = ( ToolParserManager.get_tool_parser_cls(tool_call_parser) if tool_call_parser @@ -493,53 +495,109 @@ def _get_tokenizer(self, tokenizer_name: str): return tokenizer + def _build_model_metadata(self, name: str) -> TritonModelMetadata: + model = self.server.model(name) + backend = model.config()["backend"] + if not backend and model.config()["platform"] == "ensemble": + backend = "ensemble" + print(f"Found model: {name=}, {backend=}") + + lora_configs = _parse_lora_configs( + self.server.options.model_repository, + name, + model.version, + backend if self.backend is None else self.backend, + ) + + echo_tensor_name = None + for input in model.config()["input"]: + if input["name"] in [ + "exclude_input_in_output", + "sampling_param_exclude_input_from_output", + ]: + echo_tensor_name = input["name"] + break + + return TritonModelMetadata( + name=name, + backend=backend, + model=model, + tokenizer=self.tokenizer, + lora_configs=lora_configs, + echo_tensor_name=echo_tensor_name, + create_time=int(time.time()), + inference_request_converter=self._determine_request_converter( + backend, RequestKind.GENERATION + ), + embedding_request_converter=self._determine_request_converter( + backend, RequestKind.EMBEDDING + ), + ) + def _get_model_metadata(self) -> Dict[str, TritonModelMetadata]: - # One tokenizer and creation time shared for all loaded models for now. model_metadata = {} + for name, _ in self.server.models(exclude_not_ready=True).keys(): + model_metadata[name] = self._build_model_metadata(name) + return model_metadata - # Read all triton models and store the necessary metadata for each - for name, _ in self.server.models().keys(): - model = self.server.model(name) - backend = model.config()["backend"] - # Explicitly handle ensembles to avoid any runtime validation errors - if not backend and model.config()["platform"] == "ensemble": - backend = "ensemble" - print(f"Found model: {name=}, {backend=}") - - lora_configs = _parse_lora_configs( - self.server.options.model_repository, - name, - model.version, - backend if self.backend is None else self.backend, + async def load_model(self, model_name: str) -> Model: + if ( + self.server.options.model_control_mode + != tritonserver.ModelControlMode.EXPLICIT + ): + raise ClientError( + "Model load/unload requires --model-control-mode=explicit" ) - echo_tensor_name = None - for input in model.config()["input"]: - if input["name"] in [ - "exclude_input_in_output", - "sampling_param_exclude_input_from_output", - ]: - echo_tensor_name = input["name"] - break - - metadata = TritonModelMetadata( - name=name, - backend=backend, - model=model, - tokenizer=self.tokenizer, - lora_configs=lora_configs, - echo_tensor_name=echo_tensor_name, - create_time=self.create_time, - inference_request_converter=self._determine_request_converter( - backend, RequestKind.GENERATION - ), - embedding_request_converter=self._determine_request_converter( - backend, RequestKind.EMBEDDING - ), + async with self._metadata_lock: + if model_name in self.model_metadata: + raise ClientError(f"Model '{model_name}' is already loaded") + + # Blocking C API call dispatched to thread pool to avoid blocking + # the event loop. The C API blocks until model is fully loaded and + # ready, matching standard Triton server behavior. + try: + metadata = await asyncio.to_thread(self._load_model_sync, model_name) + except tritonserver.TritonError as e: + raise ServerError(f"Failed to load model '{model_name}': {e}") + + self.model_metadata[model_name] = metadata + + return Model( + id=model_name, + created=metadata.create_time, + object=ObjectType.model, + owned_by="Triton Inference Server", + ) + + def _load_model_sync(self, model_name: str) -> TritonModelMetadata: + self.server.load(model_name) + return self._build_model_metadata(model_name) + + async def unload_model(self, model_name: str) -> None: + if ( + self.server.options.model_control_mode + != tritonserver.ModelControlMode.EXPLICIT + ): + raise ClientError( + "Model load/unload requires --model-control-mode=explicit" ) - model_metadata[name] = metadata - return model_metadata + async with self._metadata_lock: + if model_name not in self.model_metadata: + raise ClientError(f"Unknown model: {model_name}") + + # Blocking C API call dispatched to thread pool. The C API handles + # in-flight request draining and conflict resolution internally. + try: + await asyncio.to_thread(self._unload_model_sync, model_name) + except tritonserver.TritonError as e: + raise ServerError(f"Failed to unload model '{model_name}': {e}") + + del self.model_metadata[model_name] + + def _unload_model_sync(self, model_name: str) -> None: + self.server.unload(model_name) def _get_streaming_chat_response_chunk( self, diff --git a/python/openai/openai_frontend/frontend/fastapi/middleware/api_restriction.py b/python/openai/openai_frontend/frontend/fastapi/middleware/api_restriction.py index c2216f9e1d..8fa1e37477 100644 --- a/python/openai/openai_frontend/frontend/fastapi/middleware/api_restriction.py +++ b/python/openai/openai_frontend/frontend/fastapi/middleware/api_restriction.py @@ -1,4 +1,4 @@ -# Copyright 2025, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# Copyright 2025-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. # # Redistribution and use in source and binary forms, with or without # modification, are permitted provided that the following conditions @@ -37,6 +37,9 @@ "POST /v1/embeddings", ], "model-repository": ["GET /v1/models"], + "model-management": [ + "POST /v1/models/", + ], "metrics": ["GET /metrics"], "health": ["GET /health/ready"], } diff --git a/python/openai/openai_frontend/frontend/fastapi/routers/model_management.py b/python/openai/openai_frontend/frontend/fastapi/routers/model_management.py new file mode 100644 index 0000000000..2ddecb8883 --- /dev/null +++ b/python/openai/openai_frontend/frontend/fastapi/routers/model_management.py @@ -0,0 +1,90 @@ +# Copyright 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# +# Redistribution and use in source and binary forms, with or without +# modification, are permitted provided that the following conditions +# are met: +# * Redistributions of source code must retain the above copyright +# notice, this list of conditions and the following disclaimer. +# * Redistributions in binary form must reproduce the above copyright +# notice, this list of conditions and the following disclaimer in the +# documentation and/or other materials provided with the distribution. +# * Neither the name of NVIDIA CORPORATION nor the names of its +# contributors may be used to endorse or promote products derived +# from this software without specific prior written permission. +# +# THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS ``AS IS'' AND ANY +# EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE +# IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR +# PURPOSE ARE DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT OWNER OR +# CONTRIBUTORS BE LIABLE FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, +# EXEMPLARY, OR CONSEQUENTIAL DAMAGES (INCLUDING, BUT NOT LIMITED TO, +# PROCUREMENT OF SUBSTITUTE GOODS OR SERVICES; LOSS OF USE, DATA, OR +# PROFITS; OR BUSINESS INTERRUPTION) HOWEVER CAUSED AND ON ANY THEORY +# OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY, OR TORT +# (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE +# OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE. + +import traceback + +from fastapi import APIRouter, HTTPException, Request +from schemas.openai import Model +from utils.utils import ClientError, ServerError, StatusCode + +router = APIRouter() + + +@router.post( + "/v1/models/{model_name}/load", + response_model=Model, + tags=["Model Management"], +) +async def load_model(model_name: str, raw_request: Request) -> Model: + """ + Loads a model by name. Only available in EXPLICIT model control mode. + Blocks until the model is fully loaded and ready. + """ + if not raw_request.app.engine: + raise HTTPException( + status_code=StatusCode.SERVER_ERROR, + detail="No attached inference engine", + ) + + try: + return await raw_request.app.engine.load_model(model_name) + except ClientError as e: + raise HTTPException(status_code=StatusCode.CLIENT_ERROR, detail=f"{e}") + except ServerError as e: + print(traceback.format_exc()) + raise HTTPException(status_code=StatusCode.SERVER_ERROR, detail=f"{e}") + except Exception as e: + print(traceback.format_exc()) + raise HTTPException(status_code=StatusCode.SERVER_ERROR, detail=f"{e}") + + +@router.post( + "/v1/models/{model_name}/unload", + tags=["Model Management"], +) +async def unload_model(model_name: str, raw_request: Request) -> dict: + """ + Unloads a model by name. Only available in EXPLICIT model control mode. + Blocks until the model is fully unloaded. In-flight requests are allowed + to complete before the model is removed. + """ + if not raw_request.app.engine: + raise HTTPException( + status_code=StatusCode.SERVER_ERROR, + detail="No attached inference engine", + ) + + try: + await raw_request.app.engine.unload_model(model_name) + return {"status": "success", "model": model_name} + except ClientError as e: + raise HTTPException(status_code=StatusCode.CLIENT_ERROR, detail=f"{e}") + except ServerError as e: + print(traceback.format_exc()) + raise HTTPException(status_code=StatusCode.SERVER_ERROR, detail=f"{e}") + except Exception as e: + print(traceback.format_exc()) + raise HTTPException(status_code=StatusCode.SERVER_ERROR, detail=f"{e}") diff --git a/python/openai/openai_frontend/frontend/fastapi_frontend.py b/python/openai/openai_frontend/frontend/fastapi_frontend.py index e5fb01deae..f7fa8aab3c 100644 --- a/python/openai/openai_frontend/frontend/fastapi_frontend.py +++ b/python/openai/openai_frontend/frontend/fastapi_frontend.py @@ -1,4 +1,4 @@ -# Copyright 2024-2025, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# Copyright 2024-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. # # Redistribution and use in source and binary forms, with or without # modification, are permitted provided that the following conditions @@ -38,6 +38,7 @@ chat, completions, embeddings, + model_management, models, observability, ) @@ -101,6 +102,7 @@ def _create_app(self): app.include_router(observability.router) app.include_router(models.router) + app.include_router(model_management.router) app.include_router(completions.router) app.include_router(chat.router) app.include_router(embeddings.router) diff --git a/python/openai/openai_frontend/main.py b/python/openai/openai_frontend/main.py index 88aed600d4..ab3effb6c3 100755 --- a/python/openai/openai_frontend/main.py +++ b/python/openai/openai_frontend/main.py @@ -1,6 +1,6 @@ #!/usr/bin/env python3 -# Copyright 2024-2025, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# Copyright 2024-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. # # Redistribution and use in source and binary forms, with or without # modification, are permitted provided that the following conditions @@ -150,6 +150,31 @@ def parse_args(): default=16, help="The default maximum number of tokens to generate if not specified in the request. The default is 16.", ) + triton_group.add_argument( + "--model-control-mode", + type=str, + default="none", + choices=["none", "explicit"], + help="Specify the mode for model management. Options are 'none', and 'explicit'. " + "The default is 'none'. For 'none', the server will load all models in the model " + "repository at startup and will not make any changes to the load " + "models after that. For 'explicit', model load and unload is initiated by using the " + "model control APIs, and only models specified with --load-model will " + "be loaded at startup.", + ) + triton_group.add_argument( + "--load-model", + type=str, + action="append", + default=None, + help="Name of the model to be loaded on server startup. It may be specified " + "multiple times to add multiple models. To load ALL models at startup, " + "specify '*' as the model name with --load-model=* as the ONLY " + "--load-model argument, this does not imply any pattern matching. " + "Specifying --load-model=* in conjunction with another --load-model " + "argument will result in error. Note that this option will only take " + "effect if --model-control-mode=explicit is true.", + ) # OpenAI-Compatible Frontend (FastAPI) openai_group = parser.add_argument_group("Triton OpenAI-Compatible Frontend") @@ -200,8 +225,25 @@ def main(): args = parse_args() # Initialize a Triton Inference Server pointing at LLM models + model_control_mode = ( + tritonserver.ModelControlMode.EXPLICIT + if args.model_control_mode == "explicit" + else tritonserver.ModelControlMode.NONE + ) + + load_models = args.load_model or [] + if load_models and model_control_mode != tritonserver.ModelControlMode.EXPLICIT: + print( + "Error: Use of '--load-model' requires setting " + "'--model-control-mode=explicit' as well.", + file=sys.stderr, + ) + sys.exit(1) + server: tritonserver.Server = tritonserver.Server( model_repository=args.model_repository, + model_control_mode=model_control_mode, + startup_models=load_models, log_verbose=args.tritonserver_log_verbose_level, log_info=True, log_warn=True, diff --git a/python/openai/tests/test_model_management.py b/python/openai/tests/test_model_management.py new file mode 100644 index 0000000000..bd4162c5c8 --- /dev/null +++ b/python/openai/tests/test_model_management.py @@ -0,0 +1,642 @@ +# Copyright 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# +# Redistribution and use in source and binary forms, with or without +# modification, are permitted provided that the following conditions +# are met: +# * Redistributions of source code must retain the above copyright +# notice, this list of conditions and the following disclaimer. +# * Redistributions in binary form must reproduce the above copyright +# notice, this list of conditions and the following disclaimer in the +# documentation and/or other materials provided with the distribution. +# * Neither the name of NVIDIA CORPORATION nor the names of its +# contributors may be used to endorse or promote products derived +# from this software without specific prior written permission. +# +# THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS ``AS IS'' AND ANY +# EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE +# IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR +# PURPOSE ARE DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT OWNER OR +# CONTRIBUTORS BE LIABLE FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, +# EXEMPLARY, OR CONSEQUENTIAL DAMAGES (INCLUDING, BUT NOT LIMITED TO, +# PROCUREMENT OF SUBSTITUTE GOODS OR SERVICES; LOSS OF USE, DATA, OR +# PROFITS; OR BUSINESS INTERRUPTION) HOWEVER CAUSED AND ON ANY THEORY +# OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY, OR TORT +# (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE +# OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE. + +import concurrent.futures +from pathlib import Path + +import pytest +import requests +from fastapi.testclient import TestClient +from tests.utils import ( + OpenAIServer, + setup_fastapi_app, + setup_server, + setup_server_explicit, +) + +TEST_MODEL_REPOSITORY = str(Path(__file__).parent / "test_models") +TEST_MODEL = "mock_llm" +TEST_MODEL_2 = "identity_py" + + +def _get_model_names(client: TestClient) -> list[str]: + response = client.get("/v1/models") + assert response.status_code == 200 + return [m["id"] for m in response.json()["data"]] + + +def _assert_model_fields(model_data: dict): + assert model_data["id"] + assert model_data["object"] == "model" + assert model_data["created"] > 0 + assert model_data["owned_by"] == "Triton Inference Server" + + +# Test Mode enforcement – NONE mode rejects load/unload API calls +class TestModelManagementNoneMode: + @pytest.fixture(scope="class") + def client(self): + server = setup_server(TEST_MODEL_REPOSITORY) + app = setup_fastapi_app(tokenizer="", server=server, backend=None) + with TestClient(app) as test_client: + yield test_client + server.stop() + + def test_load_and_unload_rejected_in_none_mode(self, client): + for api in ["load", "unload"]: + response = client.post(f"/v1/models/{TEST_MODEL}/{api}") + assert response.status_code == 400 + assert ( + "model load/unload requires --model-control-mode=explicit" + in response.json()["detail"].lower() + ) + + +# Test load / unload operations – EXPLICIT mode +class TestModelManagement: + @pytest.fixture(scope="class") + def client(self): + server = setup_server_explicit(TEST_MODEL_REPOSITORY) + app = setup_fastapi_app(tokenizer="", server=server, backend=None) + with TestClient(app) as test_client: + yield test_client + server.stop() + + def test_load_model(self, client): + # Initially no models loaded + assert _get_model_names(client) == [] + + # Load a model – 200, response has correct OpenAI Model fields + response = client.post(f"/v1/models/{TEST_MODEL}/load") + assert response.status_code == 200 + _assert_model_fields(response.json()) + assert response.json()["id"] == TEST_MODEL + + assert TEST_MODEL in _get_model_names(client) + + # GET /v1/models/{model_name} returns correct info + response = client.get(f"/v1/models/{TEST_MODEL}") + assert response.status_code == 200 + _assert_model_fields(response.json()) + + def test_unload_model(self, client): + if TEST_MODEL not in _get_model_names(client): + response = client.post(f"/v1/models/{TEST_MODEL}/load") + assert response.status_code == 200 + + response = client.post(f"/v1/models/{TEST_MODEL}/unload") + assert response.status_code == 200 + body = response.json() + assert body["status"] == "success" + assert body["model"] == TEST_MODEL + + assert TEST_MODEL not in _get_model_names(client) + + assert client.get(f"/v1/models/{TEST_MODEL}").status_code == 404 + + def test_load_already_loaded_model(self, client): + client.post(f"/v1/models/{TEST_MODEL}/unload") + assert client.post(f"/v1/models/{TEST_MODEL}/load").status_code == 200 + + # Second load of the same model should fail + response = client.post(f"/v1/models/{TEST_MODEL}/load") + assert response.status_code == 400 + assert "already loaded" in response.json()["detail"].lower() + + client.post(f"/v1/models/{TEST_MODEL}/unload") + + def test_unload_unknown_model(self, client): + response = client.post("/v1/models/nonexistent_model/unload") + assert response.status_code == 400 + assert "unknown model" in response.json()["detail"].lower() + + def test_load_nonexistent_model(self, client): + response = client.post("/v1/models/model_not_in_repo/load") + assert response.status_code == 500 + + def test_load_unload_reload(self, client): + assert client.post(f"/v1/models/{TEST_MODEL}/load").status_code == 200 + assert TEST_MODEL in _get_model_names(client) + + assert client.post(f"/v1/models/{TEST_MODEL}/unload").status_code == 200 + assert TEST_MODEL not in _get_model_names(client) + + assert client.post(f"/v1/models/{TEST_MODEL}/load").status_code == 200 + assert TEST_MODEL in _get_model_names(client) + + client.post(f"/v1/models/{TEST_MODEL}/unload") + + def test_load_multiple_models(self, client): + assert client.post(f"/v1/models/{TEST_MODEL}/load").status_code == 200 + assert client.post(f"/v1/models/{TEST_MODEL_2}/load").status_code == 200 + + names = _get_model_names(client) + assert TEST_MODEL in names + assert TEST_MODEL_2 in names + + # Unload one; other remains + client.post(f"/v1/models/{TEST_MODEL}/unload") + names = _get_model_names(client) + assert TEST_MODEL not in names + assert TEST_MODEL_2 in names + + client.post(f"/v1/models/{TEST_MODEL_2}/unload") + + +# Test Sequential and concurrent load/unload +class TestModelManagementConcurrency: + @pytest.fixture(scope="class") + def client(self): + server = setup_server_explicit(TEST_MODEL_REPOSITORY) + app = setup_fastapi_app(tokenizer="", server=server, backend=None) + with TestClient(app) as test_client: + yield test_client + server.stop() + + def test_rapid_sequential_load_unload(self, client): + for _ in range(3): + r = client.post(f"/v1/models/{TEST_MODEL}/load") + assert r.status_code == 200 + assert TEST_MODEL in _get_model_names(client) + + r = client.post(f"/v1/models/{TEST_MODEL}/unload") + assert r.status_code == 200 + assert TEST_MODEL not in _get_model_names(client) + + def test_concurrent_load_different_models(self, client): + with concurrent.futures.ThreadPoolExecutor(max_workers=2) as pool: + f1 = pool.submit(client.post, f"/v1/models/{TEST_MODEL}/load") + f2 = pool.submit(client.post, f"/v1/models/{TEST_MODEL_2}/load") + + assert f1.result().status_code == 200 + assert f2.result().status_code == 200 + names = _get_model_names(client) + assert TEST_MODEL in names + assert TEST_MODEL_2 in names + + client.post(f"/v1/models/{TEST_MODEL}/unload") + client.post(f"/v1/models/{TEST_MODEL_2}/unload") + + def test_concurrent_load_same_model(self, client): + with concurrent.futures.ThreadPoolExecutor(max_workers=2) as pool: + f1 = pool.submit(client.post, f"/v1/models/{TEST_MODEL}/load") + f2 = pool.submit(client.post, f"/v1/models/{TEST_MODEL}/load") + + codes = sorted([f1.result().status_code, f2.result().status_code]) + assert codes == [ + 200, + 400, + ], f"Expected one 200 and one 400 for concurrent loads of the same model, got {codes}" + assert TEST_MODEL in _get_model_names(client) + + client.post(f"/v1/models/{TEST_MODEL}/unload") + + def test_concurrent_unload_same_model(self, client): + client.post(f"/v1/models/{TEST_MODEL}/load") + + with concurrent.futures.ThreadPoolExecutor(max_workers=2) as pool: + f1 = pool.submit(client.post, f"/v1/models/{TEST_MODEL}/unload") + f2 = pool.submit(client.post, f"/v1/models/{TEST_MODEL}/unload") + + # One succeeds (200), the other fails (400 "unknown model") — order is non-deterministic + codes = sorted([f1.result().status_code, f2.result().status_code]) + assert codes == [ + 200, + 400, + ], f"Expected one 200 and one 400 for concurrent unloads of the same model, got {codes}" + assert TEST_MODEL not in _get_model_names(client) + + +# Test "--load-model" and "--model-control-mode" CLI options +@pytest.mark.openai +class TestModelManagementCLIOptions: + def _assert_server_startup_fails(self, args, expected_error: str): + """Helper: verify server fails to start and stderr contains expected_error.""" + with pytest.raises(Exception) as exc_info: + with OpenAIServer(args): + pass + assert expected_error in str(exc_info.value) + + def test_load_model_without_explicit_mode_is_error(self): + """--load-model without --model-control-mode=explicit must exit with error. + Error message matches native tritonserver exactly.""" + self._assert_server_startup_fails( + args=[ + "--model-repository", + TEST_MODEL_REPOSITORY, + "--load-model", + TEST_MODEL, + ], + expected_error="Error: Use of '--load-model' requires setting '--model-control-mode=explicit' as well.", + ) + + def test_explicit_no_load_model_starts_with_no_models(self): + args = [ + "--model-repository", + TEST_MODEL_REPOSITORY, + "--model-control-mode", + "explicit", + ] + with OpenAIServer(args) as openai_server: + r = requests.get(openai_server.url_for("v1", "models"), timeout=10) + assert r.status_code == 200 + assert len(r.json()["data"]) == 0 + + def test_explicit_single_load_model_at_startup(self): + args = [ + "--model-repository", + TEST_MODEL_REPOSITORY, + "--model-control-mode", + "explicit", + "--load-model", + TEST_MODEL, + ] + with OpenAIServer(args) as openai_server: + r = requests.get(openai_server.url_for("v1", "models"), timeout=10) + names = [m["id"] for m in r.json()["data"]] + assert TEST_MODEL in names + assert TEST_MODEL_2 not in names + + def test_explicit_multiple_load_model_at_startup(self): + args = [ + "--model-repository", + TEST_MODEL_REPOSITORY, + "--model-control-mode", + "explicit", + "--load-model", + TEST_MODEL, + "--load-model", + TEST_MODEL_2, + ] + with OpenAIServer(args) as openai_server: + r = requests.get(openai_server.url_for("v1", "models"), timeout=10) + names = [m["id"] for m in r.json()["data"]] + assert TEST_MODEL in names + assert TEST_MODEL_2 in names + + def test_explicit_load_model_wildcard_loads_all(self): + args = [ + "--model-repository", + TEST_MODEL_REPOSITORY, + "--model-control-mode", + "explicit", + "--load-model", + "*", + ] + with OpenAIServer(args) as openai_server: + r = requests.get(openai_server.url_for("v1", "models"), timeout=10) + names = [m["id"] for m in r.json()["data"]] + assert TEST_MODEL in names + assert TEST_MODEL_2 in names + + def test_explicit_load_model_wildcard_with_other_is_error(self): + self._assert_server_startup_fails( + args=[ + "--model-repository", + TEST_MODEL_REPOSITORY, + "--model-control-mode", + "explicit", + "--load-model", + "*", + "--load-model", + TEST_MODEL, + ], + expected_error="Wildcard model name '*' must be the ONLY startup model if specified at all.", + ) + + def test_explicit_dynamic_load_unload_via_api(self): + args = [ + "--model-repository", + TEST_MODEL_REPOSITORY, + "--model-control-mode", + "explicit", + ] + with OpenAIServer(args) as server: + base = server.url_root + + assert ( + requests.post( + f"{base}/v1/models/{TEST_MODEL}/load", timeout=30 + ).status_code + == 200 + ) + assert ( + requests.post( + f"{base}/v1/models/{TEST_MODEL_2}/load", timeout=30 + ).status_code + == 200 + ) + + names = [ + m["id"] + for m in requests.get(f"{base}/v1/models", timeout=10).json()["data"] + ] + assert TEST_MODEL in names + assert TEST_MODEL_2 in names + + assert ( + requests.post( + f"{base}/v1/models/{TEST_MODEL}/unload", timeout=30 + ).status_code + == 200 + ) + + names = [ + m["id"] + for m in requests.get(f"{base}/v1/models", timeout=10).json()["data"] + ] + assert TEST_MODEL not in names + assert TEST_MODEL_2 in names + + +# Test Inference with real LLM backend (vLLM / TRT-LLM) after load/unload +@pytest.mark.openai +class TestModelManagementInference: + @pytest.fixture(scope="class") + def managed_server( + self, + model_repository: str, + tokenizer_model: str, + backend: str, + ): + args = [ + "--model-repository", + model_repository, + "--tokenizer", + tokenizer_model, + "--backend", + backend, + "--model-control-mode", + "explicit", + ] + with OpenAIServer(args) as openai_server: + yield openai_server + + @pytest.fixture(autouse=True) + def ensure_model_unloaded(self, managed_server, model: str): + """Guarantee clean state before and after every test: unload silently + (model may already be unloaded -- that is fine).""" + requests.post(f"{managed_server.url_root}/v1/models/{model}/unload", timeout=60) + yield + requests.post(f"{managed_server.url_root}/v1/models/{model}/unload", timeout=60) + + @staticmethod + def _load(base_url, model_name): + r = requests.post(f"{base_url}/v1/models/{model_name}/load", timeout=120) + assert r.status_code == 200, f"Load failed: {r.text}" + + @staticmethod + def _unload(base_url, model_name): + r = requests.post(f"{base_url}/v1/models/{model_name}/unload", timeout=60) + assert r.status_code == 200, f"Unload failed: {r.text}" + + @staticmethod + def _assert_unknown_model(response): + assert response.status_code == 400 + assert "unknown model" in response.json()["detail"].lower() + + @staticmethod + def _assert_usage(usage): + assert usage is not None + assert usage["prompt_tokens"] > 0 + assert usage["completion_tokens"] > 0 + assert ( + usage["total_tokens"] == usage["prompt_tokens"] + usage["completion_tokens"] + ) + + @staticmethod + def _completions(base_url, model_name, **kwargs): + return requests.post( + f"{base_url}/v1/completions", + json={ + "model": model_name, + "prompt": "What is machine learning?", + "max_tokens": 10, + **kwargs, + }, + timeout=30, + ) + + @staticmethod + def _chat_completions(base_url, model_name, **kwargs): + return requests.post( + f"{base_url}/v1/chat/completions", + json={ + "model": model_name, + "messages": [{"role": "user", "content": "What is machine learning?"}], + "max_tokens": 10, + **kwargs, + }, + timeout=30, + ) + + def test_load_completions(self, managed_server, model: str): + base = managed_server.url_root + self._assert_unknown_model(self._completions(base, model)) + + self._load(base, model) + r = self._completions(base, model) + assert r.status_code == 200 + data = r.json() + assert data["choices"][0]["text"].strip() + assert data["choices"][0]["finish_reason"] == "stop" + self._assert_usage(data["usage"]) + + self._unload(base, model) + + def test_load_chat_completions(self, managed_server, model: str): + base = managed_server.url_root + self._assert_unknown_model(self._chat_completions(base, model)) + + self._load(base, model) + r = self._chat_completions(base, model) + assert r.status_code == 200 + data = r.json() + msg = data["choices"][0]["message"] + assert msg["content"].strip() + assert msg["role"] == "assistant" + assert data["choices"][0]["finish_reason"] == "stop" + self._assert_usage(data["usage"]) + + self._unload(base, model) + + def test_load_streaming_completions(self, managed_server, model: str): + base = managed_server.url_root + self._load(base, model) + + r = requests.post( + f"{base}/v1/completions", + json={ + "model": model, + "prompt": "What is machine learning?", + "max_tokens": 10, + "stream": True, + }, + stream=True, + timeout=30, + ) + assert r.status_code == 200 + chunks = [ + line.removeprefix("data: ") + for line in r.iter_lines(decode_unicode=True) + if line.startswith("data: ") and line.strip() != "data: [DONE]" + ] + assert len(chunks) > 0 + + self._unload(base, model) + + def test_load_streaming_chat_completions(self, managed_server, model: str): + base = managed_server.url_root + self._load(base, model) + + r = requests.post( + f"{base}/v1/chat/completions", + json={ + "model": model, + "messages": [{"role": "user", "content": "What is machine learning?"}], + "max_tokens": 10, + "stream": True, + }, + stream=True, + timeout=30, + ) + assert r.status_code == 200 + chunks = [ + line.removeprefix("data: ") + for line in r.iter_lines(decode_unicode=True) + if line.startswith("data: ") and line.strip() != "data: [DONE]" + ] + assert len(chunks) > 0 + + self._unload(base, model) + + def test_unload_rejects_inference(self, managed_server, model: str): + base = managed_server.url_root + self._load(base, model) + + assert self._completions(base, model).status_code == 200 + assert self._chat_completions(base, model).status_code == 200 + + self._unload(base, model) + + self._assert_unknown_model(self._completions(base, model)) + self._assert_unknown_model(self._chat_completions(base, model)) + + def test_reload_inference(self, managed_server, model: str): + base = managed_server.url_root + + self._load(base, model) + assert self._completions(base, model).status_code == 200 + assert self._chat_completions(base, model).status_code == 200 + + self._unload(base, model) + self._assert_unknown_model(self._completions(base, model)) + + self._load(base, model) + r = self._completions(base, model) + assert r.status_code == 200 + assert r.json()["choices"][0]["text"].strip() + + r = self._chat_completions(base, model) + assert r.status_code == 200 + assert r.json()["choices"][0]["message"]["content"].strip() + + self._unload(base, model) + + def test_model_list_after_load_unload(self, managed_server, model: str): + base = managed_server.url_root + + # No models initially + r = requests.get(f"{base}/v1/models", timeout=10) + assert r.status_code == 200 + assert len(r.json()["data"]) == 0 + + self._load(base, model) + names = [ + m["id"] + for m in requests.get(f"{base}/v1/models", timeout=10).json()["data"] + ] + assert model in names + + self._unload(base, model) + names = [ + m["id"] + for m in requests.get(f"{base}/v1/models", timeout=10).json()["data"] + ] + assert model not in names + + +# Test API restriction for model-management endpoints +@pytest.mark.openai +class TestModelManagementRestriction: + @pytest.fixture(scope="class") + def restricted_server(self): + args = [ + "--model-repository", + TEST_MODEL_REPOSITORY, + "--model-control-mode", + "explicit", + "--load-model", + TEST_MODEL, + "--openai-restricted-api", + "model-management", + "mgmt-key", + "mgmt-secret", + ] + with OpenAIServer(args) as openai_server: + yield openai_server + + def test_load_without_auth_rejected(self, restricted_server): + r = requests.post( + f"{restricted_server.url_root}/v1/models/{TEST_MODEL_2}/load", timeout=10 + ) + assert r.status_code == 401 + + def test_unload_without_auth_rejected(self, restricted_server): + r = requests.post( + f"{restricted_server.url_root}/v1/models/{TEST_MODEL}/unload", timeout=10 + ) + assert r.status_code == 401 + + def test_load_with_valid_auth(self, restricted_server): + headers = {"mgmt-key": "mgmt-secret"} + r = requests.post( + f"{restricted_server.url_root}/v1/models/{TEST_MODEL_2}/load", + headers=headers, + timeout=30, + ) + assert r.status_code == 200 + + # Cleanup + requests.post( + f"{restricted_server.url_root}/v1/models/{TEST_MODEL_2}/unload", + headers=headers, + timeout=30, + ) + + def test_model_list_unrestricted(self, restricted_server): + r = requests.get(f"{restricted_server.url_root}/v1/models", timeout=10) + assert r.status_code == 200 diff --git a/python/openai/tests/utils.py b/python/openai/tests/utils.py index 0088fc0571..4068614634 100644 --- a/python/openai/tests/utils.py +++ b/python/openai/tests/utils.py @@ -53,6 +53,21 @@ def setup_server(model_repository: str): return server +def setup_server_explicit( + model_repository: str, load_models: Optional[List[str]] = None +): + server: tritonserver.Server = tritonserver.Server( + model_repository=model_repository, + model_control_mode=tritonserver.ModelControlMode.EXPLICIT, + startup_models=load_models or [], + log_verbose=0, + log_info=True, + log_warn=True, + log_error=True, + ).start(wait_until_ready=True) + return server + + def setup_fastapi_app( tokenizer: str, server: tritonserver.Server, @@ -78,11 +93,12 @@ def __init__( self, cli_args: List[str], *, + port: int = 9000, env_dict: Optional[Dict[str, str]] = None, ) -> None: # TODO: Incorporate caller's cli_args passed to this instance instead self.host = "localhost" - self.port = 9000 + self.port = port env = os.environ.copy() if env_dict is not None: From e56b855bf4bcd20151720dc4effdbea3a6b78bae Mon Sep 17 00:00:00 2001 From: Sai Kiran Polisetty Date: Tue, 3 Mar 2026 22:51:26 +0530 Subject: [PATCH 02/24] Update --- python/openai/openai_frontend/engine/triton_engine.py | 2 -- 1 file changed, 2 deletions(-) diff --git a/python/openai/openai_frontend/engine/triton_engine.py b/python/openai/openai_frontend/engine/triton_engine.py index 867dd7eb8c..941d3d79a6 100644 --- a/python/openai/openai_frontend/engine/triton_engine.py +++ b/python/openai/openai_frontend/engine/triton_engine.py @@ -138,8 +138,6 @@ def __init__( self.lora_separator = lora_separator self.default_max_tokens = default_max_tokens - # NOTE: Creation time and model metadata will be static at startup for - # now, and won't account for dynamically loading/unloading models. self.create_time = int(time.time()) self.model_metadata = self._get_model_metadata() self._metadata_lock = asyncio.Lock() From 3fadc999f0ad78856b805bb3d4c42ed5a47ed36d Mon Sep 17 00:00:00 2001 From: Sai Kiran Polisetty Date: Tue, 3 Mar 2026 22:54:04 +0530 Subject: [PATCH 03/24] Update --- python/openai/openai_frontend/engine/triton_engine.py | 1 + 1 file changed, 1 insertion(+) diff --git a/python/openai/openai_frontend/engine/triton_engine.py b/python/openai/openai_frontend/engine/triton_engine.py index 941d3d79a6..cfd0126ed1 100644 --- a/python/openai/openai_frontend/engine/triton_engine.py +++ b/python/openai/openai_frontend/engine/triton_engine.py @@ -533,6 +533,7 @@ def _build_model_metadata(self, name: str) -> TritonModelMetadata: ) def _get_model_metadata(self) -> Dict[str, TritonModelMetadata]: + # One tokenizer and creation time shared for all loaded models for now. model_metadata = {} for name, _ in self.server.models(exclude_not_ready=True).keys(): model_metadata[name] = self._build_model_metadata(name) From 10a77ef5431f5a05431a72dedffad9b1c1d826a5 Mon Sep 17 00:00:00 2001 From: Sai Kiran Polisetty Date: Tue, 3 Mar 2026 22:56:49 +0530 Subject: [PATCH 04/24] Update python/openai/openai_frontend/main.py Co-authored-by: Copilot <175728472+Copilot@users.noreply.github.com> --- python/openai/openai_frontend/main.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/python/openai/openai_frontend/main.py b/python/openai/openai_frontend/main.py index ab3effb6c3..c96d154fc2 100755 --- a/python/openai/openai_frontend/main.py +++ b/python/openai/openai_frontend/main.py @@ -173,7 +173,7 @@ def parse_args(): "--load-model argument, this does not imply any pattern matching. " "Specifying --load-model=* in conjunction with another --load-model " "argument will result in error. Note that this option will only take " - "effect if --model-control-mode=explicit is true.", + "effect if --model-control-mode is set to 'explicit'.", ) # OpenAI-Compatible Frontend (FastAPI) From 7c20f6e08d68206b97ea1f204e8dc405e76fb39e Mon Sep 17 00:00:00 2001 From: Sai Kiran Polisetty Date: Tue, 3 Mar 2026 23:02:01 +0530 Subject: [PATCH 05/24] Update --- .../openai/openai_frontend/frontend/fastapi/routers/models.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/python/openai/openai_frontend/frontend/fastapi/routers/models.py b/python/openai/openai_frontend/frontend/fastapi/routers/models.py index 39e2f26d53..a0f5bad9e0 100644 --- a/python/openai/openai_frontend/frontend/fastapi/routers/models.py +++ b/python/openai/openai_frontend/frontend/fastapi/routers/models.py @@ -36,7 +36,7 @@ @router.get("/v1/models", response_model=ListModelsResponse, tags=["Models"]) -def list_models(request: Request) -> ListModelsResponse: +async def list_models(request: Request) -> ListModelsResponse: """ Lists the currently available models, and provides basic information about each one such as the owner and availability. """ @@ -50,7 +50,7 @@ def list_models(request: Request) -> ListModelsResponse: @router.get("/v1/models/{model_name}", response_model=Model, tags=["Models"]) -def retrieve_model(request: Request, model_name: str) -> Model: +async def retrieve_model(request: Request, model_name: str) -> Model: """ Retrieves a model instance, providing basic information about the model such as the owner and permissioning. """ From 3415d92f668d3172beeee3eba3f96d381c02ef18 Mon Sep 17 00:00:00 2001 From: Sai Kiran Polisetty Date: Tue, 3 Mar 2026 23:08:03 +0530 Subject: [PATCH 06/24] Update --- .../openai/openai_frontend/frontend/fastapi/routers/models.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/python/openai/openai_frontend/frontend/fastapi/routers/models.py b/python/openai/openai_frontend/frontend/fastapi/routers/models.py index a0f5bad9e0..871e300276 100644 --- a/python/openai/openai_frontend/frontend/fastapi/routers/models.py +++ b/python/openai/openai_frontend/frontend/fastapi/routers/models.py @@ -1,4 +1,4 @@ -# Copyright 2024-2025, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# Copyright 2024-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. # # Redistribution and use in source and binary forms, with or without # modification, are permitted provided that the following conditions From 3fa4695e41a8389e4a91fdb68ae59cabe0f8a94c Mon Sep 17 00:00:00 2001 From: Sai Kiran Polisetty Date: Thu, 5 Mar 2026 23:36:23 +0530 Subject: [PATCH 07/24] Update --- python/openai/tests/test_model_management.py | 25 ++++++++++++++++---- 1 file changed, 21 insertions(+), 4 deletions(-) diff --git a/python/openai/tests/test_model_management.py b/python/openai/tests/test_model_management.py index bd4162c5c8..39ccc5047c 100644 --- a/python/openai/tests/test_model_management.py +++ b/python/openai/tests/test_model_management.py @@ -25,6 +25,7 @@ # OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE. import concurrent.futures +import os from pathlib import Path import pytest @@ -40,6 +41,10 @@ TEST_MODEL_REPOSITORY = str(Path(__file__).parent / "test_models") TEST_MODEL = "mock_llm" TEST_MODEL_2 = "identity_py" +MODEL_MGMT_LOAD_TIMEOUT_S = int(os.getenv("TEST_MODEL_MANAGEMENT_LOAD_TIMEOUT", "300")) +MODEL_MGMT_UNLOAD_TIMEOUT_S = int( + os.getenv("TEST_MODEL_MANAGEMENT_UNLOAD_TIMEOUT", "120") +) def _get_model_names(client: TestClient) -> list[str]: @@ -399,18 +404,30 @@ def managed_server( def ensure_model_unloaded(self, managed_server, model: str): """Guarantee clean state before and after every test: unload silently (model may already be unloaded -- that is fine).""" - requests.post(f"{managed_server.url_root}/v1/models/{model}/unload", timeout=60) + requests.post( + f"{managed_server.url_root}/v1/models/{model}/unload", + timeout=MODEL_MGMT_UNLOAD_TIMEOUT_S, + ) yield - requests.post(f"{managed_server.url_root}/v1/models/{model}/unload", timeout=60) + requests.post( + f"{managed_server.url_root}/v1/models/{model}/unload", + timeout=MODEL_MGMT_UNLOAD_TIMEOUT_S, + ) @staticmethod def _load(base_url, model_name): - r = requests.post(f"{base_url}/v1/models/{model_name}/load", timeout=120) + r = requests.post( + f"{base_url}/v1/models/{model_name}/load", + timeout=MODEL_MGMT_LOAD_TIMEOUT_S, + ) assert r.status_code == 200, f"Load failed: {r.text}" @staticmethod def _unload(base_url, model_name): - r = requests.post(f"{base_url}/v1/models/{model_name}/unload", timeout=60) + r = requests.post( + f"{base_url}/v1/models/{model_name}/unload", + timeout=MODEL_MGMT_UNLOAD_TIMEOUT_S, + ) assert r.status_code == 200, f"Unload failed: {r.text}" @staticmethod From bb6538a39be0f4dceb87b2badd7d8091ce74ad68 Mon Sep 17 00:00:00 2001 From: Sai Kiran Polisetty Date: Thu, 12 Mar 2026 15:41:10 +0530 Subject: [PATCH 08/24] Update --- python/openai/tests/test_model_management.py | 11 ++++++++--- python/openai/tests/utils.py | 19 +++++-------------- 2 files changed, 13 insertions(+), 17 deletions(-) diff --git a/python/openai/tests/test_model_management.py b/python/openai/tests/test_model_management.py index 39ccc5047c..d47eb48f36 100644 --- a/python/openai/tests/test_model_management.py +++ b/python/openai/tests/test_model_management.py @@ -35,7 +35,6 @@ OpenAIServer, setup_fastapi_app, setup_server, - setup_server_explicit, ) TEST_MODEL_REPOSITORY = str(Path(__file__).parent / "test_models") @@ -84,7 +83,10 @@ def test_load_and_unload_rejected_in_none_mode(self, client): class TestModelManagement: @pytest.fixture(scope="class") def client(self): - server = setup_server_explicit(TEST_MODEL_REPOSITORY) + server = setup_server( + TEST_MODEL_REPOSITORY, + model_control_mode=tritonserver.ModelControlMode.EXPLICIT, + ) app = setup_fastapi_app(tokenizer="", server=server, backend=None) with TestClient(app) as test_client: yield test_client @@ -175,7 +177,10 @@ def test_load_multiple_models(self, client): class TestModelManagementConcurrency: @pytest.fixture(scope="class") def client(self): - server = setup_server_explicit(TEST_MODEL_REPOSITORY) + server = setup_server( + TEST_MODEL_REPOSITORY, + model_control_mode=tritonserver.ModelControlMode.EXPLICIT, + ) app = setup_fastapi_app(tokenizer="", server=server, backend=None) with TestClient(app) as test_client: yield test_client diff --git a/python/openai/tests/utils.py b/python/openai/tests/utils.py index 4068614634..6d59e399d1 100644 --- a/python/openai/tests/utils.py +++ b/python/openai/tests/utils.py @@ -42,23 +42,14 @@ # TODO: Cleanup, refactor, mock, etc. -def setup_server(model_repository: str): - server: tritonserver.Server = tritonserver.Server( - model_repository=model_repository, - log_verbose=0, - log_info=True, - log_warn=True, - log_error=True, - ).start(wait_until_ready=True) - return server - - -def setup_server_explicit( - model_repository: str, load_models: Optional[List[str]] = None +def setup_server( + model_repository: str, + model_control_mode: tritonserver.ModelControlMode = tritonserver.ModelControlMode.NONE, + load_models: Optional[List[str]] = None, ): server: tritonserver.Server = tritonserver.Server( model_repository=model_repository, - model_control_mode=tritonserver.ModelControlMode.EXPLICIT, + model_control_mode=model_control_mode, startup_models=load_models or [], log_verbose=0, log_info=True, From c37b6390cc950ca0b0bdefc450e1c590cf2917e2 Mon Sep 17 00:00:00 2001 From: Sai Kiran Polisetty Date: Thu, 12 Mar 2026 15:42:51 +0530 Subject: [PATCH 09/24] Fix pre-commit errors --- python/openai/tests/test_model_management.py | 5 +---- 1 file changed, 1 insertion(+), 4 deletions(-) diff --git a/python/openai/tests/test_model_management.py b/python/openai/tests/test_model_management.py index d47eb48f36..d855b452f5 100644 --- a/python/openai/tests/test_model_management.py +++ b/python/openai/tests/test_model_management.py @@ -31,10 +31,7 @@ import pytest import requests from fastapi.testclient import TestClient -from tests.utils import ( - OpenAIServer, - setup_fastapi_app, - setup_server, +from tests.utils import OpenAIServer, setup_fastapi_app, setup_server, ) TEST_MODEL_REPOSITORY = str(Path(__file__).parent / "test_models") From b9e85cef3098cb546f0726d3c259b467699e645f Mon Sep 17 00:00:00 2001 From: Sai Kiran Polisetty Date: Thu, 12 Mar 2026 15:46:53 +0530 Subject: [PATCH 10/24] Update --- python/openai/tests/test_model_management.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/python/openai/tests/test_model_management.py b/python/openai/tests/test_model_management.py index d855b452f5..99f53e0fad 100644 --- a/python/openai/tests/test_model_management.py +++ b/python/openai/tests/test_model_management.py @@ -30,9 +30,9 @@ import pytest import requests +import tritonserver from fastapi.testclient import TestClient -from tests.utils import OpenAIServer, setup_fastapi_app, setup_server, -) +from tests.utils import OpenAIServer, setup_fastapi_app, setup_server TEST_MODEL_REPOSITORY = str(Path(__file__).parent / "test_models") TEST_MODEL = "mock_llm" From ded522a40055c4c527594232a24afa27e6ad6618 Mon Sep 17 00:00:00 2001 From: Sai Kiran Polisetty Date: Thu, 12 Mar 2026 19:57:31 +0530 Subject: [PATCH 11/24] Update --- python/openai/README.md | 15 ++-- .../fastapi/middleware/api_restriction.py | 4 +- python/openai/tests/test_model_management.py | 88 ++++--------------- .../tests/test_openai_restricted_apis.py | 60 +++++++++++++ 4 files changed, 87 insertions(+), 80 deletions(-) diff --git a/python/openai/README.md b/python/openai/README.md index 061d5c6187..103a8bcfb0 100644 --- a/python/openai/README.md +++ b/python/openai/README.md @@ -603,7 +603,7 @@ python3 openai_frontend/main.py \ > [!IMPORTANT] > `--load-model` requires `--model-control-mode=explicit`. Using `--load-model` > without it is an error. `--load-model=*` must be the **only** `--load-model` -> argument; combining it with a named model is an error. +> argument and combining it with a named model is an error. ### Dynamic Load / Unload API @@ -666,11 +666,10 @@ curl -s -X POST http://localhost:9000/v1/models/${MODEL}/unload | jq | Model not loaded (unload unknown) | 400 | `Unknown model: ` | | Model not found in repository | 500 | Triton error: `failed to poll from model repository` | -### Restricting the Management API +### Restricting the Model Management APIs The model management endpoints can be protected with authentication headers using -`--openai-restricted-api model-management`. This is separate from the inference -API restriction so operations teams can lock down management access independently. +`--openai-restricted-api model-repository`. ```bash python3 openai_frontend/main.py \ @@ -678,10 +677,11 @@ python3 openai_frontend/main.py \ --tokenizer meta-llama/Meta-Llama-3.1-8B-Instruct \ --model-control-mode explicit \ --load-model llama-3.1-8b-instruct \ - --openai-restricted-api model-management admin-key admin-secret + --openai-restricted-api model-repository admin-key admin-secret ``` -Clients must then include the header for load/unload requests: +Clients must then include the header for model management requests such as +load/unload: ```bash curl -H "admin-key: admin-secret" \ @@ -943,10 +943,9 @@ Use the `--openai-restricted-api` command-line argument to configure endpoint re - `POST /v1/completions` - **embedding**: Embedding endpoint - `POST /v1/embeddings` - - **model-repository**: Model listing and information endpoints + - **model-repository**: Model listing, information, and dynamic load/unload endpoints - `GET /v1/models` - `GET /v1/models/{model_name}` - - **model-management**: Dynamic model load/unload endpoints (requires `--model-control-mode=explicit`) - `POST /v1/models/{model_name}/load` - `POST /v1/models/{model_name}/unload` - **metrics**: Server metrics endpoint diff --git a/python/openai/openai_frontend/frontend/fastapi/middleware/api_restriction.py b/python/openai/openai_frontend/frontend/fastapi/middleware/api_restriction.py index 8fa1e37477..085434ad46 100644 --- a/python/openai/openai_frontend/frontend/fastapi/middleware/api_restriction.py +++ b/python/openai/openai_frontend/frontend/fastapi/middleware/api_restriction.py @@ -36,8 +36,8 @@ "POST /v1/completions", "POST /v1/embeddings", ], - "model-repository": ["GET /v1/models"], - "model-management": [ + "model-repository": [ + "GET /v1/models", "POST /v1/models/", ], "metrics": ["GET /metrics"], diff --git a/python/openai/tests/test_model_management.py b/python/openai/tests/test_model_management.py index 99f53e0fad..b085940744 100644 --- a/python/openai/tests/test_model_management.py +++ b/python/openai/tests/test_model_management.py @@ -383,7 +383,7 @@ def test_explicit_dynamic_load_unload_via_api(self): @pytest.mark.openai class TestModelManagementInference: @pytest.fixture(scope="class") - def managed_server( + def server_with_explicit_mode( self, model_repository: str, tokenizer_model: str, @@ -403,16 +403,16 @@ def managed_server( yield openai_server @pytest.fixture(autouse=True) - def ensure_model_unloaded(self, managed_server, model: str): + def ensure_model_unloaded(self, server_with_explicit_mode, model: str): """Guarantee clean state before and after every test: unload silently (model may already be unloaded -- that is fine).""" requests.post( - f"{managed_server.url_root}/v1/models/{model}/unload", + f"{server_with_explicit_mode.url_root}/v1/models/{model}/unload", timeout=MODEL_MGMT_UNLOAD_TIMEOUT_S, ) yield requests.post( - f"{managed_server.url_root}/v1/models/{model}/unload", + f"{server_with_explicit_mode.url_root}/v1/models/{model}/unload", timeout=MODEL_MGMT_UNLOAD_TIMEOUT_S, ) @@ -472,8 +472,8 @@ def _chat_completions(base_url, model_name, **kwargs): timeout=30, ) - def test_load_completions(self, managed_server, model: str): - base = managed_server.url_root + def test_load_completions(self, server_with_explicit_mode, model: str): + base = server_with_explicit_mode.url_root self._assert_unknown_model(self._completions(base, model)) self._load(base, model) @@ -486,8 +486,8 @@ def test_load_completions(self, managed_server, model: str): self._unload(base, model) - def test_load_chat_completions(self, managed_server, model: str): - base = managed_server.url_root + def test_load_chat_completions(self, server_with_explicit_mode, model: str): + base = server_with_explicit_mode.url_root self._assert_unknown_model(self._chat_completions(base, model)) self._load(base, model) @@ -502,8 +502,8 @@ def test_load_chat_completions(self, managed_server, model: str): self._unload(base, model) - def test_load_streaming_completions(self, managed_server, model: str): - base = managed_server.url_root + def test_load_streaming_completions(self, server_with_explicit_mode, model: str): + base = server_with_explicit_mode.url_root self._load(base, model) r = requests.post( @@ -527,8 +527,8 @@ def test_load_streaming_completions(self, managed_server, model: str): self._unload(base, model) - def test_load_streaming_chat_completions(self, managed_server, model: str): - base = managed_server.url_root + def test_load_streaming_chat_completions(self, server_with_explicit_mode, model: str): + base = server_with_explicit_mode.url_root self._load(base, model) r = requests.post( @@ -552,8 +552,8 @@ def test_load_streaming_chat_completions(self, managed_server, model: str): self._unload(base, model) - def test_unload_rejects_inference(self, managed_server, model: str): - base = managed_server.url_root + def test_unload_rejects_inference(self, server_with_explicit_mode, model: str): + base = server_with_explicit_mode.url_root self._load(base, model) assert self._completions(base, model).status_code == 200 @@ -564,8 +564,8 @@ def test_unload_rejects_inference(self, managed_server, model: str): self._assert_unknown_model(self._completions(base, model)) self._assert_unknown_model(self._chat_completions(base, model)) - def test_reload_inference(self, managed_server, model: str): - base = managed_server.url_root + def test_reload_inference(self, server_with_explicit_mode, model: str): + base = server_with_explicit_mode.url_root self._load(base, model) assert self._completions(base, model).status_code == 200 @@ -585,8 +585,8 @@ def test_reload_inference(self, managed_server, model: str): self._unload(base, model) - def test_model_list_after_load_unload(self, managed_server, model: str): - base = managed_server.url_root + def test_model_list_after_load_unload(self, server_with_explicit_mode, model: str): + base = server_with_explicit_mode.url_root # No models initially r = requests.get(f"{base}/v1/models", timeout=10) @@ -607,55 +607,3 @@ def test_model_list_after_load_unload(self, managed_server, model: str): ] assert model not in names - -# Test API restriction for model-management endpoints -@pytest.mark.openai -class TestModelManagementRestriction: - @pytest.fixture(scope="class") - def restricted_server(self): - args = [ - "--model-repository", - TEST_MODEL_REPOSITORY, - "--model-control-mode", - "explicit", - "--load-model", - TEST_MODEL, - "--openai-restricted-api", - "model-management", - "mgmt-key", - "mgmt-secret", - ] - with OpenAIServer(args) as openai_server: - yield openai_server - - def test_load_without_auth_rejected(self, restricted_server): - r = requests.post( - f"{restricted_server.url_root}/v1/models/{TEST_MODEL_2}/load", timeout=10 - ) - assert r.status_code == 401 - - def test_unload_without_auth_rejected(self, restricted_server): - r = requests.post( - f"{restricted_server.url_root}/v1/models/{TEST_MODEL}/unload", timeout=10 - ) - assert r.status_code == 401 - - def test_load_with_valid_auth(self, restricted_server): - headers = {"mgmt-key": "mgmt-secret"} - r = requests.post( - f"{restricted_server.url_root}/v1/models/{TEST_MODEL_2}/load", - headers=headers, - timeout=30, - ) - assert r.status_code == 200 - - # Cleanup - requests.post( - f"{restricted_server.url_root}/v1/models/{TEST_MODEL_2}/unload", - headers=headers, - timeout=30, - ) - - def test_model_list_unrestricted(self, restricted_server): - r = requests.get(f"{restricted_server.url_root}/v1/models", timeout=10) - assert r.status_code == 200 diff --git a/python/openai/tests/test_openai_restricted_apis.py b/python/openai/tests/test_openai_restricted_apis.py index 8e26addb11..f89d63fe38 100755 --- a/python/openai/tests/test_openai_restricted_apis.py +++ b/python/openai/tests/test_openai_restricted_apis.py @@ -150,6 +150,25 @@ def verify_model_repository_endpoints( ) +def verify_model_repository_management_endpoints( + base_url, model, headers, expected_success, description_prefix +): + for endpoint in ["load", "unload"]: + response = requests.post( + f"{base_url}/v1/models/{model}/{endpoint}", headers=headers, timeout=120 + ) + if expected_success: + assert_response_success( + response, + description=f"{description_prefix} - Model {endpoint} endpoint", + ) + else: + assert_response_unauthorized( + response, + description=f"{description_prefix} - Model {endpoint} endpoint", + ) + + def verify_metrics_endpoint(base_url, headers, expected_success, description_prefix): # Test metrics endpoint response = make_get_request(base_url, "/metrics", headers=headers) @@ -331,6 +350,47 @@ def test_unrestricted_endpoints(self, server_with_restrictions): ) +@pytest.mark.openai +class TestOpenAIServerModelRepositoryRestriction: + """Model-repository restriction for dynamic model load/unload.""" + + @pytest.fixture(scope="class") + def server_with_restrictions_explicit_mode(self): + args = [ + "--model-repository", + str(Path(__file__).parent / "test_models"), + "--model-control-mode", + "explicit", + "--openai-restricted-api", + "model-repository", + "mgmt-key", + "mgmt-secret", + ] + with OpenAIServer(args) as openai_server: + yield openai_server + + @pytest.mark.parametrize( + "headers, expected_success, description", + [ + (None, False, "No auth"), + ({"mgmt-key": "mgmt-secret"}, True, "Valid auth"), + ({"mgmt-key": "wrong-secret"}, False, "Invalid auth value"), + ({"wrong-key": "mgmt-secret"}, False, "Invalid auth key"), + ], + ) + def test_model_repository_management_endpoints_with_auth( + self, + server_with_restrictions_explicit_mode, + headers, + expected_success, + description, + ): + base_url = server_with_restrictions_explicit_mode.url_root + verify_model_repository_management_endpoints( + base_url, "mock_llm", headers, expected_success, description + ) + + @pytest.mark.openai class TestOpenAIServerMultipleRestrictions: """Test cases for OpenAI server with multiple restriction groups.""" From 0e68c9d8f95bb93f3f9b3cb2808da7edeb2e27bd Mon Sep 17 00:00:00 2001 From: Sai Kiran Polisetty Date: Thu, 12 Mar 2026 20:59:48 +0530 Subject: [PATCH 12/24] Update --- python/openai/tests/test_model_management.py | 5 +++-- python/openai/tests/test_openai_restricted_apis.py | 2 +- 2 files changed, 4 insertions(+), 3 deletions(-) diff --git a/python/openai/tests/test_model_management.py b/python/openai/tests/test_model_management.py index b085940744..d470077412 100644 --- a/python/openai/tests/test_model_management.py +++ b/python/openai/tests/test_model_management.py @@ -527,7 +527,9 @@ def test_load_streaming_completions(self, server_with_explicit_mode, model: str) self._unload(base, model) - def test_load_streaming_chat_completions(self, server_with_explicit_mode, model: str): + def test_load_streaming_chat_completions( + self, server_with_explicit_mode, model: str + ): base = server_with_explicit_mode.url_root self._load(base, model) @@ -606,4 +608,3 @@ def test_model_list_after_load_unload(self, server_with_explicit_mode, model: st for m in requests.get(f"{base}/v1/models", timeout=10).json()["data"] ] assert model not in names - diff --git a/python/openai/tests/test_openai_restricted_apis.py b/python/openai/tests/test_openai_restricted_apis.py index f89d63fe38..825270900c 100755 --- a/python/openai/tests/test_openai_restricted_apis.py +++ b/python/openai/tests/test_openai_restricted_apis.py @@ -1,6 +1,6 @@ #!/usr/bin/env python3 -# Copyright 2025, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# Copyright 2025-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. # # Redistribution and use in source and binary forms, with or without # modification, are permitted provided that the following conditions From ec73a43974fdbfe6cf8025c0cff238eba7834a8c Mon Sep 17 00:00:00 2001 From: Sai Kiran Polisetty Date: Fri, 13 Mar 2026 00:01:46 +0530 Subject: [PATCH 13/24] Update --- .../openai_frontend/engine/triton_engine.py | 6 +++++- python/openai/openai_frontend/utils/utils.py | 13 +++++++++++++ python/openai/tests/test_model_management.py | 16 ++++++++++++++++ 3 files changed, 34 insertions(+), 1 deletion(-) diff --git a/python/openai/openai_frontend/engine/triton_engine.py b/python/openai/openai_frontend/engine/triton_engine.py index cfd0126ed1..e54f74931e 100644 --- a/python/openai/openai_frontend/engine/triton_engine.py +++ b/python/openai/openai_frontend/engine/triton_engine.py @@ -94,7 +94,7 @@ Model, ObjectType, ) -from utils.utils import ClientError, ServerError +from utils.utils import ClientError, ServerError, validate_model_name # TODO: Improve type hints @@ -540,6 +540,8 @@ def _get_model_metadata(self) -> Dict[str, TritonModelMetadata]: return model_metadata async def load_model(self, model_name: str) -> Model: + validate_model_name(model_name) + if ( self.server.options.model_control_mode != tritonserver.ModelControlMode.EXPLICIT @@ -574,6 +576,8 @@ def _load_model_sync(self, model_name: str) -> TritonModelMetadata: return self._build_model_metadata(model_name) async def unload_model(self, model_name: str) -> None: + validate_model_name(model_name) + if ( self.server.options.model_control_mode != tritonserver.ModelControlMode.EXPLICIT diff --git a/python/openai/openai_frontend/utils/utils.py b/python/openai/openai_frontend/utils/utils.py index c8fd3609f6..55bf40d37c 100644 --- a/python/openai/openai_frontend/utils/utils.py +++ b/python/openai/openai_frontend/utils/utils.py @@ -24,8 +24,11 @@ # (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE # OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE. +import re from enum import IntEnum +_INVALID_MODEL_NAME_PATTERN = re.compile(r"/|\.\.") + class ServerError(Exception): """Exception raised for server errors.""" @@ -45,3 +48,13 @@ class StatusCode(IntEnum): AUTHORIZATION_ERROR = 401 NOT_FOUND = 404 SERVER_ERROR = 500 + + +def validate_model_name(model_name: str) -> None: + if not model_name or model_name.isspace(): + raise ClientError("Model name must not be empty or whitespace") + if _INVALID_MODEL_NAME_PATTERN.search(model_name): + raise ClientError( + f"Invalid model name '{model_name}': " + "must not contain path traversal characters ('..', '/')" + ) diff --git a/python/openai/tests/test_model_management.py b/python/openai/tests/test_model_management.py index d470077412..02d96cff9d 100644 --- a/python/openai/tests/test_model_management.py +++ b/python/openai/tests/test_model_management.py @@ -141,6 +141,22 @@ def test_load_nonexistent_model(self, client): response = client.post("/v1/models/model_not_in_repo/load") assert response.status_code == 500 + @pytest.mark.parametrize( + "invalid_name", + [ + "..", + "..mock_llm", + "mock_llm..", + "mock..llm", + "..%2f..%2fetc%2fpasswd", + ], + ) + def test_load_and_unload_invalid_model_name(self, client, invalid_name): + for endpoint in ["load", "unload"]: + response = client.post(f"/v1/models/{invalid_name}/{endpoint}") + assert response.status_code == 400 + assert "Invalid model name" in response.json()["detail"] + def test_load_unload_reload(self, client): assert client.post(f"/v1/models/{TEST_MODEL}/load").status_code == 200 assert TEST_MODEL in _get_model_names(client) From 074bd468ae1a912126c29922dd965a259044574b Mon Sep 17 00:00:00 2001 From: Sai Kiran Polisetty Date: Fri, 13 Mar 2026 00:17:55 +0530 Subject: [PATCH 14/24] Update --- python/openai/openai_frontend/utils/utils.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/python/openai/openai_frontend/utils/utils.py b/python/openai/openai_frontend/utils/utils.py index 55bf40d37c..69db5bdfaa 100644 --- a/python/openai/openai_frontend/utils/utils.py +++ b/python/openai/openai_frontend/utils/utils.py @@ -1,4 +1,4 @@ -# Copyright 2025, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# Copyright 2025-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. # # Redistribution and use in source and binary forms, with or without # modification, are permitted provided that the following conditions From dc70f3ca806f563632ce2bb3dcaa914773d8d941 Mon Sep 17 00:00:00 2001 From: Sai Kiran Polisetty Date: Fri, 13 Mar 2026 00:45:20 +0530 Subject: [PATCH 15/24] Update --- python/openai/tests/test_model_management.py | 60 ++++++++++---------- 1 file changed, 31 insertions(+), 29 deletions(-) diff --git a/python/openai/tests/test_model_management.py b/python/openai/tests/test_model_management.py index 02d96cff9d..4e19e49344 100644 --- a/python/openai/tests/test_model_management.py +++ b/python/openai/tests/test_model_management.py @@ -56,8 +56,9 @@ def _assert_model_fields(model_data: dict): assert model_data["owned_by"] == "Triton Inference Server" -# Test Mode enforcement – NONE mode rejects load/unload API calls class TestModelManagementNoneMode: + """Test NONE mode rejects load/unload API calls.""" + @pytest.fixture(scope="class") def client(self): server = setup_server(TEST_MODEL_REPOSITORY) @@ -76,8 +77,9 @@ def test_load_and_unload_rejected_in_none_mode(self, client): ) -# Test load / unload operations – EXPLICIT mode class TestModelManagement: + """Test load/unload operations in EXPLICIT mode.""" + @pytest.fixture(scope="class") def client(self): server = setup_server( @@ -89,11 +91,18 @@ def client(self): yield test_client server.stop() + @pytest.fixture(autouse=True) + def _clean_models(self, client): + """Ensure clean state before and after every test by unloading the models.""" + for name in [TEST_MODEL, TEST_MODEL_2]: + client.post(f"/v1/models/{name}/unload") + yield + for name in [TEST_MODEL, TEST_MODEL_2]: + client.post(f"/v1/models/{name}/unload") + def test_load_model(self, client): - # Initially no models loaded assert _get_model_names(client) == [] - # Load a model – 200, response has correct OpenAI Model fields response = client.post(f"/v1/models/{TEST_MODEL}/load") assert response.status_code == 200 _assert_model_fields(response.json()) @@ -101,15 +110,13 @@ def test_load_model(self, client): assert TEST_MODEL in _get_model_names(client) - # GET /v1/models/{model_name} returns correct info response = client.get(f"/v1/models/{TEST_MODEL}") assert response.status_code == 200 _assert_model_fields(response.json()) + assert client.get("/health/ready").status_code == 200 def test_unload_model(self, client): - if TEST_MODEL not in _get_model_names(client): - response = client.post(f"/v1/models/{TEST_MODEL}/load") - assert response.status_code == 200 + assert client.post(f"/v1/models/{TEST_MODEL}/load").status_code == 200 response = client.post(f"/v1/models/{TEST_MODEL}/unload") assert response.status_code == 200 @@ -118,20 +125,16 @@ def test_unload_model(self, client): assert body["model"] == TEST_MODEL assert TEST_MODEL not in _get_model_names(client) - assert client.get(f"/v1/models/{TEST_MODEL}").status_code == 404 + assert client.get("/health/ready").status_code == 200 def test_load_already_loaded_model(self, client): - client.post(f"/v1/models/{TEST_MODEL}/unload") assert client.post(f"/v1/models/{TEST_MODEL}/load").status_code == 200 - # Second load of the same model should fail response = client.post(f"/v1/models/{TEST_MODEL}/load") assert response.status_code == 400 assert "already loaded" in response.json()["detail"].lower() - client.post(f"/v1/models/{TEST_MODEL}/unload") - def test_unload_unknown_model(self, client): response = client.post("/v1/models/nonexistent_model/unload") assert response.status_code == 400 @@ -140,6 +143,7 @@ def test_unload_unknown_model(self, client): def test_load_nonexistent_model(self, client): response = client.post("/v1/models/model_not_in_repo/load") assert response.status_code == 500 + assert "Unknown model" in response.json()["detail"] @pytest.mark.parametrize( "invalid_name", @@ -167,8 +171,6 @@ def test_load_unload_reload(self, client): assert client.post(f"/v1/models/{TEST_MODEL}/load").status_code == 200 assert TEST_MODEL in _get_model_names(client) - client.post(f"/v1/models/{TEST_MODEL}/unload") - def test_load_multiple_models(self, client): assert client.post(f"/v1/models/{TEST_MODEL}/load").status_code == 200 assert client.post(f"/v1/models/{TEST_MODEL_2}/load").status_code == 200 @@ -177,17 +179,15 @@ def test_load_multiple_models(self, client): assert TEST_MODEL in names assert TEST_MODEL_2 in names - # Unload one; other remains client.post(f"/v1/models/{TEST_MODEL}/unload") names = _get_model_names(client) assert TEST_MODEL not in names assert TEST_MODEL_2 in names - client.post(f"/v1/models/{TEST_MODEL_2}/unload") - -# Test Sequential and concurrent load/unload class TestModelManagementConcurrency: + """Test sequential and concurrent load/unload.""" + @pytest.fixture(scope="class") def client(self): server = setup_server( @@ -199,6 +199,15 @@ def client(self): yield test_client server.stop() + @pytest.fixture(autouse=True) + def _clean_models(self, client): + """Ensures clean state before and after every test.""" + for name in [TEST_MODEL, TEST_MODEL_2]: + client.post(f"/v1/models/{name}/unload") + yield + for name in [TEST_MODEL, TEST_MODEL_2]: + client.post(f"/v1/models/{name}/unload") + def test_rapid_sequential_load_unload(self, client): for _ in range(3): r = client.post(f"/v1/models/{TEST_MODEL}/load") @@ -220,9 +229,6 @@ def test_concurrent_load_different_models(self, client): assert TEST_MODEL in names assert TEST_MODEL_2 in names - client.post(f"/v1/models/{TEST_MODEL}/unload") - client.post(f"/v1/models/{TEST_MODEL_2}/unload") - def test_concurrent_load_same_model(self, client): with concurrent.futures.ThreadPoolExecutor(max_workers=2) as pool: f1 = pool.submit(client.post, f"/v1/models/{TEST_MODEL}/load") @@ -232,11 +238,9 @@ def test_concurrent_load_same_model(self, client): assert codes == [ 200, 400, - ], f"Expected one 200 and one 400 for concurrent loads of the same model, got {codes}" + ], f"Expected one 200 and one 400 for concurrent loads, got {codes}" assert TEST_MODEL in _get_model_names(client) - client.post(f"/v1/models/{TEST_MODEL}/unload") - def test_concurrent_unload_same_model(self, client): client.post(f"/v1/models/{TEST_MODEL}/load") @@ -244,12 +248,11 @@ def test_concurrent_unload_same_model(self, client): f1 = pool.submit(client.post, f"/v1/models/{TEST_MODEL}/unload") f2 = pool.submit(client.post, f"/v1/models/{TEST_MODEL}/unload") - # One succeeds (200), the other fails (400 "unknown model") — order is non-deterministic codes = sorted([f1.result().status_code, f2.result().status_code]) assert codes == [ 200, 400, - ], f"Expected one 200 and one 400 for concurrent unloads of the same model, got {codes}" + ], f"Expected one 200 and one 400 for concurrent unloads, got {codes}" assert TEST_MODEL not in _get_model_names(client) @@ -420,8 +423,7 @@ def server_with_explicit_mode( @pytest.fixture(autouse=True) def ensure_model_unloaded(self, server_with_explicit_mode, model: str): - """Guarantee clean state before and after every test: unload silently - (model may already be unloaded -- that is fine).""" + """Ensure clean state before and after every test by unloading the model.""" requests.post( f"{server_with_explicit_mode.url_root}/v1/models/{model}/unload", timeout=MODEL_MGMT_UNLOAD_TIMEOUT_S, From 647f74e2fdae086262a6505a2c13069d1e835e7f Mon Sep 17 00:00:00 2001 From: Sai Kiran Polisetty Date: Fri, 13 Mar 2026 21:32:27 +0530 Subject: [PATCH 16/24] Update --- python/openai/tests/utils.py | 3 +-- 1 file changed, 1 insertion(+), 2 deletions(-) diff --git a/python/openai/tests/utils.py b/python/openai/tests/utils.py index 6d59e399d1..c73453ee90 100644 --- a/python/openai/tests/utils.py +++ b/python/openai/tests/utils.py @@ -84,12 +84,11 @@ def __init__( self, cli_args: List[str], *, - port: int = 9000, env_dict: Optional[Dict[str, str]] = None, ) -> None: # TODO: Incorporate caller's cli_args passed to this instance instead self.host = "localhost" - self.port = port + self.port = 9000 env = os.environ.copy() if env_dict is not None: From 4bfb5b98026fdaaadc756aabde25f38fc1ba000a Mon Sep 17 00:00:00 2001 From: Sai Kiran Polisetty Date: Thu, 19 Mar 2026 18:02:25 +0530 Subject: [PATCH 17/24] Remove model name validation from the frontend, as the Core handles it. --- .../openai_frontend/engine/triton_engine.py | 5 +- python/openai/openai_frontend/utils/utils.py | 13 ----- python/openai/tests/test_model_management.py | 51 ++++++++++++------- 3 files changed, 35 insertions(+), 34 deletions(-) diff --git a/python/openai/openai_frontend/engine/triton_engine.py b/python/openai/openai_frontend/engine/triton_engine.py index e54f74931e..6c58bc568f 100644 --- a/python/openai/openai_frontend/engine/triton_engine.py +++ b/python/openai/openai_frontend/engine/triton_engine.py @@ -94,7 +94,7 @@ Model, ObjectType, ) -from utils.utils import ClientError, ServerError, validate_model_name +from utils.utils import ClientError, ServerError # TODO: Improve type hints @@ -540,7 +540,6 @@ def _get_model_metadata(self) -> Dict[str, TritonModelMetadata]: return model_metadata async def load_model(self, model_name: str) -> Model: - validate_model_name(model_name) if ( self.server.options.model_control_mode @@ -576,8 +575,6 @@ def _load_model_sync(self, model_name: str) -> TritonModelMetadata: return self._build_model_metadata(model_name) async def unload_model(self, model_name: str) -> None: - validate_model_name(model_name) - if ( self.server.options.model_control_mode != tritonserver.ModelControlMode.EXPLICIT diff --git a/python/openai/openai_frontend/utils/utils.py b/python/openai/openai_frontend/utils/utils.py index 69db5bdfaa..5426584a43 100644 --- a/python/openai/openai_frontend/utils/utils.py +++ b/python/openai/openai_frontend/utils/utils.py @@ -24,11 +24,8 @@ # (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE # OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE. -import re from enum import IntEnum -_INVALID_MODEL_NAME_PATTERN = re.compile(r"/|\.\.") - class ServerError(Exception): """Exception raised for server errors.""" @@ -48,13 +45,3 @@ class StatusCode(IntEnum): AUTHORIZATION_ERROR = 401 NOT_FOUND = 404 SERVER_ERROR = 500 - - -def validate_model_name(model_name: str) -> None: - if not model_name or model_name.isspace(): - raise ClientError("Model name must not be empty or whitespace") - if _INVALID_MODEL_NAME_PATTERN.search(model_name): - raise ClientError( - f"Invalid model name '{model_name}': " - "must not contain path traversal characters ('..', '/')" - ) diff --git a/python/openai/tests/test_model_management.py b/python/openai/tests/test_model_management.py index 4e19e49344..6e9f0afde2 100644 --- a/python/openai/tests/test_model_management.py +++ b/python/openai/tests/test_model_management.py @@ -126,7 +126,6 @@ def test_unload_model(self, client): assert TEST_MODEL not in _get_model_names(client) assert client.get(f"/v1/models/{TEST_MODEL}").status_code == 404 - assert client.get("/health/ready").status_code == 200 def test_load_already_loaded_model(self, client): assert client.post(f"/v1/models/{TEST_MODEL}/load").status_code == 200 @@ -143,23 +142,41 @@ def test_unload_unknown_model(self, client): def test_load_nonexistent_model(self, client): response = client.post("/v1/models/model_not_in_repo/load") assert response.status_code == 500 - assert "Unknown model" in response.json()["detail"] - - @pytest.mark.parametrize( - "invalid_name", - [ + assert "failed to poll from model repository" in response.json()["detail"] + + def test_load_and_unload_invalid_model_name(self, client): + # Names containing '/' or '..' are intercepted by FastAPI/HTTP routing + # before reaching the handler. It return 404 Not Found or 405 Method Not Allowed. + INVALID_TRAVERSAL_NAMES = [ + "../etc", + "../../etc/passwd", + "../../../../etc", + "model/..", "..", - "..mock_llm", - "mock_llm..", - "mock..llm", - "..%2f..%2fetc%2fpasswd", - ], - ) - def test_load_and_unload_invalid_model_name(self, client, invalid_name): - for endpoint in ["load", "unload"]: - response = client.post(f"/v1/models/{invalid_name}/{endpoint}") - assert response.status_code == 400 - assert "Invalid model name" in response.json()["detail"] + "/etc/passwd", + "model/subdir", + "model/", + ] + for model_name in INVALID_TRAVERSAL_NAMES: + for endpoint in ["load", "unload"]: + response = client.post(f"/v1/models/{model_name}/{endpoint}") + assert response.status_code in (404, 405), ( + f"Expected 404 or 405 for model name {model_name!r}, " + f"got {response.status_code}" + ) + + INVALID_WHITESPACE_NAMES = [" ", " ", "\t", "\n", "\r", "\f", "\v", " \t \n "] + for model_name in INVALID_WHITESPACE_NAMES: + for endpoint in ["load", "unload"]: + response = client.post(f"/v1/models/{model_name}/{endpoint}") + assert response.status_code == 500, ( + f"Expected 500 for model name {model_name!r}, " + f"got {response.status_code}" + ) + assert ( + "model name must not contain only whitespace characters" + in response.json()["detail"].lower() + ) def test_load_unload_reload(self, client): assert client.post(f"/v1/models/{TEST_MODEL}/load").status_code == 200 From 88a96581c1a05fbc9db622f24a1b8a26df0460d9 Mon Sep 17 00:00:00 2001 From: Sai Kiran Polisetty Date: Thu, 19 Mar 2026 18:06:00 +0530 Subject: [PATCH 18/24] Fix pre-commit --- python/openai/openai_frontend/engine/triton_engine.py | 1 - 1 file changed, 1 deletion(-) diff --git a/python/openai/openai_frontend/engine/triton_engine.py b/python/openai/openai_frontend/engine/triton_engine.py index 6c58bc568f..cfd0126ed1 100644 --- a/python/openai/openai_frontend/engine/triton_engine.py +++ b/python/openai/openai_frontend/engine/triton_engine.py @@ -540,7 +540,6 @@ def _get_model_metadata(self) -> Dict[str, TritonModelMetadata]: return model_metadata async def load_model(self, model_name: str) -> Model: - if ( self.server.options.model_control_mode != tritonserver.ModelControlMode.EXPLICIT From f9a8cf62bca3e630149f1f3afe30c73cb770bd0b Mon Sep 17 00:00:00 2001 From: Sai Kiran Polisetty Date: Thu, 19 Mar 2026 18:26:35 +0530 Subject: [PATCH 19/24] Update --- python/openai/tests/test_model_management.py | 13 ------------- 1 file changed, 13 deletions(-) diff --git a/python/openai/tests/test_model_management.py b/python/openai/tests/test_model_management.py index 6e9f0afde2..7c9bb1192a 100644 --- a/python/openai/tests/test_model_management.py +++ b/python/openai/tests/test_model_management.py @@ -165,19 +165,6 @@ def test_load_and_unload_invalid_model_name(self, client): f"got {response.status_code}" ) - INVALID_WHITESPACE_NAMES = [" ", " ", "\t", "\n", "\r", "\f", "\v", " \t \n "] - for model_name in INVALID_WHITESPACE_NAMES: - for endpoint in ["load", "unload"]: - response = client.post(f"/v1/models/{model_name}/{endpoint}") - assert response.status_code == 500, ( - f"Expected 500 for model name {model_name!r}, " - f"got {response.status_code}" - ) - assert ( - "model name must not contain only whitespace characters" - in response.json()["detail"].lower() - ) - def test_load_unload_reload(self, client): assert client.post(f"/v1/models/{TEST_MODEL}/load").status_code == 200 assert TEST_MODEL in _get_model_names(client) From dca17facf989c8e8bd1edf1e770a9e9f9cafbf87 Mon Sep 17 00:00:00 2001 From: Yingge He Date: Thu, 19 Mar 2026 17:03:11 -0700 Subject: [PATCH 20/24] Fix trtllm generate_engine.py --- qa/L0_openai/generate_engine.py | 6 ++++-- 1 file changed, 4 insertions(+), 2 deletions(-) diff --git a/qa/L0_openai/generate_engine.py b/qa/L0_openai/generate_engine.py index b454896cfb..207932e3ee 100644 --- a/qa/L0_openai/generate_engine.py +++ b/qa/L0_openai/generate_engine.py @@ -1,4 +1,4 @@ -# Copyright (c) 2025, NVIDIA CORPORATION. All rights reserved. +# Copyright (c) 2025-2026, NVIDIA CORPORATION. All rights reserved. # # Redistribution and use in source and binary forms, with or without # modification, are permitted provided that the following conditions @@ -32,7 +32,9 @@ def generate_model_engine(model: str, engines_path: str): - config = BuildConfig(plugin_config=PluginConfig.from_dict({"_gemm_plugin": "auto"})) + config = BuildConfig( + plugin_config=PluginConfig.from_dict(**{"gemm_plugin": "auto"}) + ) lora_config = LoraConfig( lora_target_modules=["attn_q", "attn_k", "attn_v"], From 02eb203dd1a3a2846e6b1a0ffb08b8889f9a3989 Mon Sep 17 00:00:00 2001 From: Yingge He Date: Thu, 19 Mar 2026 18:13:03 -0700 Subject: [PATCH 21/24] asfas --- qa/L0_openai/generate_engine.py | 4 +--- 1 file changed, 1 insertion(+), 3 deletions(-) diff --git a/qa/L0_openai/generate_engine.py b/qa/L0_openai/generate_engine.py index 207932e3ee..22c21d2209 100644 --- a/qa/L0_openai/generate_engine.py +++ b/qa/L0_openai/generate_engine.py @@ -32,9 +32,7 @@ def generate_model_engine(model: str, engines_path: str): - config = BuildConfig( - plugin_config=PluginConfig.from_dict(**{"gemm_plugin": "auto"}) - ) + config = BuildConfig(plugin_config=PluginConfig(gemm_plugin="auto")) lora_config = LoraConfig( lora_target_modules=["attn_q", "attn_k", "attn_v"], From 55af90c0508aae9a496f4860e8a7e08ad1257ebd Mon Sep 17 00:00:00 2001 From: Yingge He Date: Tue, 24 Mar 2026 05:26:03 -0700 Subject: [PATCH 22/24] Update tests --- .../openai_frontend/engine/triton_engine.py | 7 +- python/openai/openai_frontend/main.py | 4 +- python/openai/openai_frontend/utils/utils.py | 2 +- python/openai/tests/test_model_management.py | 860 +++++++++--------- .../tests/test_openai_restricted_apis.py | 63 +- src/command_line_parser.cc | 6 +- 6 files changed, 442 insertions(+), 500 deletions(-) diff --git a/python/openai/openai_frontend/engine/triton_engine.py b/python/openai/openai_frontend/engine/triton_engine.py index cfd0126ed1..19a5b2f157 100644 --- a/python/openai/openai_frontend/engine/triton_engine.py +++ b/python/openai/openai_frontend/engine/triton_engine.py @@ -138,7 +138,6 @@ def __init__( self.lora_separator = lora_separator self.default_max_tokens = default_max_tokens - self.create_time = int(time.time()) self.model_metadata = self._get_model_metadata() self._metadata_lock = asyncio.Lock() self.tool_call_parser = ( @@ -533,7 +532,7 @@ def _build_model_metadata(self, name: str) -> TritonModelMetadata: ) def _get_model_metadata(self) -> Dict[str, TritonModelMetadata]: - # One tokenizer and creation time shared for all loaded models for now. + # One tokenizer is shared for all loaded models; creation time is per model. model_metadata = {} for name, _ in self.server.models(exclude_not_ready=True).keys(): model_metadata[name] = self._build_model_metadata(name) @@ -557,6 +556,8 @@ async def load_model(self, model_name: str) -> Model: # ready, matching standard Triton server behavior. try: metadata = await asyncio.to_thread(self._load_model_sync, model_name) + except tritonserver.InvalidArgumentError as e: + raise ClientError(f"Failed to load model '{model_name}': {e}") except tritonserver.TritonError as e: raise ServerError(f"Failed to load model '{model_name}': {e}") @@ -590,6 +591,8 @@ async def unload_model(self, model_name: str) -> None: # in-flight request draining and conflict resolution internally. try: await asyncio.to_thread(self._unload_model_sync, model_name) + except tritonserver.InvalidArgumentError as e: + raise ClientError(f"Failed to unload model '{model_name}': {e}") except tritonserver.TritonError as e: raise ServerError(f"Failed to unload model '{model_name}': {e}") diff --git a/python/openai/openai_frontend/main.py b/python/openai/openai_frontend/main.py index c96d154fc2..ac310e662c 100755 --- a/python/openai/openai_frontend/main.py +++ b/python/openai/openai_frontend/main.py @@ -157,8 +157,8 @@ def parse_args(): choices=["none", "explicit"], help="Specify the mode for model management. Options are 'none', and 'explicit'. " "The default is 'none'. For 'none', the server will load all models in the model " - "repository at startup and will not make any changes to the load " - "models after that. For 'explicit', model load and unload is initiated by using the " + "repository at startup and will not make any changes to the loaded " + "models after that. For 'explicit', model load and unload are initiated by using the " "model control APIs, and only models specified with --load-model will " "be loaded at startup.", ) diff --git a/python/openai/openai_frontend/utils/utils.py b/python/openai/openai_frontend/utils/utils.py index 5426584a43..c8fd3609f6 100644 --- a/python/openai/openai_frontend/utils/utils.py +++ b/python/openai/openai_frontend/utils/utils.py @@ -1,4 +1,4 @@ -# Copyright 2025-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# Copyright 2025, NVIDIA CORPORATION & AFFILIATES. All rights reserved. # # Redistribution and use in source and binary forms, with or without # modification, are permitted provided that the following conditions diff --git a/python/openai/tests/test_model_management.py b/python/openai/tests/test_model_management.py index 7c9bb1192a..85fcee0617 100644 --- a/python/openai/tests/test_model_management.py +++ b/python/openai/tests/test_model_management.py @@ -25,255 +25,37 @@ # OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE. import concurrent.futures +import http.client import os +import random +import time from pathlib import Path +from urllib.parse import urlparse import pytest import requests -import tritonserver -from fastapi.testclient import TestClient -from tests.utils import OpenAIServer, setup_fastapi_app, setup_server +from tests.utils import OpenAIServer TEST_MODEL_REPOSITORY = str(Path(__file__).parent / "test_models") TEST_MODEL = "mock_llm" TEST_MODEL_2 = "identity_py" -MODEL_MGMT_LOAD_TIMEOUT_S = int(os.getenv("TEST_MODEL_MANAGEMENT_LOAD_TIMEOUT", "300")) -MODEL_MGMT_UNLOAD_TIMEOUT_S = int( - os.getenv("TEST_MODEL_MANAGEMENT_UNLOAD_TIMEOUT", "120") -) - - -def _get_model_names(client: TestClient) -> list[str]: - response = client.get("/v1/models") - assert response.status_code == 200 - return [m["id"] for m in response.json()["data"]] - - -def _assert_model_fields(model_data: dict): - assert model_data["id"] - assert model_data["object"] == "model" - assert model_data["created"] > 0 - assert model_data["owned_by"] == "Triton Inference Server" - - -class TestModelManagementNoneMode: - """Test NONE mode rejects load/unload API calls.""" - - @pytest.fixture(scope="class") - def client(self): - server = setup_server(TEST_MODEL_REPOSITORY) - app = setup_fastapi_app(tokenizer="", server=server, backend=None) - with TestClient(app) as test_client: - yield test_client - server.stop() - - def test_load_and_unload_rejected_in_none_mode(self, client): - for api in ["load", "unload"]: - response = client.post(f"/v1/models/{TEST_MODEL}/{api}") - assert response.status_code == 400 - assert ( - "model load/unload requires --model-control-mode=explicit" - in response.json()["detail"].lower() - ) - - -class TestModelManagement: - """Test load/unload operations in EXPLICIT mode.""" - - @pytest.fixture(scope="class") - def client(self): - server = setup_server( - TEST_MODEL_REPOSITORY, - model_control_mode=tritonserver.ModelControlMode.EXPLICIT, - ) - app = setup_fastapi_app(tokenizer="", server=server, backend=None) - with TestClient(app) as test_client: - yield test_client - server.stop() - - @pytest.fixture(autouse=True) - def _clean_models(self, client): - """Ensure clean state before and after every test by unloading the models.""" - for name in [TEST_MODEL, TEST_MODEL_2]: - client.post(f"/v1/models/{name}/unload") - yield - for name in [TEST_MODEL, TEST_MODEL_2]: - client.post(f"/v1/models/{name}/unload") - - def test_load_model(self, client): - assert _get_model_names(client) == [] - - response = client.post(f"/v1/models/{TEST_MODEL}/load") - assert response.status_code == 200 - _assert_model_fields(response.json()) - assert response.json()["id"] == TEST_MODEL - - assert TEST_MODEL in _get_model_names(client) - - response = client.get(f"/v1/models/{TEST_MODEL}") - assert response.status_code == 200 - _assert_model_fields(response.json()) - assert client.get("/health/ready").status_code == 200 - - def test_unload_model(self, client): - assert client.post(f"/v1/models/{TEST_MODEL}/load").status_code == 200 - - response = client.post(f"/v1/models/{TEST_MODEL}/unload") - assert response.status_code == 200 - body = response.json() - assert body["status"] == "success" - assert body["model"] == TEST_MODEL - - assert TEST_MODEL not in _get_model_names(client) - assert client.get(f"/v1/models/{TEST_MODEL}").status_code == 404 - - def test_load_already_loaded_model(self, client): - assert client.post(f"/v1/models/{TEST_MODEL}/load").status_code == 200 - - response = client.post(f"/v1/models/{TEST_MODEL}/load") - assert response.status_code == 400 - assert "already loaded" in response.json()["detail"].lower() - - def test_unload_unknown_model(self, client): - response = client.post("/v1/models/nonexistent_model/unload") - assert response.status_code == 400 - assert "unknown model" in response.json()["detail"].lower() - - def test_load_nonexistent_model(self, client): - response = client.post("/v1/models/model_not_in_repo/load") - assert response.status_code == 500 - assert "failed to poll from model repository" in response.json()["detail"] - - def test_load_and_unload_invalid_model_name(self, client): - # Names containing '/' or '..' are intercepted by FastAPI/HTTP routing - # before reaching the handler. It return 404 Not Found or 405 Method Not Allowed. - INVALID_TRAVERSAL_NAMES = [ - "../etc", - "../../etc/passwd", - "../../../../etc", - "model/..", - "..", - "/etc/passwd", - "model/subdir", - "model/", - ] - for model_name in INVALID_TRAVERSAL_NAMES: - for endpoint in ["load", "unload"]: - response = client.post(f"/v1/models/{model_name}/{endpoint}") - assert response.status_code in (404, 405), ( - f"Expected 404 or 405 for model name {model_name!r}, " - f"got {response.status_code}" - ) - - def test_load_unload_reload(self, client): - assert client.post(f"/v1/models/{TEST_MODEL}/load").status_code == 200 - assert TEST_MODEL in _get_model_names(client) - - assert client.post(f"/v1/models/{TEST_MODEL}/unload").status_code == 200 - assert TEST_MODEL not in _get_model_names(client) - - assert client.post(f"/v1/models/{TEST_MODEL}/load").status_code == 200 - assert TEST_MODEL in _get_model_names(client) - - def test_load_multiple_models(self, client): - assert client.post(f"/v1/models/{TEST_MODEL}/load").status_code == 200 - assert client.post(f"/v1/models/{TEST_MODEL_2}/load").status_code == 200 - - names = _get_model_names(client) - assert TEST_MODEL in names - assert TEST_MODEL_2 in names - - client.post(f"/v1/models/{TEST_MODEL}/unload") - names = _get_model_names(client) - assert TEST_MODEL not in names - assert TEST_MODEL_2 in names - - -class TestModelManagementConcurrency: - """Test sequential and concurrent load/unload.""" - - @pytest.fixture(scope="class") - def client(self): - server = setup_server( - TEST_MODEL_REPOSITORY, - model_control_mode=tritonserver.ModelControlMode.EXPLICIT, - ) - app = setup_fastapi_app(tokenizer="", server=server, backend=None) - with TestClient(app) as test_client: - yield test_client - server.stop() - - @pytest.fixture(autouse=True) - def _clean_models(self, client): - """Ensures clean state before and after every test.""" - for name in [TEST_MODEL, TEST_MODEL_2]: - client.post(f"/v1/models/{name}/unload") - yield - for name in [TEST_MODEL, TEST_MODEL_2]: - client.post(f"/v1/models/{name}/unload") - - def test_rapid_sequential_load_unload(self, client): - for _ in range(3): - r = client.post(f"/v1/models/{TEST_MODEL}/load") - assert r.status_code == 200 - assert TEST_MODEL in _get_model_names(client) - - r = client.post(f"/v1/models/{TEST_MODEL}/unload") - assert r.status_code == 200 - assert TEST_MODEL not in _get_model_names(client) - - def test_concurrent_load_different_models(self, client): - with concurrent.futures.ThreadPoolExecutor(max_workers=2) as pool: - f1 = pool.submit(client.post, f"/v1/models/{TEST_MODEL}/load") - f2 = pool.submit(client.post, f"/v1/models/{TEST_MODEL_2}/load") - - assert f1.result().status_code == 200 - assert f2.result().status_code == 200 - names = _get_model_names(client) - assert TEST_MODEL in names - assert TEST_MODEL_2 in names - - def test_concurrent_load_same_model(self, client): - with concurrent.futures.ThreadPoolExecutor(max_workers=2) as pool: - f1 = pool.submit(client.post, f"/v1/models/{TEST_MODEL}/load") - f2 = pool.submit(client.post, f"/v1/models/{TEST_MODEL}/load") - - codes = sorted([f1.result().status_code, f2.result().status_code]) - assert codes == [ - 200, - 400, - ], f"Expected one 200 and one 400 for concurrent loads, got {codes}" - assert TEST_MODEL in _get_model_names(client) - - def test_concurrent_unload_same_model(self, client): - client.post(f"/v1/models/{TEST_MODEL}/load") - - with concurrent.futures.ThreadPoolExecutor(max_workers=2) as pool: - f1 = pool.submit(client.post, f"/v1/models/{TEST_MODEL}/unload") - f2 = pool.submit(client.post, f"/v1/models/{TEST_MODEL}/unload") - - codes = sorted([f1.result().status_code, f2.result().status_code]) - assert codes == [ - 200, - 400, - ], f"Expected one 200 and one 400 for concurrent unloads, got {codes}" - assert TEST_MODEL not in _get_model_names(client) # Test "--load-model" and "--model-control-mode" CLI options @pytest.mark.openai -class TestModelManagementCLIOptions: - def _assert_server_startup_fails(self, args, expected_error: str): +class TestModelControlCLIOptions: + @staticmethod + def _assert_server_launch_fails(args, expected_error: str): """Helper: verify server fails to start and stderr contains expected_error.""" with pytest.raises(Exception) as exc_info: with OpenAIServer(args): pass assert expected_error in str(exc_info.value) - def test_load_model_without_explicit_mode_is_error(self): + def test_non_explicit_mode_load_model_error(self): """--load-model without --model-control-mode=explicit must exit with error. Error message matches native tritonserver exactly.""" - self._assert_server_startup_fails( + self._assert_server_launch_fails( args=[ "--model-repository", TEST_MODEL_REPOSITORY, @@ -283,7 +65,7 @@ def test_load_model_without_explicit_mode_is_error(self): expected_error="Error: Use of '--load-model' requires setting '--model-control-mode=explicit' as well.", ) - def test_explicit_no_load_model_starts_with_no_models(self): + def test_explicit_mode_load_zero_model(self): args = [ "--model-repository", TEST_MODEL_REPOSITORY, @@ -295,7 +77,7 @@ def test_explicit_no_load_model_starts_with_no_models(self): assert r.status_code == 200 assert len(r.json()["data"]) == 0 - def test_explicit_single_load_model_at_startup(self): + def test_explicit_mode_load_one_model(self): args = [ "--model-repository", TEST_MODEL_REPOSITORY, @@ -310,7 +92,7 @@ def test_explicit_single_load_model_at_startup(self): assert TEST_MODEL in names assert TEST_MODEL_2 not in names - def test_explicit_multiple_load_model_at_startup(self): + def test_explicit_mode_load_multiple_models(self): args = [ "--model-repository", TEST_MODEL_REPOSITORY, @@ -327,7 +109,7 @@ def test_explicit_multiple_load_model_at_startup(self): assert TEST_MODEL in names assert TEST_MODEL_2 in names - def test_explicit_load_model_wildcard_loads_all(self): + def test_explicit_mode_load_all_models(self): args = [ "--model-repository", TEST_MODEL_REPOSITORY, @@ -342,8 +124,8 @@ def test_explicit_load_model_wildcard_loads_all(self): assert TEST_MODEL in names assert TEST_MODEL_2 in names - def test_explicit_load_model_wildcard_with_other_is_error(self): - self._assert_server_startup_fails( + def test_explicit_mode_load_all_models_and_specific_model_error(self): + self._assert_server_launch_fails( args=[ "--model-repository", TEST_MODEL_REPOSITORY, @@ -357,107 +139,416 @@ def test_explicit_load_model_wildcard_with_other_is_error(self): expected_error="Wildcard model name '*' must be the ONLY startup model if specified at all.", ) - def test_explicit_dynamic_load_unload_via_api(self): - args = [ - "--model-repository", - TEST_MODEL_REPOSITORY, - "--model-control-mode", - "explicit", - ] - with OpenAIServer(args) as server: - base = server.url_root + def test_explicit_mode_load_nonexistent_model_error(self): + self._assert_server_launch_fails( + args=[ + "--model-repository", + TEST_MODEL_REPOSITORY, + "--model-control-mode", + "explicit", + "--load-model", + "nonexistent_model", + ], + expected_error="failed to poll model 'nonexistent_model': model not found in any model repository", + ) - assert ( - requests.post( - f"{base}/v1/models/{TEST_MODEL}/load", timeout=30 - ).status_code - == 200 - ) - assert ( - requests.post( - f"{base}/v1/models/{TEST_MODEL_2}/load", timeout=30 - ).status_code - == 200 + def test_explicit_mode_load_invalid_model_name_error(self): + invalid_model_names = [ + ( + os.path.relpath("/etc", TEST_MODEL_REPOSITORY), + "at least one version must be available under the version policy", + ), + ( + os.path.relpath("/etc/passwd", TEST_MODEL_REPOSITORY), + "Poll failed for model directory", + ), + ("model/..", "model not found in any model repository"), + ("..", "at least one version must be available under the version policy"), + ("/etc/passwd", "model not found in any model repository"), + ("", "Invalid model name"), + (" ", "model not found in any model repository"), + ("\n\t", "model not found in any model repository"), + ] + for model_name, expected_error in invalid_model_names: + self._assert_server_launch_fails( + args=[ + "--model-repository", + TEST_MODEL_REPOSITORY, + "--model-control-mode", + "explicit", + "--load-model", + model_name, + ], + expected_error=expected_error, ) - names = [ - m["id"] - for m in requests.get(f"{base}/v1/models", timeout=10).json()["data"] - ] - assert TEST_MODEL in names - assert TEST_MODEL_2 in names - assert ( - requests.post( - f"{base}/v1/models/{TEST_MODEL}/unload", timeout=30 - ).status_code - == 200 - ) - - names = [ - m["id"] - for m in requests.get(f"{base}/v1/models", timeout=10).json()["data"] - ] - assert TEST_MODEL not in names - assert TEST_MODEL_2 in names +class _ModelManagementBase: + @pytest.fixture(scope="class") + def base_url(self, server: OpenAIServer): + return server.url_root + @pytest.fixture(scope="class") + def all_models(self) -> list[str]: + return [TEST_MODEL, TEST_MODEL_2] -# Test Inference with real LLM backend (vLLM / TRT-LLM) after load/unload -@pytest.mark.openai -class TestModelManagementInference: @pytest.fixture(scope="class") - def server_with_explicit_mode( + def server( self, model_repository: str, tokenizer_model: str, backend: str, + model_control_mode: str, ): args = [ "--model-repository", model_repository, - "--tokenizer", - tokenizer_model, - "--backend", - backend, - "--model-control-mode", - "explicit", ] + if tokenizer_model: + args += ["--tokenizer", tokenizer_model] + if backend: + args += ["--backend", backend] + if model_control_mode: + args += ["--model-control-mode", model_control_mode] + with OpenAIServer(args) as openai_server: yield openai_server @pytest.fixture(autouse=True) - def ensure_model_unloaded(self, server_with_explicit_mode, model: str): - """Ensure clean state before and after every test by unloading the model.""" - requests.post( - f"{server_with_explicit_mode.url_root}/v1/models/{model}/unload", - timeout=MODEL_MGMT_UNLOAD_TIMEOUT_S, - ) + def _cleanup(self, base_url, all_models: list[str]): + """Ensure clean state before and after every test by unloading the models.""" + for name in all_models: + requests.post(f"{base_url}/v1/models/{name}/unload") yield - requests.post( - f"{server_with_explicit_mode.url_root}/v1/models/{model}/unload", - timeout=MODEL_MGMT_UNLOAD_TIMEOUT_S, - ) + for name in all_models: + requests.post(f"{base_url}/v1/models/{name}/unload") @staticmethod - def _load(base_url, model_name): - r = requests.post( - f"{base_url}/v1/models/{model_name}/load", - timeout=MODEL_MGMT_LOAD_TIMEOUT_S, - ) - assert r.status_code == 200, f"Load failed: {r.text}" + def _list_available_models(base_url: str) -> list[str]: + response = requests.get(f"{base_url}/v1/models") + assert response.status_code == 200 + return [m["id"] for m in response.json()["data"]] @staticmethod - def _unload(base_url, model_name): - r = requests.post( - f"{base_url}/v1/models/{model_name}/unload", - timeout=MODEL_MGMT_UNLOAD_TIMEOUT_S, - ) - assert r.status_code == 200, f"Unload failed: {r.text}" + def _assert_unknown_model(response: requests.Response): + assert response.status_code == 400 + assert "unknown model" in response.json()["detail"].lower() + + +class TestModelControlModeNone(_ModelManagementBase): + """Test NONE mode rejects load/unload API calls.""" + + @pytest.fixture(scope="class") + def model_repository(self): + return TEST_MODEL_REPOSITORY + + @pytest.mark.parametrize("model_control_mode", [None, "none"], indirect=True) + def test_load_and_unload_rejected(self, base_url): + """Test NONE mode rejects load/unload API calls.""" + + for api in ["load", "unload"]: + response = requests.post(f"{base_url}/v1/models/{TEST_MODEL}/{api}") + assert response.status_code == 400 + assert ( + "model load/unload requires --model-control-mode=explicit" + in response.json()["detail"].lower() + ) + + +class TestModelManagement(_ModelManagementBase): + """Test load/unload operations in EXPLICIT mode.""" + + @pytest.fixture(scope="class") + def model_control_mode(self): + return "explicit" + + @pytest.fixture(scope="class") + def model_repository(self): + return TEST_MODEL_REPOSITORY @staticmethod - def _assert_unknown_model(response): + def _assert_model_metadata(model_data: dict): + assert model_data["id"] + assert model_data["object"] == "model" + assert model_data["created"] > 0 + assert model_data["owned_by"] == "Triton Inference Server" + + def test_load_model(self, base_url): + # Server should start with no models loaded + assert self._list_available_models(base_url) == [] + + response = requests.post(f"{base_url}/v1/models/{TEST_MODEL}/load") + assert response.status_code == 200 + self._assert_model_metadata(response.json()) + assert response.json()["id"] == TEST_MODEL + assert TEST_MODEL in self._list_available_models(base_url) + + response = requests.get(f"{base_url}/v1/models/{TEST_MODEL}") + assert response.status_code == 200 + self._assert_model_metadata(response.json()) + + def test_unload_model(self, base_url): + assert ( + requests.post(f"{base_url}/v1/models/{TEST_MODEL}/load").status_code == 200 + ) + + response = requests.post(f"{base_url}/v1/models/{TEST_MODEL}/unload") + assert response.status_code == 200 + body = response.json() + assert body["status"] == "success" + assert body["model"] == TEST_MODEL + + assert TEST_MODEL not in self._list_available_models(base_url) + assert requests.get(f"{base_url}/v1/models/{TEST_MODEL}").status_code == 404 + + def test_load_rejects_duplicate(self, base_url): + assert ( + requests.post(f"{base_url}/v1/models/{TEST_MODEL}/load").status_code == 200 + ) + + response = requests.post(f"{base_url}/v1/models/{TEST_MODEL}/load") assert response.status_code == 400 - assert "unknown model" in response.json()["detail"].lower() + assert "already loaded" in response.json()["detail"].lower() + + def test_load_unload_unknown_model(self, base_url): + response = requests.post(f"{base_url}/v1/models/unknown_model/load") + assert response.status_code == 500 + assert "failed to poll from model repository" in response.json()["detail"] + + response = requests.post(f"{base_url}/v1/models/unknown_model/unload") + self._assert_unknown_model(response) + + def test_load_unload_invalid_model_name(self, base_url): + invalid_model_names = [ + (os.path.relpath("/etc", TEST_MODEL_REPOSITORY), 404), + (os.path.relpath("/etc/passwd", TEST_MODEL_REPOSITORY), 404), + ("model/..", 404), + ("..", 400), + ("/etc/passwd", 404), + ("model/subdir", 404), + ("model/", 404), + ("", 404), + ("%20%20", 400), + ("%0A%09", 400), + ] + parsed = urlparse(base_url) + for model_name, expected_status in invalid_model_names: + for endpoint in ["load", "unload"]: + conn = http.client.HTTPConnection(parsed.hostname, parsed.port) + conn.request("POST", f"/v1/models/{model_name}/{endpoint}") + response = conn.getresponse() + + assert response.status == expected_status, ( + f"Expected {expected_status} for model name {model_name!r}, " + f"got {response.status} {response.read().decode()}" + ) + + def test_load_unload_reload(self, base_url): + assert ( + requests.post(f"{base_url}/v1/models/{TEST_MODEL}/load").status_code == 200 + ) + assert TEST_MODEL in self._list_available_models(base_url) + + assert ( + requests.post(f"{base_url}/v1/models/{TEST_MODEL}/unload").status_code + == 200 + ) + assert TEST_MODEL not in self._list_available_models(base_url) + + assert ( + requests.post(f"{base_url}/v1/models/{TEST_MODEL}/load").status_code == 200 + ) + assert TEST_MODEL in self._list_available_models(base_url) + + def test_load_multiple_models(self, base_url): + assert ( + requests.post(f"{base_url}/v1/models/{TEST_MODEL}/load").status_code == 200 + ) + assert ( + requests.post(f"{base_url}/v1/models/{TEST_MODEL_2}/load").status_code + == 200 + ) + + names = self._list_available_models(base_url) + assert TEST_MODEL in names + assert TEST_MODEL_2 in names + + # Unload the first model + assert ( + requests.post(f"{base_url}/v1/models/{TEST_MODEL}/unload").status_code + == 200 + ) + names = self._list_available_models(base_url) + assert TEST_MODEL not in names + assert TEST_MODEL_2 in names + + # Unload the second model + assert ( + requests.post(f"{base_url}/v1/models/{TEST_MODEL_2}/unload").status_code + == 200 + ) + names = self._list_available_models(base_url) + assert TEST_MODEL not in names + assert TEST_MODEL_2 not in names + + +class TestConcurrentModelManagement(_ModelManagementBase): + """Test sequential and concurrent load/unload.""" + + @pytest.fixture(scope="class") + def model_repository(self): + return TEST_MODEL_REPOSITORY + + @pytest.fixture(scope="class") + def model_control_mode(self): + return "explicit" + + def test_concurrent_load_model(self, base_url): + futures = [] + concurrency = 5 + with concurrent.futures.ThreadPoolExecutor(max_workers=concurrency) as pool: + for _ in range(concurrency): + futures.append( + pool.submit( + requests.post, f"{base_url}/v1/models/{TEST_MODEL}/load" + ) + ) + + codes = sorted([future.result().status_code for future in futures]) + assert codes == [200] + [400] * ( + concurrency - 1 + ), f"Expected one 200 and the rest 400 for concurrent loads, got {codes}" + assert TEST_MODEL in self._list_available_models(base_url) + + def test_concurrent_unload_model(self, base_url): + # Load the model + requests.post(f"{base_url}/v1/models/{TEST_MODEL}/load") + assert TEST_MODEL in self._list_available_models(base_url) + + futures = [] + concurrency = 5 + with concurrent.futures.ThreadPoolExecutor(max_workers=concurrency) as pool: + for _ in range(concurrency): + futures.append( + pool.submit( + requests.post, f"{base_url}/v1/models/{TEST_MODEL}/unload" + ) + ) + + codes = sorted([future.result().status_code for future in futures]) + assert codes == [200] + [400] * ( + concurrency - 1 + ), f"Expected one 200 and the rest 400 for concurrent unloads, got {codes}" + assert TEST_MODEL not in self._list_available_models(base_url) + + def test_concurrent_load_multiple_models(self, base_url): + concurrency = 10 + tasks = [] + with concurrent.futures.ThreadPoolExecutor(max_workers=concurrency) as pool: + for _ in range(concurrency // 2): + for model_name in [TEST_MODEL, TEST_MODEL_2]: + tasks.append( + pool.submit( + requests.post, f"{base_url}/v1/models/{model_name}/load" + ) + ) + + codes = sorted([task.result().status_code for task in tasks]) + assert codes == [200, 200] + [400] * ( + concurrency - 2 + ), f"Expected two 200s and the rest 400 for concurrent loads, got {codes}" + available_models = self._list_available_models(base_url) + assert TEST_MODEL in available_models + assert TEST_MODEL_2 in available_models + + def test_concurrent_unload_multiple_models(self, base_url): + # Load both models + for model_name in [TEST_MODEL, TEST_MODEL_2]: + requests.post(f"{base_url}/v1/models/{model_name}/load") + assert model_name in self._list_available_models(base_url) + + concurrency = 10 + tasks = [] + with concurrent.futures.ThreadPoolExecutor(max_workers=concurrency) as pool: + for _ in range(concurrency // 2): + for model_name in [TEST_MODEL, TEST_MODEL_2]: + tasks.append( + pool.submit( + requests.post, f"{base_url}/v1/models/{model_name}/unload" + ) + ) + + codes = sorted([task.result().status_code for task in tasks]) + assert codes == [200, 200] + [400] * ( + concurrency - 2 + ), f"Expected two 200s and the rest 400 for concurrent unloads, got {codes}" + available_models = self._list_available_models(base_url) + for model_name in [TEST_MODEL, TEST_MODEL_2]: + assert model_name not in available_models + + def test_concurrent_load_unload_stress(self, base_url): + futures = [] + concurrency = 50 + with concurrent.futures.ThreadPoolExecutor(max_workers=concurrency) as pool: + for _ in range(concurrency): + action = random.choice(["load", "unload"]) + model_name = random.choice([TEST_MODEL, TEST_MODEL_2]) + futures.append( + pool.submit( + requests.post, f"{base_url}/v1/models/{model_name}/{action}" + ) + ) + done, _ = concurrent.futures.wait( + futures, return_when=concurrent.futures.ALL_COMPLETED + ) + assert ( + len(done) == concurrency + ), f"Expected {concurrency} requests to be completed, got {len(done)}" + for future in done: + response = future.result() + assert response.status_code in ( + 200, + 400, + ), f"Unexpected status code: {response.status_code}" + + # Wait for server to be ready + for _ in range(10): + response = requests.get(f"{base_url}/health/ready") + if response.status_code == 200: + break + time.sleep(1) + assert response.status_code == 200 + + +# Test Inference with real LLM backend (vLLM / TRT-LLM) after load/unload +@pytest.mark.openai +class TestModelManagementInference(_ModelManagementBase): + @pytest.fixture(scope="class") + def model_control_mode(self): + return "explicit" + + @pytest.fixture(scope="class") + def all_models(self, model: str) -> list[str]: + if model == "tensorrt_llm_bls": + return ["postprocessing", "preprocessing", "tensorrt_llm", model] + else: + return [model] + + @staticmethod + def _assert_load(base_url, model_list: list[str]): + for model_name in model_list: + r = requests.post( + f"{base_url}/v1/models/{model_name}/load", + ) + assert r.status_code == 200, f"Load {model_name} failed: {r.text}" + + @staticmethod + def _assert_unload(base_url, model_list: list[str]): + for model_name in model_list: + r = requests.post( + f"{base_url}/v1/models/{model_name}/unload", + ) + assert r.status_code == 200, f"Unload {model_name} failed: {r.text}" @staticmethod def _assert_usage(usage): @@ -469,7 +560,7 @@ def _assert_usage(usage): ) @staticmethod - def _completions(base_url, model_name, **kwargs): + def _completions(base_url, model_name: str, **kwargs): return requests.post( f"{base_url}/v1/completions", json={ @@ -478,155 +569,38 @@ def _completions(base_url, model_name, **kwargs): "max_tokens": 10, **kwargs, }, - timeout=30, ) - @staticmethod - def _chat_completions(base_url, model_name, **kwargs): - return requests.post( - f"{base_url}/v1/chat/completions", - json={ - "model": model_name, - "messages": [{"role": "user", "content": "What is machine learning?"}], - "max_tokens": 10, - **kwargs, - }, - timeout=30, - ) + def test_load_completions(self, base_url, all_models: list[str], model: str): + self._assert_unknown_model(self._completions(base_url, model)) - def test_load_completions(self, server_with_explicit_mode, model: str): - base = server_with_explicit_mode.url_root - self._assert_unknown_model(self._completions(base, model)) - - self._load(base, model) - r = self._completions(base, model) + self._assert_load(base_url, all_models) + r = self._completions(base_url, model) assert r.status_code == 200 data = r.json() assert data["choices"][0]["text"].strip() assert data["choices"][0]["finish_reason"] == "stop" self._assert_usage(data["usage"]) - self._unload(base, model) - - def test_load_chat_completions(self, server_with_explicit_mode, model: str): - base = server_with_explicit_mode.url_root - self._assert_unknown_model(self._chat_completions(base, model)) + self._assert_unload(base_url, all_models) - self._load(base, model) - r = self._chat_completions(base, model) - assert r.status_code == 200 - data = r.json() - msg = data["choices"][0]["message"] - assert msg["content"].strip() - assert msg["role"] == "assistant" - assert data["choices"][0]["finish_reason"] == "stop" - self._assert_usage(data["usage"]) - - self._unload(base, model) - - def test_load_streaming_completions(self, server_with_explicit_mode, model: str): - base = server_with_explicit_mode.url_root - self._load(base, model) - - r = requests.post( - f"{base}/v1/completions", - json={ - "model": model, - "prompt": "What is machine learning?", - "max_tokens": 10, - "stream": True, - }, - stream=True, - timeout=30, - ) - assert r.status_code == 200 - chunks = [ - line.removeprefix("data: ") - for line in r.iter_lines(decode_unicode=True) - if line.startswith("data: ") and line.strip() != "data: [DONE]" - ] - assert len(chunks) > 0 - - self._unload(base, model) - - def test_load_streaming_chat_completions( - self, server_with_explicit_mode, model: str + def test_unload_rejects_inference( + self, base_url, all_models: list[str], model: str ): - base = server_with_explicit_mode.url_root - self._load(base, model) - - r = requests.post( - f"{base}/v1/chat/completions", - json={ - "model": model, - "messages": [{"role": "user", "content": "What is machine learning?"}], - "max_tokens": 10, - "stream": True, - }, - stream=True, - timeout=30, - ) - assert r.status_code == 200 - chunks = [ - line.removeprefix("data: ") - for line in r.iter_lines(decode_unicode=True) - if line.startswith("data: ") and line.strip() != "data: [DONE]" - ] - assert len(chunks) > 0 - - self._unload(base, model) - - def test_unload_rejects_inference(self, server_with_explicit_mode, model: str): - base = server_with_explicit_mode.url_root - self._load(base, model) - - assert self._completions(base, model).status_code == 200 - assert self._chat_completions(base, model).status_code == 200 + self._assert_load(base_url, all_models) + assert self._completions(base_url, model).status_code == 200 - self._unload(base, model) + self._assert_unload(base_url, all_models) + self._assert_unknown_model(self._completions(base_url, model)) - self._assert_unknown_model(self._completions(base, model)) - self._assert_unknown_model(self._chat_completions(base, model)) + def test_reload_inference(self, base_url, all_models: list[str], model: str): + self._assert_load(base_url, all_models) + assert self._completions(base_url, model).status_code == 200 - def test_reload_inference(self, server_with_explicit_mode, model: str): - base = server_with_explicit_mode.url_root + self._assert_unload(base_url, all_models) + self._assert_unknown_model(self._completions(base_url, model)) - self._load(base, model) - assert self._completions(base, model).status_code == 200 - assert self._chat_completions(base, model).status_code == 200 - - self._unload(base, model) - self._assert_unknown_model(self._completions(base, model)) - - self._load(base, model) - r = self._completions(base, model) + self._assert_load(base_url, all_models) + r = self._completions(base_url, model) assert r.status_code == 200 assert r.json()["choices"][0]["text"].strip() - - r = self._chat_completions(base, model) - assert r.status_code == 200 - assert r.json()["choices"][0]["message"]["content"].strip() - - self._unload(base, model) - - def test_model_list_after_load_unload(self, server_with_explicit_mode, model: str): - base = server_with_explicit_mode.url_root - - # No models initially - r = requests.get(f"{base}/v1/models", timeout=10) - assert r.status_code == 200 - assert len(r.json()["data"]) == 0 - - self._load(base, model) - names = [ - m["id"] - for m in requests.get(f"{base}/v1/models", timeout=10).json()["data"] - ] - assert model in names - - self._unload(base, model) - names = [ - m["id"] - for m in requests.get(f"{base}/v1/models", timeout=10).json()["data"] - ] - assert model not in names diff --git a/python/openai/tests/test_openai_restricted_apis.py b/python/openai/tests/test_openai_restricted_apis.py index 825270900c..200412df34 100755 --- a/python/openai/tests/test_openai_restricted_apis.py +++ b/python/openai/tests/test_openai_restricted_apis.py @@ -40,7 +40,7 @@ def assert_response_success( """Assert that a response was successful.""" assert ( response.status_code == expected_status - ), f"{description} should return {expected_status}, got {response.status_code}" + ), f"{description} should return {expected_status}, got {response.status_code} {response.text}" def assert_response_unauthorized( @@ -49,7 +49,7 @@ def assert_response_unauthorized( """Assert that a response was unauthorized.""" assert ( response.status_code == expected_status - ), f"{description} should be unauthorized with {expected_status}, got {response.status_code}" + ), f"{description} should be unauthorized with {expected_status}, got {response.status_code} {response.text}" def make_get_request( @@ -129,6 +129,7 @@ def make_completion_request( def verify_model_repository_endpoints( base_url, model, headers, expected_success, description_prefix ): + # Verify model repository endpoints response = make_get_request(base_url, "/v1/models", headers=headers) if expected_success: assert_response_success( @@ -149,13 +150,10 @@ def verify_model_repository_endpoints( response, description=f"{description_prefix} Specific model endpoint" ) - -def verify_model_repository_management_endpoints( - base_url, model, headers, expected_success, description_prefix -): - for endpoint in ["load", "unload"]: + # Verify model management endpoints + for endpoint in ["unload", "load"]: response = requests.post( - f"{base_url}/v1/models/{model}/{endpoint}", headers=headers, timeout=120 + f"{base_url}/v1/models/{model}/{endpoint}", headers=headers ) if expected_success: assert_response_success( @@ -307,6 +305,10 @@ def server_with_restrictions(self, model_repository, tokenizer_model, backend): tokenizer_model, "--backend", backend, + "--model-control-mode", + "explicit", + "--load-model", + "*", "--openai-restricted-api", "inference,model-repository", "admin-key", @@ -350,47 +352,6 @@ def test_unrestricted_endpoints(self, server_with_restrictions): ) -@pytest.mark.openai -class TestOpenAIServerModelRepositoryRestriction: - """Model-repository restriction for dynamic model load/unload.""" - - @pytest.fixture(scope="class") - def server_with_restrictions_explicit_mode(self): - args = [ - "--model-repository", - str(Path(__file__).parent / "test_models"), - "--model-control-mode", - "explicit", - "--openai-restricted-api", - "model-repository", - "mgmt-key", - "mgmt-secret", - ] - with OpenAIServer(args) as openai_server: - yield openai_server - - @pytest.mark.parametrize( - "headers, expected_success, description", - [ - (None, False, "No auth"), - ({"mgmt-key": "mgmt-secret"}, True, "Valid auth"), - ({"mgmt-key": "wrong-secret"}, False, "Invalid auth value"), - ({"wrong-key": "mgmt-secret"}, False, "Invalid auth key"), - ], - ) - def test_model_repository_management_endpoints_with_auth( - self, - server_with_restrictions_explicit_mode, - headers, - expected_success, - description, - ): - base_url = server_with_restrictions_explicit_mode.url_root - verify_model_repository_management_endpoints( - base_url, "mock_llm", headers, expected_success, description - ) - - @pytest.mark.openai class TestOpenAIServerMultipleRestrictions: """Test cases for OpenAI server with multiple restriction groups.""" @@ -405,6 +366,10 @@ def server_multiple_restrictions(self, model_repository, tokenizer_model, backen tokenizer_model, "--backend", backend, + "--model-control-mode", + "explicit", + "--load-model", + "*", "--openai-restricted-api", "model-repository", "model-key", diff --git a/src/command_line_parser.cc b/src/command_line_parser.cc index cc11803ae8..359a02ea6e 100644 --- a/src/command_line_parser.cc +++ b/src/command_line_parser.cc @@ -1,4 +1,4 @@ -// Copyright 2022-2025, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +// Copyright 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. // // Redistribution and use in source and binary forms, with or without // modification, are permitted provided that the following conditions @@ -424,11 +424,11 @@ TritonParser::SetupOptions() "Specify the mode for model management. Options are \"none\", \"poll\" " "and \"explicit\". The default is \"none\". " "For \"none\", the server will load all models in the model " - "repository(s) at startup and will not make any changes to the load " + "repository(s) at startup and will not make any changes to the loaded " "models after that. For \"poll\", the server will poll the model " "repository(s) to detect changes and will load/unload models based on " "those changes. The poll rate is controlled by 'repository-poll-secs'. " - "For \"explicit\", model load and unload is initiated by using the " + "For \"explicit\", model load and unload are initiated by using the " "model control APIs, and only models specified with --load-model will " "be loaded at startup."}); model_repo_options_.push_back( From ea5273f4b645d9f5b8d9b27cb6040e9bb0b6b441 Mon Sep 17 00:00:00 2001 From: Yingge He Date: Tue, 24 Mar 2026 06:24:59 -0700 Subject: [PATCH 23/24] Update docs --- python/openai/README.md | 48 ++++---------------- python/openai/tests/test_model_management.py | 8 +++- 2 files changed, 15 insertions(+), 41 deletions(-) diff --git a/python/openai/README.md b/python/openai/README.md index 103a8bcfb0..538d086d73 100644 --- a/python/openai/README.md +++ b/python/openai/README.md @@ -564,12 +564,15 @@ runtime. The default is `none`. > See [Triton Model Management](../../docs/user_guide/model_management.md) for > more details. -### Loading Models at Startup (Explicit Mode) +### Load Models at Startup (Explicit Mode) When using `--model-control-mode=explicit`, use `--load-model` to specify which models should be loaded at startup. It may be specified multiple times to load multiple models. +
+Example + ```bash # Start in explicit mode with no models loaded python3 openai_frontend/main.py \ @@ -600,15 +603,16 @@ python3 openai_frontend/main.py \ --load-model '*' ``` +
+ > [!IMPORTANT] -> `--load-model` requires `--model-control-mode=explicit`. Using `--load-model` -> without it is an error. `--load-model=*` must be the **only** `--load-model` -> argument and combining it with a named model is an error. +> - `--load-model` requires `--model-control-mode=explicit`. +> - `--load-model=*` can not be used together with loading specific models. ### Dynamic Load / Unload API Once the server is running in `explicit` mode, models can be loaded and unloaded -at runtime via these endpoints: +at runtime via the following endpoints: | Method | Endpoint | Description | |--------|----------|-------------| @@ -657,40 +661,6 @@ curl -s -X POST http://localhost:9000/v1/models/${MODEL}/unload | jq -#### Error cases - -| Scenario | HTTP Status | Detail | -|----------|-------------|--------| -| `--model-control-mode` is not `explicit` | 400 | `Model load/unload requires --model-control-mode=explicit` | -| Model already loaded (duplicate load) | 400 | `Model '' is already loaded` | -| Model not loaded (unload unknown) | 400 | `Unknown model: ` | -| Model not found in repository | 500 | Triton error: `failed to poll from model repository` | - -### Restricting the Model Management APIs - -The model management endpoints can be protected with authentication headers using -`--openai-restricted-api model-repository`. - -```bash -python3 openai_frontend/main.py \ - --model-repository /path/to/models \ - --tokenizer meta-llama/Meta-Llama-3.1-8B-Instruct \ - --model-control-mode explicit \ - --load-model llama-3.1-8b-instruct \ - --openai-restricted-api model-repository admin-key admin-secret -``` - -Clients must then include the header for model management requests such as -load/unload: - -```bash -curl -H "admin-key: admin-secret" \ - -X POST http://localhost:9000/v1/models/llama-3.1-8b-instruct/load -``` - -See [Limit Endpoint Access](#limit-endpoint-access) for more details on -restricting API groups. - ## Model Parallelism Support - [x] vLLM ([EngineArgs](https://github.com/triton-inference-server/vllm_backend/blob/main/README.md#using-the-vllm-backend)) diff --git a/python/openai/tests/test_model_management.py b/python/openai/tests/test_model_management.py index 85fcee0617..c20dd7d4e4 100644 --- a/python/openai/tests/test_model_management.py +++ b/python/openai/tests/test_model_management.py @@ -188,7 +188,7 @@ class _ModelManagementBase: def base_url(self, server: OpenAIServer): return server.url_root - @pytest.fixture(scope="class") + @pytest.fixture def all_models(self) -> list[str]: return [TEST_MODEL, TEST_MODEL_2] @@ -242,6 +242,10 @@ class TestModelControlModeNone(_ModelManagementBase): def model_repository(self): return TEST_MODEL_REPOSITORY + @pytest.fixture(scope="class") + def model_control_mode(self, request): + return request.param + @pytest.mark.parametrize("model_control_mode", [None, "none"], indirect=True) def test_load_and_unload_rejected(self, base_url): """Test NONE mode rejects load/unload API calls.""" @@ -527,7 +531,7 @@ class TestModelManagementInference(_ModelManagementBase): def model_control_mode(self): return "explicit" - @pytest.fixture(scope="class") + @pytest.fixture def all_models(self, model: str) -> list[str]: if model == "tensorrt_llm_bls": return ["postprocessing", "preprocessing", "tensorrt_llm", model] From 9db96f12bc4fbbda3bbe5a9b6c054907542c9b91 Mon Sep 17 00:00:00 2001 From: Yingge He Date: Tue, 24 Mar 2026 06:55:51 -0700 Subject: [PATCH 24/24] update test --- python/openai/tests/test_model_management.py | 61 ++++++++++---------- 1 file changed, 32 insertions(+), 29 deletions(-) diff --git a/python/openai/tests/test_model_management.py b/python/openai/tests/test_model_management.py index c20dd7d4e4..7cf420eed1 100644 --- a/python/openai/tests/test_model_management.py +++ b/python/openai/tests/test_model_management.py @@ -192,6 +192,10 @@ def base_url(self, server: OpenAIServer): def all_models(self) -> list[str]: return [TEST_MODEL, TEST_MODEL_2] + @pytest.fixture(scope="class") + def load_model(self) -> list[str]: + return [] + @pytest.fixture(scope="class") def server( self, @@ -199,6 +203,7 @@ def server( tokenizer_model: str, backend: str, model_control_mode: str, + load_model: list[str], ): args = [ "--model-repository", @@ -210,6 +215,9 @@ def server( args += ["--backend", backend] if model_control_mode: args += ["--model-control-mode", model_control_mode] + if load_model: + for model in load_model: + args += ["--load-model", model] with OpenAIServer(args) as openai_server: yield openai_server @@ -531,28 +539,25 @@ class TestModelManagementInference(_ModelManagementBase): def model_control_mode(self): return "explicit" - @pytest.fixture - def all_models(self, model: str) -> list[str]: + @pytest.fixture(scope="class") + def load_model(self, model: str) -> list[str]: + # For tensorrt_llm_bls, we need to load all the dependent models if model == "tensorrt_llm_bls": - return ["postprocessing", "preprocessing", "tensorrt_llm", model] - else: - return [model] + return ["postprocessing", "preprocessing", "tensorrt_llm"] @staticmethod - def _assert_load(base_url, model_list: list[str]): - for model_name in model_list: - r = requests.post( - f"{base_url}/v1/models/{model_name}/load", - ) - assert r.status_code == 200, f"Load {model_name} failed: {r.text}" + def _assert_load(base_url, model_name: str): + r = requests.post( + f"{base_url}/v1/models/{model_name}/load", + ) + assert r.status_code == 200, f"Load {model_name} failed: {r.text}" @staticmethod - def _assert_unload(base_url, model_list: list[str]): - for model_name in model_list: - r = requests.post( - f"{base_url}/v1/models/{model_name}/unload", - ) - assert r.status_code == 200, f"Unload {model_name} failed: {r.text}" + def _assert_unload(base_url, model_name: str): + r = requests.post( + f"{base_url}/v1/models/{model_name}/unload", + ) + assert r.status_code == 200, f"Unload {model_name} failed: {r.text}" @staticmethod def _assert_usage(usage): @@ -575,10 +580,10 @@ def _completions(base_url, model_name: str, **kwargs): }, ) - def test_load_completions(self, base_url, all_models: list[str], model: str): + def test_load_completions(self, base_url, model: str): self._assert_unknown_model(self._completions(base_url, model)) - self._assert_load(base_url, all_models) + self._assert_load(base_url, model) r = self._completions(base_url, model) assert r.status_code == 200 data = r.json() @@ -586,25 +591,23 @@ def test_load_completions(self, base_url, all_models: list[str], model: str): assert data["choices"][0]["finish_reason"] == "stop" self._assert_usage(data["usage"]) - self._assert_unload(base_url, all_models) + self._assert_unload(base_url, model) - def test_unload_rejects_inference( - self, base_url, all_models: list[str], model: str - ): - self._assert_load(base_url, all_models) + def test_unload_rejects_inference(self, base_url, model: str): + self._assert_load(base_url, model) assert self._completions(base_url, model).status_code == 200 - self._assert_unload(base_url, all_models) + self._assert_unload(base_url, model) self._assert_unknown_model(self._completions(base_url, model)) - def test_reload_inference(self, base_url, all_models: list[str], model: str): - self._assert_load(base_url, all_models) + def test_reload_inference(self, base_url, model: str): + self._assert_load(base_url, model) assert self._completions(base_url, model).status_code == 200 - self._assert_unload(base_url, all_models) + self._assert_unload(base_url, model) self._assert_unknown_model(self._completions(base_url, model)) - self._assert_load(base_url, all_models) + self._assert_load(base_url, model) r = self._completions(base_url, model) assert r.status_code == 200 assert r.json()["choices"][0]["text"].strip()