fix(proxy): restrict /nvidia_nim route to NIM-backed model groups and inject router in tests

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
yassin 2026-09-15 22:55:14 +00:00
parent ebb9a3ceeb
commit b50a22b370
5 changed files with 124 additions and 71 deletions

View file

@ -1,7 +1,7 @@
from __future__ import annotations
import re
from collections.abc import Collection, Mapping, Sequence
from collections.abc import Collection, Iterable, Mapping, Sequence
from typing import TYPE_CHECKING, Final
import httpx
@ -14,7 +14,8 @@ from litellm.llms.base_llm.passthrough.transformation import (
)
from litellm.secret_managers.main import get_secret_str
from litellm.types.llms.openai import AllMessageValues
from litellm.types.utils import StandardPassThroughResponseObject
from litellm.types.router import DeploymentTypedDict
from litellm.types.utils import LlmProviders, StandardPassThroughResponseObject
if TYPE_CHECKING:
from httpx import URL, Response
@ -25,6 +26,16 @@ if TYPE_CHECKING:
API_VERSION_SEGMENT: Final = re.compile(r"^v\d+$")
NVIDIA_NIM_MODEL_PREFIX: Final = f"{LlmProviders.NVIDIA_NIM.value}/"
def nvidia_nim_model_groups(deployments: Iterable[DeploymentTypedDict] | None) -> frozenset[str]:
return frozenset(
deployment["model_name"]
for deployment in deployments or ()
if deployment["litellm_params"].get("custom_llm_provider") == LlmProviders.NVIDIA_NIM.value
or deployment["litellm_params"].get("model", "").startswith(NVIDIA_NIM_MODEL_PREFIX)
)
def nvidia_nim_router_model_in_endpoint(endpoint: str, router_models: Collection[str]) -> str | None:

View file

