mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-12 23:01:41 +00:00
fix(proxy): preserve BYOK listing metadata
This commit is contained in:
parent
e1a3fca8fa
commit
9d78277a0f
3 changed files with 196 additions and 96 deletions
101
litellm/proxy/common_utils/model_listing_utils.py
Normal file
101
litellm/proxy/common_utils/model_listing_utils.py
Normal file
|
|
@ -0,0 +1,101 @@
|
|||
from collections.abc import Mapping
|
||||
from typing import TYPE_CHECKING, cast
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.router import Router
|
||||
|
||||
|
||||
def get_model_listing_entries(
|
||||
model_names: list[str],
|
||||
llm_router: "Router | None",
|
||||
general_settings: object,
|
||||
) -> tuple[tuple[str, str], ...]:
|
||||
"""Build `/v1/models` entries while keeping router lookup keys intact.
|
||||
|
||||
Team-scoped BYOK deployments use internal router keys like
|
||||
`model_name_{team_id}_{uuid}`. The listing response should surface the
|
||||
public model name, but metadata lookups such as fallbacks still need the
|
||||
internal key because the router indexes those configs by routing key.
|
||||
"""
|
||||
if not _should_use_team_public_model_name(general_settings) or llm_router is None:
|
||||
return _default_model_listing_entries(model_names)
|
||||
|
||||
router_model_list: object = llm_router.get_model_list()
|
||||
if not isinstance(router_model_list, list):
|
||||
return _default_model_listing_entries(model_names)
|
||||
router_models = cast(list[object], router_model_list) # any-ok: checked
|
||||
|
||||
team_name_pairs = tuple(
|
||||
pair
|
||||
for model in router_models
|
||||
for pair in (_team_public_name_pair(model),)
|
||||
if pair is not None
|
||||
)
|
||||
if not team_name_pairs:
|
||||
return _default_model_listing_entries(model_names)
|
||||
|
||||
return _dedupe_model_listing_entries(
|
||||
tuple(
|
||||
_model_listing_entry(model_name, team_name_pairs)
|
||||
for model_name in model_names
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
def _default_model_listing_entries(
|
||||
model_names: list[str],
|
||||
) -> tuple[tuple[str, str], ...]:
|
||||
return tuple((model_name, model_name) for model_name in model_names)
|
||||
|
||||
|
||||
def _should_use_team_public_model_name(general_settings: object) -> bool:
|
||||
if not isinstance(general_settings, Mapping):
|
||||
return True
|
||||
settings = cast(Mapping[str, object], general_settings) # any-ok: checked
|
||||
return settings.get("use_team_public_model_name", True) is not False
|
||||
|
||||
|
||||
def _team_public_name_pair(model: object) -> tuple[str, str] | None:
|
||||
if not isinstance(model, Mapping):
|
||||
return None
|
||||
model_dict = cast(Mapping[str, object], model) # any-ok: checked
|
||||
model_info_raw: object = model_dict.get("model_info")
|
||||
if not isinstance(model_info_raw, Mapping):
|
||||
return None
|
||||
|
||||
model_info = cast(Mapping[str, object], model_info_raw) # any-ok: checked
|
||||
team_id = model_info.get("team_id")
|
||||
team_public = model_info.get("team_public_model_name")
|
||||
name = model_dict.get("model_name")
|
||||
if (
|
||||
isinstance(team_id, str)
|
||||
and isinstance(team_public, str)
|
||||
and isinstance(name, str)
|
||||
and name.startswith(f"model_name_{team_id}_")
|
||||
):
|
||||
return name, team_public
|
||||
return None
|
||||
|
||||
|
||||
def _model_listing_entry(
|
||||
model_name: str,
|
||||
team_name_pairs: tuple[tuple[str, str], ...],
|
||||
) -> tuple[str, str]:
|
||||
for internal_name, public_name in team_name_pairs:
|
||||
if internal_name == model_name:
|
||||
return public_name, internal_name
|
||||
return model_name, model_name
|
||||
|
||||
|
||||
def _dedupe_model_listing_entries(
|
||||
entries: tuple[tuple[str, str], ...],
|
||||
) -> tuple[tuple[str, str], ...]:
|
||||
deduped_entries: list[tuple[str, str]] = []
|
||||
seen_response_ids: set[str] = set()
|
||||
for entry in entries:
|
||||
response_model_id: str = entry[0]
|
||||
if response_model_id in seen_response_ids:
|
||||
continue
|
||||
seen_response_ids.add(response_model_id)
|
||||
deduped_entries.append(entry)
|
||||
return tuple(deduped_entries)
|
||||
|
|
@ -15,7 +15,6 @@ import threading
|
|||
import time
|
||||
import traceback
|
||||
import warnings
|
||||
from collections.abc import Mapping
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from typing import (
|
||||
TYPE_CHECKING,
|
||||
|
|
@ -302,6 +301,9 @@ from litellm.proxy.common_utils.load_config_utils import (
|
|||
get_config_file_contents_from_gcs,
|
||||
get_file_contents_from_s3,
|
||||
)
|
||||
from litellm.proxy.common_utils.model_listing_utils import (
|
||||
get_model_listing_entries,
|
||||
)
|
||||
from litellm.proxy.common_utils.openai_endpoint_utils import (
|
||||
remove_sensitive_info_from_deployment,
|
||||
)
|
||||
|
|
@ -8273,6 +8275,8 @@ async def model_list(
|
|||
get_available_models_for_user,
|
||||
)
|
||||
|
||||
proxy_general_settings: object = general_settings # any-ok: legacy settings
|
||||
|
||||
# Validate scope parameter if provided
|
||||
if scope is not None and scope != "expand":
|
||||
raise HTTPException(
|
||||
|
|
@ -8344,20 +8348,28 @@ async def model_list(
|
|||
if hidden_names:
|
||||
all_models = [m for m in all_models if m not in hidden_names]
|
||||
|
||||
# Surface the public name for team-scoped rows by default. Operators
|
||||
# that need legacy internal routing keys can explicitly disable this.
|
||||
all_models = _translate_team_model_names_for_listing(all_models, llm_router)
|
||||
model_entries = get_model_listing_entries(
|
||||
model_names=all_models,
|
||||
llm_router=llm_router,
|
||||
general_settings=proxy_general_settings,
|
||||
)
|
||||
|
||||
# Build response data with all proxy models
|
||||
model_data = []
|
||||
for model in all_models:
|
||||
model_info = create_model_info_response(
|
||||
model_id=model,
|
||||
provider="openai",
|
||||
include_metadata=include_metadata or False,
|
||||
fallback_type=fallback_type,
|
||||
llm_router=llm_router,
|
||||
model_data: list[dict[str, object]] = []
|
||||
for response_model_id, metadata_lookup_model_id in model_entries:
|
||||
model_info_raw: object = (
|
||||
create_model_info_response( # any-ok: legacy response helper
|
||||
model_id=metadata_lookup_model_id,
|
||||
provider="openai",
|
||||
include_metadata=include_metadata or False,
|
||||
fallback_type=fallback_type,
|
||||
llm_router=llm_router,
|
||||
)
|
||||
)
|
||||
model_info = cast(
|
||||
dict[str, object], model_info_raw
|
||||
) # any-ok: legacy response helper
|
||||
model_info["id"] = response_model_id
|
||||
model_data.append(model_info)
|
||||
|
||||
return dict(
|
||||
|
|
@ -8385,20 +8397,28 @@ async def model_list(
|
|||
if hidden_names:
|
||||
all_models = [m for m in all_models if m not in hidden_names]
|
||||
|
||||
# Surface the public name for team-scoped rows by default. Operators that
|
||||
# need legacy internal routing keys can explicitly disable this.
|
||||
all_models = _translate_team_model_names_for_listing(all_models, llm_router)
|
||||
model_entries = get_model_listing_entries(
|
||||
model_names=all_models,
|
||||
llm_router=llm_router,
|
||||
general_settings=proxy_general_settings,
|
||||
)
|
||||
|
||||
# Build response data
|
||||
model_data = []
|
||||
for model in all_models:
|
||||
model_info = create_model_info_response(
|
||||
model_id=model,
|
||||
provider="openai",
|
||||
include_metadata=include_metadata or False,
|
||||
fallback_type=fallback_type,
|
||||
llm_router=llm_router,
|
||||
model_data: list[dict[str, object]] = [] # any-ok: typed response list
|
||||
for response_model_id, metadata_lookup_model_id in model_entries:
|
||||
model_info_raw: object = (
|
||||
create_model_info_response( # any-ok: legacy response helper
|
||||
model_id=metadata_lookup_model_id,
|
||||
provider="openai",
|
||||
include_metadata=include_metadata or False,
|
||||
fallback_type=fallback_type,
|
||||
llm_router=llm_router,
|
||||
)
|
||||
)
|
||||
model_info = cast(
|
||||
dict[str, object], model_info_raw
|
||||
) # any-ok: legacy response helper
|
||||
model_info["id"] = response_model_id
|
||||
model_data.append(model_info)
|
||||
|
||||
return dict(
|
||||
|
|
@ -12598,60 +12618,6 @@ def _translate_model_name_for_response(model: dict) -> dict:
|
|||
return {**model, "model_name": team_public}
|
||||
|
||||
|
||||
def _translate_team_model_names_for_listing(
|
||||
model_names: list[str],
|
||||
llm_router: Router | None,
|
||||
) -> list[str]:
|
||||
"""Swap internal team routing keys for their public names in list-style
|
||||
responses (e.g. `/v1/models`, `/models`).
|
||||
|
||||
`/v1/models` builds from bare model-name strings produced by access-group
|
||||
expansion (`get_model_access_groups`), which surfaces the internal routing
|
||||
key `model_name_{team_id}_{uuid}` for team-scoped (BYOK) deployments. This
|
||||
is a presentation-layer swap only -- access-group/auth semantics are
|
||||
unchanged (see issue #28382). Sibling deployments collapse to one public
|
||||
name, so the result is de-duplicated while preserving order.
|
||||
"""
|
||||
settings = cast(dict[str, object], general_settings) # any-ok: legacy settings
|
||||
if settings.get("use_team_public_model_name", True) is False or llm_router is None:
|
||||
return model_names
|
||||
|
||||
router_model_list = llm_router.get_model_list()
|
||||
if not isinstance(router_model_list, list):
|
||||
return model_names
|
||||
|
||||
internal_to_public: dict[str, str] = {}
|
||||
for model in router_model_list:
|
||||
if not isinstance(model, dict):
|
||||
continue
|
||||
model_dict = cast(dict[str, object], model) # any-ok: checked
|
||||
model_info_raw: object = model_dict.get("model_info")
|
||||
if not isinstance(model_info_raw, Mapping):
|
||||
continue
|
||||
model_info = cast(Mapping[str, object], model_info_raw) # any-ok: checked
|
||||
team_id = model_info.get("team_id")
|
||||
team_public = model_info.get("team_public_model_name")
|
||||
name = model_dict.get("model_name")
|
||||
if (
|
||||
isinstance(team_id, str)
|
||||
and isinstance(team_public, str)
|
||||
and isinstance(name, str)
|
||||
and name.startswith(f"model_name_{team_id}_")
|
||||
):
|
||||
internal_to_public[name] = team_public
|
||||
if not internal_to_public:
|
||||
return model_names
|
||||
translated_names: list[str] = []
|
||||
seen_names: set[str] = set()
|
||||
for model_name in model_names:
|
||||
translated_name = internal_to_public.get(model_name, model_name)
|
||||
if translated_name in seen_names:
|
||||
continue
|
||||
seen_names.add(translated_name)
|
||||
translated_names.append(translated_name)
|
||||
return translated_names
|
||||
|
||||
|
||||
def _get_proxy_model_info(model: dict) -> dict:
|
||||
# provided model_info in config.yaml
|
||||
model_info = model.get("model_info", {})
|
||||
|
|
|
|||
|
|
@ -15,6 +15,7 @@ import pytest
|
|||
|
||||
import litellm.proxy.proxy_server as ps
|
||||
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
|
||||
from litellm.proxy.common_utils.model_listing_utils import get_model_listing_entries
|
||||
from litellm.proxy.proxy_server import (
|
||||
_get_proxy_model_info,
|
||||
_translate_model_name_for_response,
|
||||
|
|
@ -678,7 +679,8 @@ async def test_v1_models_keeps_internal_names_when_public_name_flag_disabled(
|
|||
|
||||
@pytest.mark.asyncio
|
||||
async def test_v1_models_translates_team_model_with_metadata(monkeypatch):
|
||||
"""include_metadata=true must build metadata for the public model id."""
|
||||
"""include_metadata=true must return the public id while resolving
|
||||
fallback metadata through the internal router key."""
|
||||
team_dep = {
|
||||
"model_name": "model_name_teamX_uuid9",
|
||||
"litellm_params": {"model": "azure/gpt-4.1"},
|
||||
|
|
@ -695,6 +697,7 @@ async def test_v1_models_translates_team_model_with_metadata(monkeypatch):
|
|||
router.get_fully_blocked_model_names.return_value = set()
|
||||
router.model_list = [team_dep]
|
||||
router.get_model_list.return_value = [team_dep]
|
||||
router.fallbacks = [{"model_name_teamX_uuid9": ["fallback-gpt-4.1"]}]
|
||||
|
||||
monkeypatch.setattr(ps, "llm_router", router)
|
||||
monkeypatch.setattr(ps, "user_model", None)
|
||||
|
|
@ -711,15 +714,14 @@ async def test_v1_models_translates_team_model_with_metadata(monkeypatch):
|
|||
"object": "model",
|
||||
"created": 1677610602,
|
||||
"owned_by": "openai",
|
||||
"metadata": {"fallbacks": []},
|
||||
"metadata": {"fallbacks": ["fallback-gpt-4.1"]},
|
||||
}
|
||||
]
|
||||
|
||||
|
||||
def test_translate_team_model_names_for_listing_swaps_and_dedupes(monkeypatch):
|
||||
def test_model_listing_entries_swaps_and_dedupes(monkeypatch):
|
||||
"""Internal team routing keys -> public name; sibling deployments sharing a
|
||||
public name collapse to one entry (order preserved); globals untouched."""
|
||||
from litellm.proxy.proxy_server import _translate_team_model_names_for_listing
|
||||
|
||||
router = MagicMock()
|
||||
router.get_model_list.return_value = [
|
||||
|
|
@ -741,16 +743,19 @@ def test_translate_team_model_names_for_listing_swaps_and_dedupes(monkeypatch):
|
|||
]
|
||||
monkeypatch.setattr(ps, "general_settings", {})
|
||||
|
||||
out = _translate_team_model_names_for_listing(
|
||||
out = get_model_listing_entries(
|
||||
["model_name_teamX_uuidA", "model_name_teamX_uuidB", "gpt-4o"],
|
||||
router,
|
||||
{},
|
||||
)
|
||||
assert out == ["tushar-gpt-4.1", "gpt-4o"]
|
||||
assert list(out) == [
|
||||
("tushar-gpt-4.1", "model_name_teamX_uuidA"),
|
||||
("gpt-4o", "gpt-4o"),
|
||||
]
|
||||
|
||||
|
||||
def test_translate_team_model_names_for_listing_leaves_unmapped_names(monkeypatch):
|
||||
def test_model_listing_entries_leave_unmapped_names(monkeypatch):
|
||||
"""Names with no team mapping (globals, access-group keys) pass through."""
|
||||
from litellm.proxy.proxy_server import _translate_team_model_names_for_listing
|
||||
|
||||
router = MagicMock()
|
||||
router.get_model_list.return_value = [
|
||||
|
|
@ -758,22 +763,26 @@ def test_translate_team_model_names_for_listing_leaves_unmapped_names(monkeypatc
|
|||
]
|
||||
monkeypatch.setattr(ps, "general_settings", {})
|
||||
|
||||
assert _translate_team_model_names_for_listing(
|
||||
["gpt-4o", "beta-group"], router
|
||||
) == ["gpt-4o", "beta-group"]
|
||||
out = get_model_listing_entries(["gpt-4o", "beta-group"], router, {})
|
||||
assert list(out) == [
|
||||
("gpt-4o", "gpt-4o"),
|
||||
("beta-group", "beta-group"),
|
||||
]
|
||||
|
||||
|
||||
def test_translate_team_model_names_for_listing_none_router(monkeypatch):
|
||||
def test_model_listing_entries_handle_none_router(monkeypatch):
|
||||
"""No router -> return the input list unchanged."""
|
||||
from litellm.proxy.proxy_server import _translate_team_model_names_for_listing
|
||||
|
||||
monkeypatch.setattr(ps, "general_settings", {})
|
||||
assert _translate_team_model_names_for_listing(["a", "b"], None) == ["a", "b"]
|
||||
out = get_model_listing_entries(["a", "b"], None, {})
|
||||
assert list(out) == [
|
||||
("a", "a"),
|
||||
("b", "b"),
|
||||
]
|
||||
|
||||
|
||||
def test_translate_team_model_names_for_listing_respects_legacy_flag(monkeypatch):
|
||||
def test_model_listing_entries_respect_legacy_flag(monkeypatch):
|
||||
"""Operators can keep returning the legacy internal routing key."""
|
||||
from litellm.proxy.proxy_server import _translate_team_model_names_for_listing
|
||||
|
||||
router = MagicMock()
|
||||
router.get_model_list.return_value = [
|
||||
|
|
@ -787,6 +796,30 @@ def test_translate_team_model_names_for_listing_respects_legacy_flag(monkeypatch
|
|||
]
|
||||
monkeypatch.setattr(ps, "general_settings", {"use_team_public_model_name": False})
|
||||
|
||||
assert _translate_team_model_names_for_listing(
|
||||
["model_name_teamX_uuidA"], router
|
||||
) == ["model_name_teamX_uuidA"]
|
||||
out = get_model_listing_entries(
|
||||
["model_name_teamX_uuidA"],
|
||||
router,
|
||||
{"use_team_public_model_name": False},
|
||||
)
|
||||
assert list(out) == [("model_name_teamX_uuidA", "model_name_teamX_uuidA")]
|
||||
|
||||
|
||||
def test_model_listing_entries_handles_unexpected_router_model_list():
|
||||
router = MagicMock()
|
||||
router.get_model_list.return_value = {"not": "a list"}
|
||||
|
||||
out = get_model_listing_entries(["gpt-4o"], router, None)
|
||||
|
||||
assert list(out) == [("gpt-4o", "gpt-4o")]
|
||||
|
||||
|
||||
def test_model_listing_entries_ignores_malformed_router_rows():
|
||||
router = MagicMock()
|
||||
router.get_model_list.return_value = [
|
||||
object(),
|
||||
{"model_name": "model_name_teamX_uuidA", "model_info": "not-a-dict"},
|
||||
]
|
||||
|
||||
out = get_model_listing_entries(["model_name_teamX_uuidA"], router, {})
|
||||
|
||||
assert list(out) == [("model_name_teamX_uuidA", "model_name_teamX_uuidA")]
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue