mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-26 01:12:21 +00:00
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:
parent
bf0187072b
commit
7b25a151bd
5 changed files with 596 additions and 23 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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]:
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
425
tests/test_litellm/proxy/test_model_list_callback_filter.py
Normal file
425
tests/test_litellm/proxy/test_model_list_callback_filter.py
Normal 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",)
|
||||
Loading…
Add table
Reference in a new issue