@ -28,7 +28,10 @@ from litellm.litellm_core_utils.url_utils import (
validate_url,
)
from litellm.llms.azure.passthrough.transformation import azure_router_model_in_endpoint
from litellm.llms.nvidia_nim.passthrough.transformation import nvidia_nim_router_model_in_endpoint
from litellm.llms.nvidia_nim.passthrough.transformation import (
nvidia_nim_model_groups,
nvidia_nim_router_model_in_endpoint,
)
from litellm.proxy._types import *
from litellm.proxy.common_utils.http_parsing_utils import extract_nested_form_metadata
from litellm.types.passthrough_endpoints.pass_through_endpoints import (
@ -2052,7 +2055,7 @@ def _router_model_from_nvidia_nim_route(route: str, llm_router: Router | None) -
if llm_router is None:
return None
endpoint: Final = re.sub(r"^/nvidia_nim/", "", route, flags=re.IGNORECASE)
return nvidia_nim_router_model_in_endpoint(endpoint, frozenset(llm_router.get_model_names()))
return nvidia_nim_router_model_in_endpoint(endpoint, nvidia_nim_model_groups(llm_router.get_model_list()))
def _router_model_from_azure_route(route: str, llm_router: Router | None) -> str | None:

View file

@ -36,7 +36,10 @@ from litellm.litellm_core_utils.aws_partition import get_aws_dns_suffix
from litellm.llms.anthropic.common_utils import AnthropicModelInfo
from litellm.llms.azure.passthrough.transformation import foreign_azure_deployment
from litellm.llms.custom_httpx.http_handler import get_async_httpx_client
from litellm.llms.nvidia_nim.passthrough.transformation import nvidia_nim_router_model_in_endpoint
from litellm.llms.nvidia_nim.passthrough.transformation import (
nvidia_nim_model_groups,
nvidia_nim_router_model_in_endpoint,
)
from litellm.llms.vertex_ai.vertex_llm_base import VertexBase
from litellm.passthrough.main import AsyncPassthroughStreamingResponse
from litellm.proxy._types import *
@ -1635,17 +1638,34 @@ async def nvidia_nim_proxy_route(
"""
from litellm.proxy.proxy_server import llm_router
return await relay_nvidia_nim_request(
llm_router=llm_router,
endpoint=endpoint,
request=request,
request_body=await get_request_body(request),
user_api_key_dict=user_api_key_dict,
)
async def relay_nvidia_nim_request(
llm_router: litellm.Router | None,
endpoint: str,
request: Request,
request_body: Mapping[str, object],
user_api_key_dict: UserAPIKeyAuth,
) -> Response:
model_group: Final = (
nvidia_nim_router_model_in_endpoint(endpoint, llm_router.get_model_names()) if llm_router else None
nvidia_nim_router_model_in_endpoint(endpoint, nvidia_nim_model_groups(llm_router.get_model_list()))
if llm_router
else None
)
if llm_router is None or model_group is None:
rejection: Final[RelayRejection] = {
"error": "no LiteLLM model group in the path; call /nvidia_nim/{model_group}/v1/infer with a model "
"error": "no NVIDIA NIM model group in the path; call /nvidia_nim/{model_group}/v1/infer with a model "
"from your `model_list` whose `model` starts with `nvidia_nim/`"
}
raise HTTPException(status_code=400, detail=rejection)
request_body: Final = await get_request_body(request)
is_streaming_request: Final = is_passthrough_request_streaming(request_body)
return await open_sse_before_first_byte(
_relay_router_model(

View file

@ -844,6 +844,10 @@ def _nvidia_nim_relay_router():
"api_key": "k",
},
},
{
"model_name": "gpt-4o",
"litellm_params": {"model": "openai/gpt-4o", "api_key": "k"},
},
]
)
@ -868,6 +872,7 @@ NIM_INFER_BODY = {"input": [{"type": "image_url", "url": "data:image/png;base64,
("/nvidia_nim/v1/infer", NIM_INFER_BODY, None),
("/nvidia_nim/unknown-group/v1/infer", NIM_INFER_BODY, None),
("/nvidia_nim/nim-page-elements-v2/v1/infer", NIM_INFER_BODY, None),
("/nvidia_nim/gpt-4o/v1/infer", NIM_INFER_BODY, None),
],
)
def test_get_model_from_request_nvidia_nim_relay_routes_use_the_model_group_in_the_path(route, request_data, expected):

View file

@ -40,7 +40,7 @@ from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import (
llm_passthrough_factory_proxy_route,
milvus_proxy_route,
mistral_proxy_route,
nvidia_nim_proxy_route,
relay_nvidia_nim_request,
openai_proxy_route,
vertex_discovery_proxy_route,
vertex_proxy_route,
@ -5392,20 +5392,10 @@ class TestNvidiaNimProxyRoute:
request.query_params = {}
return request
def _install_router(self, monkeypatch, router, body: dict) -> None:
import litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints as ep
import litellm.proxy.proxy_server as proxy_server
async def fake_get_request_body(_request):
return dict(body)
monkeypatch.setattr(proxy_server, "llm_router", router)
monkeypatch.setattr(ep, "get_request_body", fake_get_request_body)
def _recording_router(self, captured: list[dict], model_names: tuple[str, ...]):
def _recording_router(self, captured: list[dict], deployments: dict[str, str]):
class RecordingRouter:
def get_model_names(self):
return list(model_names)
def get_model_list(self):
return [{"model_name": name, "litellm_params": {"model": model}} for name, model in deployments.items()]
async def allm_passthrough_route(self, **kwargs):
captured.append(kwargs)
@ -5415,18 +5405,31 @@ class TestNvidiaNimProxyRoute:
return RecordingRouter()
@pytest.mark.asyncio
async def test_model_group_in_the_path_selects_the_deployment_and_the_body_stays_model_free(self, monkeypatch):
captured: list[dict] = []
self._install_router(
monkeypatch, self._recording_router(captured, ("nim-page-elements", "nim-table")), NIM_INFER_BODY
async def _relay(self, llm_router, endpoint: str, body: dict, user_api_key_dict=None) -> Response:
return await relay_nvidia_nim_request(
llm_router=llm_router,
endpoint=endpoint,
request=self._request(),
request_body=dict(body),
user_api_key_dict=user_api_key_dict or UserAPIKeyAuth(api_key="hashed-token"),
)
result = await nvidia_nim_proxy_route(
endpoint="nim-page-elements/v1/infer",
request=self._request(),
fastapi_response=MagicMock(spec=Response),
user_api_key_dict=UserAPIKeyAuth(api_key="hashed-token", team_id="team-1"),
@pytest.mark.asyncio
async def test_model_group_in_the_path_selects_the_deployment_and_the_body_stays_model_free(self):
captured: list[dict] = []
router = self._recording_router(
captured,
{
"nim-page-elements": "nvidia_nim/nvidia/nemoretriever-page-elements-v2",
"nim-table": "nvidia_nim/nvidia/nemoretriever-table-structure-v1",
},
)
result = await self._relay(
router,
"nim-page-elements/v1/infer",
NIM_INFER_BODY,
UserAPIKeyAuth(api_key="hashed-token", team_id="team-1"),
)
(relay,) = captured
@ -5441,65 +5444,81 @@ class TestNvidiaNimProxyRoute:
assert result.headers["x-nim-request"] == "r1"
@pytest.mark.asyncio
async def test_model_group_with_a_slash_is_matched_as_the_longest_leading_path(self, monkeypatch):
async def test_model_group_with_a_slash_is_matched_as_the_longest_leading_path(self):
captured: list[dict] = []
self._install_router(
monkeypatch, self._recording_router(captured, ("nvidia/nemoretriever-page-elements-v2",)), NIM_INFER_BODY
router = self._recording_router(
captured, {"nvidia/nemoretriever-page-elements-v2": "nvidia_nim/nvidia/nemoretriever-page-elements-v2"}
)
await nvidia_nim_proxy_route(
endpoint="nvidia/nemoretriever-page-elements-v2/v1/infer",
request=self._request(),
fastapi_response=MagicMock(spec=Response),
user_api_key_dict=UserAPIKeyAuth(api_key="hashed-token"),
)
await self._relay(router, "nvidia/nemoretriever-page-elements-v2/v1/infer", NIM_INFER_BODY)
assert captured[0]["model"] == "nvidia/nemoretriever-page-elements-v2"
@pytest.mark.asyncio
@pytest.mark.parametrize("endpoint", ["v1/infer", "unknown-group/v1/infer", "nim-page-elements-v2/v1/infer"])
async def test_path_without_a_configured_model_group_is_rejected_before_any_upstream_call(
self, monkeypatch, endpoint
):
async def test_custom_llm_provider_marks_a_deployment_as_nim_without_the_model_prefix(self):
captured: list[dict] = []
class ProviderRouter:
def get_model_list(self):
return [
{
"model_name": "page-elements",
"litellm_params": {
"model": "nvidia/nemoretriever-page-elements-v2",
"custom_llm_provider": "nvidia_nim",
},
}
]
async def allm_passthrough_route(self, **kwargs):
captured.append(kwargs)
return httpx.Response(200, json={"data": []})
await self._relay(ProviderRouter(), "page-elements/v1/infer", NIM_INFER_BODY)
assert captured[0]["model"] == "page-elements"
@pytest.mark.asyncio
@pytest.mark.parametrize(
"endpoint",
["v1/infer", "unknown-group/v1/infer", "nim-page-elements-v2/v1/infer", "gpt-4o/v1/infer"],
)
async def test_path_without_a_nim_model_group_is_rejected_before_any_upstream_call(self, endpoint):
from fastapi import HTTPException
captured: list[dict] = []
self._install_router(monkeypatch, self._recording_router(captured, ("nim-page-elements",)), NIM_INFER_BODY)
router = self._recording_router(
captured,
{"nim-page-elements": "nvidia_nim/nvidia/nemoretriever-page-elements-v2", "gpt-4o": "openai/gpt-4o"},
)
with pytest.raises(HTTPException) as exc_info:
await nvidia_nim_proxy_route(
endpoint=endpoint,
request=self._request(),
fastapi_response=MagicMock(spec=Response),
user_api_key_dict=UserAPIKeyAuth(api_key="hashed-token"),
)
await self._relay(router, endpoint, NIM_INFER_BODY)
assert exc_info.value.status_code == 400
assert captured == []
@pytest.mark.asyncio
async def test_no_router_is_rejected_before_any_upstream_call(self, monkeypatch):
async def test_no_router_is_rejected_before_any_upstream_call(self):
from fastapi import HTTPException
self._install_router(monkeypatch, None, NIM_INFER_BODY)
with pytest.raises(HTTPException) as exc_info:
await nvidia_nim_proxy_route(
endpoint="nim-page-elements/v1/infer",
request=self._request(),
fastapi_response=MagicMock(spec=Response),
user_api_key_dict=UserAPIKeyAuth(api_key="hashed-token"),
)
await self._relay(None, "nim-page-elements/v1/infer", NIM_INFER_BODY)
assert exc_info.value.status_code == 400
@pytest.mark.asyncio
async def test_upstream_rejection_is_relayed_with_its_status_body_and_headers(self, monkeypatch):
async def test_upstream_rejection_is_relayed_with_its_status_body_and_headers(self):
upstream_body = {"detail": "input[0].url must be a data URL"}
class RejectingRouter:
def get_model_names(self):
return ["nim-page-elements"]
def get_model_list(self):
return [
{
"model_name": "nim-page-elements",
"litellm_params": {"model": "nvidia_nim/nvidia/nemoretriever-page-elements-v2"},
}
]
async def allm_passthrough_route(self, **kwargs):
upstream_request = httpx.Request("POST", "http://nim.internal:8000/v1/infer")
@ -5508,13 +5527,8 @@ class TestNvidiaNimProxyRoute:
)
raise httpx.HTTPStatusError("422", request=upstream_request, response=upstream)
self._install_router(monkeypatch, RejectingRouter(), {"input": [{"type": "image_url", "url": "x"}]})
result = await nvidia_nim_proxy_route(
endpoint="nim-page-elements/v1/infer",
request=self._request(),
fastapi_response=MagicMock(spec=Response),
user_api_key_dict=UserAPIKeyAuth(api_key="hashed-token"),
result = await self._relay(
RejectingRouter(), "nim-page-elements/v1/infer", {"input": [{"type": "image_url", "url": "x"}]}
)
assert result.status_code == 422