feat(proxy): let callbacks filter the model listing routes per caller (#43027)

* feat(proxy): let callbacks filter the model listing routes per caller

* fix(proxy): offer every listed name to the listing callback, agent groups and deployment lookups included

* fix(proxy): hide aliases of a team model by its public name and offer /model/info lookups the listed name

* fix(proxy): map a malformed model listing filter return to the proxy error contract and document legacy team names

---------

Co-authored-by: mateo-berri <277851410+mateo-berri@users.noreply.github.com>
This commit is contained in:
devin-ai-integration[bot] 2026-09-24 15:52:44 -07:00 • committed by GitHub
parent bf0187072b
commit 7b25a151bd
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
5 changed files with 596 additions and 23 deletions

View file

@ -421,6 +421,24 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac
): # raise exception if invalid, return a str for the user to receive - if rejected, or return a modified dictionary for passing into litellm
pass
async def async_filter_listed_models(
self,
user_api_key_dict: UserAPIKeyAuth,
model_names: Sequence[str],
) -> Sequence[str]:
"""Runs on the model listing routes (`/v1/models`, `/v1/models/{id}`, `/model/info`,
`/model_group/info`) with the public model names the route would otherwise return, so a
lookup of one model may offer just that name: decide per name, never by position in the
sequence. Return the names to keep as a sequence of strings; a name left out disappears
from every listing, any alias of it offered in the same call goes with it, and
`/v1/models/{id}` answers 404 for it, exactly as for a model that does not exist. Names
outside `model_names` are ignored, so a callback can only narrow the listing, never widen
it. Under `use_team_public_model_name: false`, `/v1/models` and `/model_group/info` list a
team model by its internal routing name while `/model/info` keeps its public name, so hide
both names to hide it on every route.
"""
return model_names
async def async_post_call_response_headers_hook(
self,
data: dict,

View file

@ -152,6 +152,7 @@ from litellm.router_utils.auto_router_tuning_baseline import (
snapshot_tuning_baselines,
tuning_limit_violation,
)
from litellm.router_utils.common_utils import resolve_model_group_alias
from litellm.router_utils.routing_groups import parse_routing_groups
from litellm.types.caching import RedisPipelineIncrementOperation
from litellm.types.utils import (
@ -11104,6 +11105,40 @@ class ProxyStartupEvent:
#### API ENDPOINTS ####
async def _names_hidden_by_listing_callbacks(
user_api_key_dict: UserAPIKeyAuth, model_names: Sequence[str]
) -> frozenset[str]:
hidden: Final = await proxy_logging_obj.hidden_by_listing_callbacks(user_api_key_dict, model_names)
if not hidden or llm_router is None:
return hidden
aliases: Final = llm_router.model_group_alias
internal_to_public: Final = TeamModelNameTranslator.build_internal_to_public_map(llm_router, general_settings)
return hidden | frozenset(
alias
for alias in aliases
if (target := resolve_model_group_alias(aliases, alias)) is not None
and internal_to_public.get(target, target) in hidden
)
async def _entries_kept_by_listing_callbacks(
entries: Sequence[tuple[str, str]], user_api_key_dict: UserAPIKeyAuth
) -> tuple[tuple[str, str], ...]:
hidden: Final = await _names_hidden_by_listing_callbacks(
user_api_key_dict, tuple(response_id for response_id, _ in entries)
)
if not hidden:
return tuple(entries)
return tuple(entry for entry in entries if entry[0] not in hidden)
async def _deployment_hidden_by_listing_callbacks(deployment: Deployment, user_api_key_dict: UserAPIKeyAuth) -> bool:
listed_name: Final = _translate_model_name_for_response(deployment.model_dump(exclude_none=True)).get("model_name")
if not isinstance(listed_name, str):
return False
return listed_name in await _names_hidden_by_listing_callbacks(user_api_key_dict, (listed_name,))
@router.get("/v1/models", dependencies=[Depends(user_api_key_auth)], tags=["model management"])
@router.get(
"/models", dependencies=[Depends(user_api_key_auth)], tags=["model management"]
@ -11254,7 +11289,9 @@ async def model_list(
# The internal routing key drives the metadata/fallback lookup, while the
# public name is what the client sees as the model id.
model_data = []
admin_entries: Final = TeamModelNameTranslator.listing_entries(all_models, llm_router, settings)
admin_entries: Final = await _entries_kept_by_listing_callbacks(
TeamModelNameTranslator.listing_entries(all_models, llm_router, settings), user_api_key_dict
)
for response_id, lookup_id in admin_entries:
model_info = create_model_info_response(
model_id=lookup_id,
@ -11310,7 +11347,10 @@ async def model_list(
# public name is what the client sees as the model id.
model_data = []
entries: Final = alias_listing_entries(
TeamModelNameTranslator.listing_entries(all_models, llm_router, settings), caller_aliases
await _entries_kept_by_listing_callbacks(
TeamModelNameTranslator.listing_entries(all_models, llm_router, settings), user_api_key_dict
),
caller_aliases,
)
for response_id, lookup_id in entries:
model_info = create_model_info_response(
@ -11404,13 +11444,24 @@ async def model_info(
llm_router=llm_router,
)
hidden_names: Final = blocked_names | unhealthy_names
if hidden_names:
all_models = [m for m in all_models if m not in hidden_names]
internal_to_public: Final = TeamModelNameTranslator.build_internal_to_public_map(llm_router, settings)
callback_hidden_names: Final = await _names_hidden_by_listing_callbacks(
user_api_key_dict,
tuple(
response_id
for response_id, _ in TeamModelNameTranslator.listing_entries(
tuple(m for m in all_models if m not in hidden_names), llm_router, settings
)
),
)
if hidden_names or callback_hidden_names:
all_models = [
m for m in all_models if m not in hidden_names and internal_to_public.get(m, m) not in callback_hidden_names
]
undiscoverable_names: Final = undiscoverable_model_names(
all_models, llm_router, user_api_key_dict, team_id or user_api_key_dict.team_id
)
internal_to_public: Final = TeamModelNameTranslator.build_internal_to_public_map(llm_router, settings)
aliased_model_id: Final = alias_target(
model_id,
caller_alias_maps(
@ -15730,7 +15781,7 @@ async def model_info_v1(
if litellm_model_id is not None:
# user is trying to get specific model from litellm router
deployment_info: Final = llm_router.get_deployment(model_id=litellm_model_id)
if deployment_info is None:
if deployment_info is None or await _deployment_hidden_by_listing_callbacks(deployment_info, user_api_key_dict):
raise HTTPException(
status_code=400,
detail={"error": f"Model id = {litellm_model_id} not found on litellm proxy"},
@ -15819,10 +15870,17 @@ async def model_info_v1(
general_settings=general_settings,
llm_router=llm_router,
)
visible_models: Final = discoverable_rows(
servable_rows: Final = discoverable_rows(
(model for model in all_models if model.get("model_name") not in hidden_names),
user_api_key_dict,
)
listed_names: Final = tuple(
dict.fromkeys(name for model in servable_rows if isinstance(name := model.get("model_name"), str))
)
callback_hidden_names: Final = await _names_hidden_by_listing_callbacks(user_api_key_dict, listed_names)
visible_models: Final = tuple(
model for model in servable_rows if model.get("model_name") not in callback_hidden_names
)
verbose_proxy_logger.debug("all_models: %s", visible_models)
return _model_info_json_response(visible_models)
@ -15871,7 +15929,7 @@ async def model_deprecations(
def _get_model_group_info(
llm_router: Router, all_models_str: list[str], model_group: str | None
llm_router: Router, all_models_str: Sequence[str], model_group: str | None
) -> list[ModelGroupInfoProxy]:
model_groups: Final[list[ModelGroupInfoProxy]] = []
@ -16104,23 +16162,34 @@ async def model_group_info(
undiscoverable_group_names: Final = undiscoverable_model_names(
all_models_str, llm_router, user_api_key_dict, user_api_key_dict.team_id
)
model_groups: list[ModelGroupInfoProxy] = _get_model_group_info(
llm_router=llm_router,
all_models_str=[name for name in all_models_str if name not in undiscoverable_group_names],
model_group=model_group,
)
listed_group_names: Final = tuple(name for name in all_models_str if name not in undiscoverable_group_names)
# Append A2A agents to model groups
from litellm.proxy.agent_endpoints.model_list_helpers import (
append_agents_to_model_group,
)
model_groups = await append_agents_to_model_group(
model_groups=model_groups,
model_groups: Final = await append_agents_to_model_group(
model_groups=_get_model_group_info(
llm_router=llm_router, all_models_str=listed_group_names, model_group=model_group
),
user_api_key_dict=user_api_key_dict,
)
internal_to_public: Final = TeamModelNameTranslator.build_internal_to_public_map(llm_router, general_settings)
public_group_names: Final = tuple(
internal_to_public.get(group.model_group, group.model_group) for group in model_groups
)
callback_hidden_names: Final = await _names_hidden_by_listing_callbacks(
user_api_key_dict, tuple(dict.fromkeys(public_group_names))
)
return {"data": model_groups}
return {
"data": [
group
for group, public_name in zip(model_groups, public_group_names, strict=True)
if public_name not in callback_hidden_names
]
}
@router.get(

View file

@ -36,6 +36,7 @@ from typing import (
Final,
Generic,
Literal,
NoReturn,
Optional,
Protocol,
TypeAlias,
@ -106,7 +107,7 @@ except ImportError:
raise ImportError("backoff is not installed. Please install it via 'pip install backoff'")
from fastapi import HTTPException, status
from pydantic import TypeAdapter
from pydantic import TypeAdapter, ValidationError
import litellm
import litellm.litellm_core_utils
@ -1136,11 +1137,51 @@ class _CallbackCapabilities:
# avoids the per-request ``get_custom_logger_compatible_class`` walk for
# every string entry in ``litellm.callbacks``.
resolved_callbacks: tuple[object, ...] = field(default_factory=tuple)
listed_models_filters: tuple[CustomLogger, ...] = field(default_factory=tuple)
def _overrides_hook(callback: CustomLogger, hook_name: str) -> bool:
leaf_to_base: Final = takewhile(lambda klass: klass is not CustomLogger, type(callback).__mro__)
return any(hook_name in klass.__dict__ for klass in leaf_to_base)
def _overrides_moderation_hook(callback: CustomLogger) -> bool:
leaf_to_base: Final = takewhile(lambda klass: klass is not CustomLogger, type(callback).__mro__)
return any("async_moderation_hook" in klass.__dict__ for klass in leaf_to_base)
return _overrides_hook(callback, "async_moderation_hook")
_LISTED_MODEL_NAMES: Final = TypeAdapter(tuple[str, ...])
@dataclass(frozen=True, slots=True)
class MalformedListingFilterReturn:
callback: str
tag: Literal["malformed_listing_filter_return"] = "malformed_listing_filter_return"
async def _names_kept_by_listing_callbacks(
callbacks: Sequence[CustomLogger],
user_api_key_dict: UserAPIKeyAuth,
model_names: tuple[str, ...],
) -> tuple[str, ...] | MalformedListingFilterReturn:
if not callbacks or not model_names:
return model_names
returned: Final = await callbacks[0].async_filter_listed_models(user_api_key_dict, model_names)
try:
kept: Final = frozenset(_LISTED_MODEL_NAMES.validate_python(returned))
except ValidationError:
return MalformedListingFilterReturn(callback=type(callbacks[0]).__name__)
return await _names_kept_by_listing_callbacks(
callbacks[1:], user_api_key_dict, tuple(name for name in model_names if name in kept)
)
def _raise_malformed_listing_filter_return(error: MalformedListingFilterReturn) -> NoReturn:
raise ProxyException(
message=f"{error.callback}.async_filter_listed_models must return a sequence of model names",
type=ProxyErrorTypes.internal_server_error,
param=None,
code=500,
)
class ProxyLogging:
@ -2808,6 +2849,9 @@ class ProxyLogging:
has_moderation_override=has_moderation_override,
iterator_overrides=tuple(iterator_overrides),
resolved_callbacks=tuple(resolved_callbacks),
listed_models_filters=tuple(
callback for callback in resolved_callbacks if _overrides_hook(callback, "async_filter_listed_models")
),
)
# Limit cache to handle test churn without leaking; production
# callback lists are stable so this rarely grows past 1 entry.
@ -3715,6 +3759,18 @@ class ProxyLogging:
verbose_proxy_logger.exception("Error in post_call_response_headers_hook: %s", str(e))
return merged_headers
async def hidden_by_listing_callbacks(
self, user_api_key_dict: UserAPIKeyAuth, model_names: Sequence[str]
) -> frozenset[str]:
filters: Final = ProxyLogging._callback_capabilities().listed_models_filters
if not filters:
return frozenset()
candidates: Final = tuple(model_names)
kept: Final = await _names_kept_by_listing_callbacks(filters, user_api_key_dict, candidates)
if isinstance(kept, MalformedListingFilterReturn):
_raise_malformed_listing_filter_return(kept)
return frozenset(candidates).difference(kept)
@staticmethod
def _build_litellm_call_info(data: dict, response: object) -> dict[str, object]:
"""

View file

@ -322,7 +322,9 @@ def test_get_proxy_model_info_shows_litellm_params_pricing_and_names_it_as_an_ov
def test_get_proxy_model_info_names_config_model_info_pricing_as_an_override(monkeypatch, local_model_cost_map):
"""Pricing declared under ``model_info`` in config.yaml overrides the cost map too."""
info = _enriched_model_info(
monkeypatch, {"model": "openai/gpt-5.6"}, {"id": "dep-config", "db_model": False, "output_cost_per_token": 7e-06}
monkeypatch,
{"model": "openai/gpt-5.6"},
{"id": "dep-config", "db_model": False, "output_cost_per_token": 7e-06},
)
assert info["pricing_overrides"] == ("output_cost_per_token",)
assert info["output_cost_per_token"] == 7e-06
@ -399,7 +401,9 @@ def test_model_info_reports_null_cost_for_unpriced_deployment_and_zero_for_decla
def enriched_cost(model_name: str) -> tuple:
deployment = router.get_model_list(model_name=model_name)[0]
info = proxy_server._enrich_model_info_with_litellm_data({**deployment, "model_info": dict(deployment["model_info"])})["model_info"]
info = proxy_server._enrich_model_info_with_litellm_data(
{**deployment, "model_info": dict(deployment["model_info"])}
)["model_info"]
return info.get("input_cost_per_token"), info.get("output_cost_per_token")
assert enriched_cost("vllm-unpriced") == (None, None)
@ -643,7 +647,6 @@ def model_group_info_router(monkeypatch):
monkeypatch.setattr(proxy_server, "user_model", None)
monkeypatch.setattr(proxy_server, "general_settings", {})
monkeypatch.setattr(proxy_server, "prisma_client", None)
monkeypatch.setattr(proxy_server, "proxy_logging_obj", None)
monkeypatch.setattr(proxy_server, "user_api_key_cache", None)
monkeypatch.setattr(proxy_server, "_get_model_group_info", model_group_info)
@ -671,7 +674,9 @@ def test_model_group_info_proxy_admin_ignores_key_model_restriction(
@pytest.mark.parametrize("admin_role", ["proxy_admin", "proxy_admin_viewer"])
def test_model_group_info_proxy_admin_expands_wildcard_deployments(client, auth_as, model_group_info_router, admin_role):
def test_model_group_info_proxy_admin_expands_wildcard_deployments(
client, auth_as, model_group_info_router, admin_role
):
from litellm.proxy._types import LitellmUserRoles
from litellm.proxy.auth.model_checks import get_known_models_from_wildcard

View file

@ -0,0 +1,425 @@
"""
Tests for `CustomLogger.async_filter_listed_models` on the model listing endpoints:
GET /v1/models (`model_list`, OpenAI and Anthropic shapes), GET /v1/models/{id}
(`model_info`), GET /v1/model/info (`model_info_v1`) and GET /model_group/info
(`model_group_info`). A registered callback that overrides the hook decides per
caller which of the names the route would list are kept; the rest disappear and
`/v1/models/{id}` answers 404 for them.
"""
import json
from collections.abc import Sequence
import pytest
from fastapi import HTTPException
from starlette.requests import Request
import litellm
from litellm import Router
from litellm.integrations.custom_logger import CustomLogger
from litellm.proxy import proxy_server
from litellm.proxy._types import LitellmUserRoles, ProxyException, UserAPIKeyAuth
from litellm.proxy.utils import ProxyLogging
class _Gate(CustomLogger):
def __init__(self, hidden: frozenset[str] = frozenset(), extra: tuple[str, ...] = ()) -> None:
super().__init__()
self.hidden = hidden
self.extra = extra
self.seen: list[tuple[str, ...]] = []
async def async_filter_listed_models(
self, user_api_key_dict: UserAPIKeyAuth, model_names: Sequence[str]
) -> Sequence[str]:
self.seen.append(tuple(model_names))
return [*(name for name in model_names if name not in self.hidden), *self.extra]
class _InferenceOnlyGate(CustomLogger):
async def async_pre_call_hook(self, user_api_key_dict, cache, data, call_type):
if data.get("model") == "restricted-model":
raise HTTPException(status_code=403, detail="not entitled to this model")
return data
class _RaisingGate(CustomLogger):
async def async_filter_listed_models(
self, user_api_key_dict: UserAPIKeyAuth, model_names: Sequence[str]
) -> Sequence[str]:
raise HTTPException(status_code=503, detail="entitlement service down")
class _ReversingGate(CustomLogger):
async def async_filter_listed_models(
self, user_api_key_dict: UserAPIKeyAuth, model_names: Sequence[str]
) -> Sequence[str]:
return list(reversed(model_names))
class _StringReturningGate(CustomLogger):
async def async_filter_listed_models(self, user_api_key_dict: UserAPIKeyAuth, model_names: Sequence[str]) -> str:
return "open-model"
def _deployment(model_name: str, model: str = "openai/gpt-4o", **model_info):
return {
"model_name": model_name,
"litellm_params": {"model": model, "api_key": "sk-fake"},
"model_info": {"id": f"{model_name}-id", **model_info},
}
def _install_router(monkeypatch, *deployments, **router_kwargs) -> Router:
router = Router(model_list=list(deployments), **router_kwargs)
monkeypatch.setattr(proxy_server, "llm_router", router)
monkeypatch.setattr(proxy_server, "llm_model_list", router.model_list)
monkeypatch.setattr(proxy_server, "prisma_client", None)
monkeypatch.setattr(proxy_server, "general_settings", {})
monkeypatch.setattr(proxy_server, "user_model", None)
return router
def _register(monkeypatch, *callbacks: CustomLogger) -> None:
monkeypatch.setattr(litellm, "callbacks", list(callbacks))
ProxyLogging._callback_capabilities_cache.clear()
@pytest.fixture
def two_model_router(monkeypatch) -> Router:
return _install_router(monkeypatch, _deployment("open-model"), _deployment("restricted-model"))
@pytest.fixture
def team_router(monkeypatch) -> Router:
return _install_router(
monkeypatch,
_deployment("gpt-4"),
_deployment("model_name_team1_abc", team_id="team1", team_public_model_name="team-gpt"),
_deployment("model_name_team1_def", team_id="team1", team_public_model_name="team-chat"),
)
@pytest.fixture
def team_admin_privileges(monkeypatch) -> None:
from litellm.proxy.management_endpoints import common_utils
async def _is_team_admin(**kwargs) -> bool:
return True
monkeypatch.setattr(common_utils, "_user_has_admin_privileges", _is_team_admin)
def _non_admin(**kwargs) -> UserAPIKeyAuth:
return UserAPIKeyAuth(api_key="sk-test", user_role=LitellmUserRoles.INTERNAL_USER, **kwargs)
def _admin() -> UserAPIKeyAuth:
return UserAPIKeyAuth(api_key="sk-test", user_id="u", user_role=LitellmUserRoles.PROXY_ADMIN, team_models=[])
def _team_member() -> UserAPIKeyAuth:
return UserAPIKeyAuth(
api_key="sk-test",
user_id="u",
user_role=LitellmUserRoles.INTERNAL_USER,
team_id="team1",
team_models=["model_name_team1_abc", "model_name_team1_def"],
models=["model_name_team1_abc", "model_name_team1_def"],
)
def _anthropic_request() -> Request:
return Request(
scope={
"type": "http",
"method": "GET",
"path": "/v1/models",
"query_string": b"",
"headers": [(b"anthropic-version", b"2023-06-01")],
}
)
async def _v1_models(user_api_key_dict: UserAPIKeyAuth, **kwargs) -> list[str]:
response = await proxy_server.model_list(user_api_key_dict=user_api_key_dict, **kwargs)
return [m["id"] for m in response["data"]]
async def _v1_model_info_names(user_api_key_dict: UserAPIKeyAuth) -> list[str]:
response = await proxy_server.model_info_v1(user_api_key_dict=user_api_key_dict)
return [row["model_name"] for row in json.loads(response.body)["data"]]
async def _model_groups(user_api_key_dict: UserAPIKeyAuth) -> list[str]:
response = await proxy_server.model_group_info(user_api_key_dict=user_api_key_dict)
return [group.model_group for group in response["data"]]
async def _model_by_id_status(model_id: str, user_api_key_dict: UserAPIKeyAuth) -> int:
try:
response = await proxy_server.model_info(model_id=model_id, user_api_key_dict=user_api_key_dict)
except HTTPException as error:
return error.status_code
assert response["id"] == model_id
return 200
@pytest.mark.asyncio
async def test_v1_models_lists_only_the_names_the_callback_keeps(two_model_router, monkeypatch):
_register(monkeypatch, _Gate(hidden=frozenset({"restricted-model"})))
assert await _v1_models(_non_admin()) == ["open-model"]
assert await _v1_models(_admin()) == ["open-model"]
assert await _v1_models(_non_admin(), request=_anthropic_request()) == ["open-model"]
@pytest.mark.asyncio
async def test_v1_models_scope_expand_applies_the_callback(two_model_router, team_admin_privileges, monkeypatch):
_register(monkeypatch, _Gate(hidden=frozenset({"restricted-model"})))
assert await _v1_models(_non_admin(), scope="expand") == ["open-model"]
assert await _v1_models(_admin(), scope="expand") == ["open-model"]
@pytest.mark.asyncio
async def test_v1_models_by_id_answers_404_for_a_name_the_callback_leaves_out(two_model_router, monkeypatch):
_register(monkeypatch, _Gate(hidden=frozenset({"restricted-model"})))
assert await _model_by_id_status("restricted-model", _non_admin()) == 404
assert await _model_by_id_status("open-model", _non_admin()) == 200
@pytest.mark.asyncio
async def test_v1_model_info_lists_only_the_rows_the_callback_keeps(two_model_router, monkeypatch):
_register(monkeypatch, _Gate(hidden=frozenset({"restricted-model"})))
assert await _v1_model_info_names(_non_admin()) == ["open-model"]
assert await _v1_model_info_names(_admin()) == ["open-model"]
@pytest.mark.asyncio
async def test_model_group_info_lists_only_the_groups_the_callback_keeps(two_model_router, monkeypatch):
_register(monkeypatch, _Gate(hidden=frozenset({"restricted-model"})))
assert await _model_groups(_non_admin()) == ["open-model"]
assert await _model_groups(_admin()) == ["open-model"]
@pytest.mark.asyncio
async def test_a_callback_without_the_hook_changes_no_listing(two_model_router, monkeypatch):
_register(monkeypatch, _InferenceOnlyGate())
assert await _v1_models(_non_admin()) == ["open-model", "restricted-model"]
assert await _model_by_id_status("restricted-model", _non_admin()) == 200
assert await _v1_model_info_names(_non_admin()) == ["open-model", "restricted-model"]
assert await _model_groups(_non_admin()) == ["open-model", "restricted-model"]
@pytest.mark.asyncio
async def test_callback_cannot_add_a_name_it_was_not_offered(two_model_router, monkeypatch):
_register(monkeypatch, _Gate(extra=("ghost-model",)))
assert await _v1_models(_non_admin()) == ["open-model", "restricted-model"]
assert await _model_by_id_status("ghost-model", _non_admin()) == 404
@pytest.mark.asyncio
async def test_callbacks_narrow_in_registration_order(monkeypatch):
_install_router(monkeypatch, _deployment("a"), _deployment("b"), _deployment("c"))
first: _Gate = _Gate(hidden=frozenset({"a"}))
second: _Gate = _Gate(hidden=frozenset({"b"}))
_register(monkeypatch, first, second)
assert await _v1_models(_non_admin()) == ["c"]
assert first.seen == [("a", "b", "c")]
assert second.seen == [("b", "c")]
@pytest.mark.asyncio
async def test_callback_sees_and_filters_team_models_by_their_public_name(team_router, monkeypatch):
gate: _Gate = _Gate(hidden=frozenset({"team-gpt"}))
_register(monkeypatch, gate)
assert await _v1_models(_team_member()) == ["team-chat"]
assert await _model_by_id_status("team-gpt", _team_member()) == 404
assert await _model_by_id_status("team-chat", _team_member()) == 200
assert all("team-gpt" in seen and "model_name_team1_abc" not in seen for seen in gate.seen)
@pytest.mark.asyncio
async def test_callback_sees_public_team_names_on_every_listing_route(team_router, monkeypatch):
gate: _Gate = _Gate(hidden=frozenset({"team-gpt"}))
_register(monkeypatch, gate)
assert await _v1_models(_team_member()) == ["team-chat"]
assert await _v1_model_info_names(_team_member()) == ["team-chat"]
assert await _model_groups(_team_member()) == ["model_name_team1_def"]
assert await _model_by_id_status("team-gpt", _team_member()) == 404
assert len(gate.seen) == 4
assert all(sorted(seen) == ["team-chat", "team-gpt"] for seen in gate.seen)
@pytest.mark.asyncio
async def test_router_alias_follows_its_hidden_target(monkeypatch):
_install_router(
monkeypatch,
_deployment("open-model"),
_deployment("restricted-model"),
model_group_alias={"mini": "restricted-model", "wide": "open-model"},
)
_register(monkeypatch, _Gate(hidden=frozenset({"restricted-model"})))
assert sorted(await _v1_models(_non_admin())) == ["open-model", "wide"]
assert sorted(await _v1_model_info_names(_non_admin())) == ["open-model", "wide"]
assert sorted(await _model_groups(_non_admin())) == ["open-model", "wide"]
_register(monkeypatch, _Gate(hidden=frozenset({"mini"})))
assert sorted(await _v1_models(_non_admin())) == ["open-model", "restricted-model", "wide"]
@pytest.mark.asyncio
async def test_router_alias_of_a_team_model_follows_its_hidden_public_name(monkeypatch):
_install_router(
monkeypatch,
_deployment("gpt-4"),
_deployment("model_name_team1_abc", team_id="team1", team_public_model_name="team-gpt"),
model_group_alias={"team-alias": "model_name_team1_abc"},
)
caller: UserAPIKeyAuth = _non_admin(
user_id="u",
team_id="team1",
team_models=["model_name_team1_abc", "team-alias"],
models=["model_name_team1_abc", "team-alias"],
)
_register(monkeypatch, _Gate(hidden=frozenset()))
assert sorted(await _v1_models(caller)) == ["team-alias", "team-gpt"]
_register(monkeypatch, _Gate(hidden=frozenset({"team-gpt"})))
assert await _v1_models(caller) == []
assert await _model_groups(caller) == []
@pytest.mark.asyncio
async def test_v1_model_info_offers_only_the_rows_the_caller_would_see(monkeypatch):
_install_router(monkeypatch, _deployment("open-model"), _deployment("hidden-model", discoverable=False))
gate: _Gate = _Gate()
_register(monkeypatch, gate)
assert await _v1_model_info_names(_non_admin()) == ["open-model"]
assert await _v1_model_info_names(_admin()) == ["open-model", "hidden-model"]
assert gate.seen == [("open-model",), ("open-model", "hidden-model")]
@pytest.mark.asyncio
async def test_listing_keeps_its_order_whatever_order_the_callback_returns(monkeypatch):
_install_router(monkeypatch, _deployment("a"), _deployment("b"), _deployment("c"))
_register(monkeypatch, _ReversingGate())
assert await _v1_models(_non_admin()) == ["a", "b", "c"]
assert await _v1_model_info_names(_non_admin()) == ["a", "b", "c"]
@pytest.mark.asyncio
async def test_a_callback_returning_a_string_is_an_error_not_an_empty_listing(two_model_router, monkeypatch):
_register(monkeypatch, _StringReturningGate())
with pytest.raises(ProxyException, match=r"_StringReturningGate\.async_filter_listed_models") as raised:
await _v1_models(_non_admin())
assert raised.value.code == "500"
@pytest.mark.asyncio
async def test_alias_of_a_hidden_model_is_not_listed(two_model_router, monkeypatch):
_register(monkeypatch, _Gate(hidden=frozenset({"restricted-model"})))
caller = _non_admin(aliases={"mini": "restricted-model", "wide": "open-model"})
assert await _v1_models(caller) == ["open-model", "wide"]
assert await _model_by_id_status("mini", caller) == 404
assert await _model_by_id_status("wide", caller) == 200
@pytest.mark.asyncio
async def test_callback_error_reaches_the_caller(two_model_router, monkeypatch):
_register(monkeypatch, _RaisingGate())
with pytest.raises(HTTPException) as raised:
await _v1_models(_non_admin())
assert raised.value.status_code == 503
@pytest.mark.asyncio
async def test_hidden_model_still_routes_for_direct_requests(two_model_router, monkeypatch):
_register(monkeypatch, _Gate(hidden=frozenset({"restricted-model"})))
assert "restricted-model" not in await _v1_models(_non_admin())
deployment = two_model_router.get_available_deployment(
model="restricted-model", messages=[{"role": "user", "content": "hi"}]
)
assert deployment["model_name"] == "restricted-model"
@pytest.mark.asyncio
async def test_model_group_info_offers_a2a_agent_groups_to_the_callback(two_model_router, monkeypatch):
from litellm.proxy._types import LiteLLM_ObjectPermissionTable
from litellm.proxy.agent_endpoints.agent_registry import global_agent_registry
from litellm.types.agents import AgentResponse
monkeypatch.setattr(
global_agent_registry,
"agent_list",
[AgentResponse(agent_id="agent-1", agent_name="helper", agent_card_params={})],
)
caller = _non_admin(object_permission=LiteLLM_ObjectPermissionTable(object_permission_id="p1", agents=["agent-1"]))
gate: _Gate = _Gate()
_register(monkeypatch, gate)
assert await _model_groups(caller) == ["open-model", "restricted-model", "a2a/helper"]
assert gate.seen == [("open-model", "restricted-model", "a2a/helper")]
_register(monkeypatch, _Gate(hidden=frozenset({"a2a/helper", "restricted-model"})))
assert await _model_groups(caller) == ["open-model"]
async def _v1_model_info_by_deployment_id(deployment_id: str, user_api_key_dict: UserAPIKeyAuth) -> int | list[str]:
try:
response = await proxy_server.model_info_v1(user_api_key_dict=user_api_key_dict, litellm_model_id=deployment_id)
except HTTPException as error:
return error.status_code
return [row["model_name"] for row in json.loads(response.body)["data"]]
@pytest.mark.asyncio
async def test_v1_model_info_by_deployment_id_answers_like_an_unknown_id_for_a_hidden_model(
two_model_router, monkeypatch
):
_register(monkeypatch, _Gate(hidden=frozenset({"restricted-model"})))
assert await _v1_model_info_by_deployment_id("restricted-model-id", _non_admin()) == 400
assert await _v1_model_info_by_deployment_id("no-such-id", _non_admin()) == 400
assert await _v1_model_info_by_deployment_id("open-model-id", _non_admin()) == ["open-model"]
@pytest.mark.asyncio
async def test_v1_model_info_by_deployment_id_offers_the_public_team_name(team_router, monkeypatch):
gate: _Gate = _Gate(hidden=frozenset({"team-gpt"}))
_register(monkeypatch, gate)
assert await _v1_model_info_by_deployment_id("model_name_team1_abc-id", _team_member()) == 400
assert await _v1_model_info_by_deployment_id("model_name_team1_def-id", _team_member()) == ["team-chat"]
assert gate.seen == [("team-gpt",), ("team-chat",)]
@pytest.mark.asyncio
async def test_v1_model_info_by_deployment_id_offers_the_name_its_listing_shows_in_legacy_mode(
team_router, monkeypatch
):
monkeypatch.setattr(proxy_server, "general_settings", {"use_team_public_model_name": False})
gate: _Gate = _Gate(hidden=frozenset({"team-gpt"}))
_register(monkeypatch, gate)
assert await _v1_model_info_names(_team_member()) == ["team-chat"]
assert await _v1_model_info_by_deployment_id("model_name_team1_abc-id", _team_member()) == 400
assert gate.seen[-1] == ("team-gpt",)