mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
fix(router): keep pod routing records accurate
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
afb68249ba
commit
a0966ceede
12 changed files with 268 additions and 105 deletions
|
|
@ -49,6 +49,7 @@ ROUTER_SETTINGS_MANAGED_OUTSIDE_CONFIG: Final[frozenset[str]] = frozenset(
|
|||
"fallback_access_check",
|
||||
"fallback_budget_check",
|
||||
"auto_router_capability_limit",
|
||||
"kubernetes_pod_discovery",
|
||||
}
|
||||
)
|
||||
DEFAULT_BATCH_SIZE: Final = int(os.getenv("DEFAULT_BATCH_SIZE", 512))
|
||||
|
|
|
|||
|
|
@ -6999,7 +6999,7 @@ class ProxyConfig:
|
|||
"Key '%s' is not a valid argument for Router.__init__(). Ignoring this key.", k
|
||||
)
|
||||
router = litellm.Router(
|
||||
**router_params,
|
||||
**router_params, # pyright: ignore[reportUnknownArgumentType] # router settings are filtered at runtime
|
||||
assistants_config=assistants_config,
|
||||
search_tools=search_tools,
|
||||
router_general_settings=RouterGeneralSettings(
|
||||
|
|
|
|||
|
|
@ -45,9 +45,13 @@ def _route_user_config_request(data: dict, route_type: str):
|
|||
# Filter router_config to only include valid Router.__init__ arguments
|
||||
# This prevents TypeError when invalid parameters are stored in the database
|
||||
valid_args: Final = litellm.Router.get_valid_args()
|
||||
filtered_config: Final = {k: v for k, v in router_config.items() if k in valid_args}
|
||||
filtered_config: Final = {
|
||||
k: v for k, v in router_config.items() if k in valid_args and k != "kubernetes_pod_discovery"
|
||||
}
|
||||
|
||||
user_router: Final = litellm.Router(**filtered_config)
|
||||
user_router: Final = litellm.Router(
|
||||
**filtered_config, # pyright: ignore[reportUnknownArgumentType] # config is filtered by get_valid_args
|
||||
)
|
||||
ret_val: Final = getattr(user_router, f"{route_type}")(**data)
|
||||
user_router.discard()
|
||||
return ret_val
|
||||
|
|
|
|||
|
|
@ -843,6 +843,7 @@ class Router:
|
|||
fallback_access_check: FallbackAccessCheck | None = None,
|
||||
fallback_budget_check: FallbackBudgetCheck | None = None,
|
||||
auto_router_capability_limit: AutoRouterCapabilityLimit | None = None,
|
||||
kubernetes_pod_discovery: KubernetesPodDiscovery | None = None,
|
||||
) -> None:
|
||||
"""
|
||||
Initialize the Router class with the given parameters for caching, reliability, and routing strategy.
|
||||
|
|
@ -952,7 +953,9 @@ class Router:
|
|||
cache_config: Final[dict[str, Any]] = {}
|
||||
|
||||
self.client_ttl = client_ttl
|
||||
self.kubernetes_pod_discovery: Final = KubernetesPodDiscovery()
|
||||
self.kubernetes_pod_discovery: Final = (
|
||||
kubernetes_pod_discovery if kubernetes_pod_discovery is not None else KubernetesPodDiscovery()
|
||||
)
|
||||
if redis_url is not None or (redis_host is not None and redis_port is not None):
|
||||
cache_type = "redis"
|
||||
|
||||
|
|
@ -1250,7 +1253,8 @@ class Router:
|
|||
"""
|
||||
Returns a list of valid arguments for the Router.__init__ method.
|
||||
"""
|
||||
arg_spec: Final = inspect.getfullargspec(Router.__init__)
|
||||
router_init: Final[Callable[..., None]] = Router.__init__
|
||||
arg_spec: Final = inspect.getfullargspec(router_init)
|
||||
valid_args: Final = arg_spec.args + arg_spec.kwonlyargs
|
||||
if "self" in valid_args:
|
||||
valid_args.remove("self")
|
||||
|
|
@ -4021,6 +4025,7 @@ class Router:
|
|||
metadata_variable_name: Final = _get_router_metadata_variable_name(
|
||||
function_name=function_name,
|
||||
)
|
||||
routing: Final = deployment.get(KUBERNETES_POD_ROUTING_KEY)
|
||||
|
||||
kwargs.setdefault(metadata_variable_name, {}).update(
|
||||
{
|
||||
|
|
@ -4028,9 +4033,16 @@ class Router:
|
|||
"model_info": model_info,
|
||||
"api_base": deployment_api_base,
|
||||
"deployment_model_name": deployment_model_name,
|
||||
KUBERNETES_POD_ROUTING_KEY: deployment.get(KUBERNETES_POD_ROUTING_KEY),
|
||||
KUBERNETES_POD_ROUTING_KEY: routing,
|
||||
}
|
||||
)
|
||||
other_bucket: Final = "metadata" if metadata_variable_name == "litellm_metadata" else "litellm_metadata"
|
||||
other_bucket_metadata: Final = kwargs.get(other_bucket)
|
||||
if isinstance(other_bucket_metadata, Mapping) and KUBERNETES_POD_ROUTING_KEY in other_bucket_metadata:
|
||||
kwargs[other_bucket] = {
|
||||
**other_bucket_metadata,
|
||||
KUBERNETES_POD_ROUTING_KEY: routing,
|
||||
}
|
||||
|
||||
# A retry/fallback reuses this same kwargs dict for the next deployment.
|
||||
# Refund and clear any reservation the previous deployment attempt left
|
||||
|
|
|
|||
|
|
@ -297,7 +297,7 @@ async def test_aresponses_with_streaming_fallbacks_wraps_streaming_iterator():
|
|||
# -------- every fallback entry stays reachable across hops --------
|
||||
|
||||
|
||||
def _make_three_tier_router(**router_kwargs) -> Router:
|
||||
def _make_three_tier_router(fallbacks: list[dict[str, list[str]]] | None = None) -> Router:
|
||||
return Router(
|
||||
model_list=[
|
||||
{"model_name": "primary", "litellm_params": {"model": "openai/primary-model", "api_key": "sk-test"}},
|
||||
|
|
@ -305,7 +305,7 @@ def _make_three_tier_router(**router_kwargs) -> Router:
|
|||
{"model_name": "fb2", "litellm_params": {"model": "openai/fb2-model", "api_key": "sk-test"}},
|
||||
],
|
||||
num_retries=0,
|
||||
**router_kwargs,
|
||||
fallbacks=[] if fallbacks is None else fallbacks,
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -25,8 +25,16 @@ from litellm.router_utils.fallback_event_handlers import _trigger_cooldown_for_f
|
|||
from litellm.types.router import AllowedFailsPolicy
|
||||
|
||||
|
||||
def _make_router(model_list: list, **kwargs) -> Router:
|
||||
return Router(model_list=model_list, **kwargs)
|
||||
def _make_router(
|
||||
model_list: list,
|
||||
allowed_fails: int | None = None,
|
||||
allowed_fails_policy: AllowedFailsPolicy | None = None,
|
||||
) -> Router:
|
||||
return Router(
|
||||
model_list=model_list,
|
||||
allowed_fails=allowed_fails,
|
||||
allowed_fails_policy=allowed_fails_policy,
|
||||
)
|
||||
|
||||
|
||||
class TestDeploymentLevelAllowedFails:
|
||||
|
|
|
|||
|
|
@ -12,6 +12,7 @@ from litellm.proxy.common_utils.discoverable_model_filter import (
|
|||
discoverable_rows,
|
||||
undiscoverable_model_names,
|
||||
)
|
||||
from litellm.types.router import RouterModelGroupAliasItem
|
||||
|
||||
|
||||
def _deployment(model_name: str, model: str = "openai/gpt-4o", **model_info):
|
||||
|
|
@ -22,8 +23,11 @@ def _deployment(model_name: str, model: str = "openai/gpt-4o", **model_info):
|
|||
}
|
||||
|
||||
|
||||
def _router(*deployments, **router_kwargs) -> Router:
|
||||
return Router(model_list=list(deployments), **router_kwargs)
|
||||
def _router(
|
||||
*deployments,
|
||||
model_group_alias: dict[str, str | RouterModelGroupAliasItem] | None = None,
|
||||
) -> Router:
|
||||
return Router(model_list=list(deployments), model_group_alias=model_group_alias)
|
||||
|
||||
|
||||
def _non_admin() -> UserAPIKeyAuth:
|
||||
|
|
|
|||
|
|
@ -17,6 +17,7 @@ from litellm.proxy.guardrails.guardrail_hooks.llm_as_a_judge import (
|
|||
initialize_guardrail,
|
||||
)
|
||||
from litellm.types.guardrails import GuardrailEventHooks, Mode
|
||||
from litellm.types.router import RouterModelGroupAliasItem
|
||||
from litellm.types.utils import LLM_AS_A_JUDGE_GUARDRAIL_CALL_ORIGIN
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
|
|
@ -598,19 +599,22 @@ def _judge_response_mock() -> MagicMock:
|
|||
return MagicMock(choices=[MagicMock(message=MagicMock(content=json.dumps(_make_verdict_response(90.0))))])
|
||||
|
||||
|
||||
def _real_router(model_list, **router_kwargs):
|
||||
def _real_router(
|
||||
model_list,
|
||||
model_group_alias: dict[str, str | RouterModelGroupAliasItem] | None = None,
|
||||
):
|
||||
"""Build a real Router so the router-membership decision is exercised for
|
||||
real (wildcards, model_group_alias, exact names), stubbing only the outbound
|
||||
completion so no network call is made."""
|
||||
from litellm import Router
|
||||
|
||||
router = Router(model_list=model_list, **router_kwargs)
|
||||
router = Router(model_list=model_list, model_group_alias=model_group_alias)
|
||||
router.acompletion = AsyncMock(return_value=_judge_response_mock())
|
||||
return router
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"model_list, router_kwargs, judge_model",
|
||||
"model_list, model_group_alias, judge_model",
|
||||
[
|
||||
(
|
||||
[
|
||||
|
|
@ -619,12 +623,12 @@ def _real_router(model_list, **router_kwargs):
|
|||
"litellm_params": {"model": "anthropic/claude-sonnet-4-6", "api_key": "sk-ant-test"},
|
||||
}
|
||||
],
|
||||
{},
|
||||
None,
|
||||
"my-judge-alias",
|
||||
),
|
||||
(
|
||||
[{"model_name": "anthropic/*", "litellm_params": {"model": "anthropic/*", "api_key": "sk-ant-test"}}],
|
||||
{},
|
||||
None,
|
||||
"anthropic/claude-sonnet-4-6",
|
||||
),
|
||||
(
|
||||
|
|
@ -634,7 +638,7 @@ def _real_router(model_list, **router_kwargs):
|
|||
"litellm_params": {"model": "anthropic/claude-sonnet-4-6", "api_key": "sk-ant-test"},
|
||||
}
|
||||
],
|
||||
{"model_group_alias": {"my-judge-alias": "backing-group"}},
|
||||
{"my-judge-alias": "backing-group"},
|
||||
"my-judge-alias",
|
||||
),
|
||||
(
|
||||
|
|
@ -644,7 +648,7 @@ def _real_router(model_list, **router_kwargs):
|
|||
"litellm_params": {"model": "anthropic/claude-sonnet-4-6", "api_key": "sk-ant-test"},
|
||||
}
|
||||
],
|
||||
{"model_group_alias": {"my-judge-alias": {"model": "backing-group", "hidden": True}}},
|
||||
{"my-judge-alias": {"model": "backing-group", "hidden": True}},
|
||||
"my-judge-alias",
|
||||
),
|
||||
],
|
||||
|
|
@ -653,13 +657,16 @@ def _real_router(model_list, **router_kwargs):
|
|||
@pytest.mark.asyncio
|
||||
@patch("litellm.proxy.guardrails.guardrail_hooks.llm_as_a_judge.litellm.acompletion", new_callable=AsyncMock)
|
||||
async def test_judge_routes_through_router_for_router_served_model(
|
||||
mock_sdk_completion, model_list, router_kwargs, judge_model
|
||||
mock_sdk_completion,
|
||||
model_list,
|
||||
model_group_alias: dict[str, str | RouterModelGroupAliasItem] | None,
|
||||
judge_model,
|
||||
):
|
||||
"""Any judge_model the Router can serve must resolve its credentials via the
|
||||
Router. Wildcard and alias shapes regress the naive `judge_model in
|
||||
get_model_names()` check, which reports patterns/aliases literally and so
|
||||
routes a servable model to the SDK, where deployment creds do not resolve."""
|
||||
router = _real_router(model_list, **router_kwargs)
|
||||
router = _real_router(model_list, model_group_alias)
|
||||
guardrail = _make_guardrail(judge_model=judge_model, router_provider=lambda: router)
|
||||
inputs = {"texts": ["good response"]}
|
||||
request_data: dict = {"messages": [{"role": "user", "content": "hi"}], "metadata": {}}
|
||||
|
|
|
|||
|
|
@ -20,6 +20,7 @@ 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
|
||||
from litellm.types.router import RouterModelGroupAliasItem
|
||||
|
||||
|
||||
class _Gate(CustomLogger):
|
||||
|
|
@ -70,8 +71,12 @@ def _deployment(model_name: str, model: str = "openai/gpt-4o", **model_info):
|
|||
}
|
||||
|
||||
|
||||
def _install_router(monkeypatch, *deployments, **router_kwargs) -> Router:
|
||||
router = Router(model_list=list(deployments), **router_kwargs)
|
||||
def _install_router(
|
||||
monkeypatch,
|
||||
*deployments,
|
||||
model_group_alias: dict[str, str | RouterModelGroupAliasItem] | None = None,
|
||||
) -> Router:
|
||||
router = Router(model_list=list(deployments), model_group_alias=model_group_alias)
|
||||
monkeypatch.setattr(proxy_server, "llm_router", router)
|
||||
monkeypatch.setattr(proxy_server, "llm_model_list", router.model_list)
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", None)
|
||||
|
|
|
|||
|
|
@ -100,31 +100,6 @@ def _empty_proxy_environment() -> Mapping[str, str]:
|
|||
return {}
|
||||
|
||||
|
||||
def _patch_router_pod_discovery(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
*,
|
||||
clock: Callable[[], float] | None = None,
|
||||
refresh_interval_seconds: float | None = None,
|
||||
) -> None:
|
||||
def create_discovery() -> KubernetesPodDiscovery:
|
||||
if refresh_interval_seconds is not None and clock is not None:
|
||||
return KubernetesPodDiscovery(
|
||||
refresh_interval_seconds=refresh_interval_seconds,
|
||||
clock=clock,
|
||||
proxy_environment=_empty_proxy_environment,
|
||||
)
|
||||
if refresh_interval_seconds is not None:
|
||||
return KubernetesPodDiscovery(
|
||||
refresh_interval_seconds=refresh_interval_seconds,
|
||||
proxy_environment=_empty_proxy_environment,
|
||||
)
|
||||
if clock is not None:
|
||||
return KubernetesPodDiscovery(clock=clock, proxy_environment=_empty_proxy_environment)
|
||||
return KubernetesPodDiscovery(proxy_environment=_empty_proxy_environment)
|
||||
|
||||
monkeypatch.setattr("litellm.router.KubernetesPodDiscovery", create_discovery)
|
||||
|
||||
|
||||
def _stub_sync_dns(monkeypatch: pytest.MonkeyPatch, *ips: str) -> None:
|
||||
def getaddrinfo(
|
||||
host: str | None,
|
||||
|
|
@ -775,16 +750,20 @@ async def test_router_sends_pod_hosts_without_forwarding_discovery_flag(
|
|||
monkeypatch.setattr(loop, "getaddrinfo", async_getaddrinfo)
|
||||
monkeypatch.setattr(socket, "getaddrinfo", sync_getaddrinfo)
|
||||
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
|
||||
_patch_router_pod_discovery(monkeypatch, clock=clock)
|
||||
|
||||
with respx.mock(assert_all_called=True) as respx_mock:
|
||||
route: Final = respx_mock.post(
|
||||
url__regex=re.compile(
|
||||
r"http://(?:10\.0\.0\.1|10\.0\.0\.2|vllm-headless\.ns\.svc\.cluster\.local):8000/v1/chat/completions"
|
||||
)
|
||||
).mock(return_value=httpx.Response(200, json=_CHAT_RESPONSE))
|
||||
discovery_enabled_router: Final = Router(model_list=[_deployment()])
|
||||
control_router: Final = Router(model_list=[_deployment(kubernetes_pod_discovery=None)])
|
||||
discovery_enabled_router: Final = Router(
|
||||
model_list=[_deployment()],
|
||||
kubernetes_pod_discovery=KubernetesPodDiscovery(clock=clock, proxy_environment=_empty_proxy_environment),
|
||||
)
|
||||
control_router: Final = Router(
|
||||
model_list=[_deployment(kubernetes_pod_discovery=None)],
|
||||
kubernetes_pod_discovery=KubernetesPodDiscovery(clock=clock, proxy_environment=_empty_proxy_environment),
|
||||
)
|
||||
cached_async_client: Final = AsyncOpenAI(api_key="fake", base_url=_SERVICE_URL)
|
||||
cached_sync_client: Final = OpenAI(api_key="fake", base_url=_SERVICE_URL)
|
||||
discovery_enabled_router.cache.set_cache(
|
||||
|
|
@ -855,14 +834,18 @@ async def test_router_session_requests_stick_to_one_pod_for_sync_and_async(
|
|||
monkeypatch.setattr(loop, "getaddrinfo", async_getaddrinfo)
|
||||
monkeypatch.setattr(socket, "getaddrinfo", sync_getaddrinfo)
|
||||
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
|
||||
_patch_router_pod_discovery(monkeypatch, clock=lambda: 0.0)
|
||||
|
||||
with respx.mock(assert_all_called=True) as respx_mock:
|
||||
route: Final = respx_mock.post(url__regex=_THREE_POD_CHAT_PATTERN).mock(
|
||||
return_value=httpx.Response(200, json=_CHAT_RESPONSE)
|
||||
)
|
||||
async_router: Final = Router(model_list=[_deployment()])
|
||||
sync_router: Final = Router(model_list=[_deployment()])
|
||||
async_router: Final = Router(
|
||||
model_list=[_deployment()],
|
||||
kubernetes_pod_discovery=KubernetesPodDiscovery(clock=lambda: 0.0, proxy_environment=_empty_proxy_environment),
|
||||
)
|
||||
sync_router: Final = Router(
|
||||
model_list=[_deployment()],
|
||||
kubernetes_pod_discovery=KubernetesPodDiscovery(clock=lambda: 0.0, proxy_environment=_empty_proxy_environment),
|
||||
)
|
||||
async_client: Final = _cache_async_client(async_router)
|
||||
sync_client: Final = _cache_sync_client(sync_router)
|
||||
for _ in range(6):
|
||||
|
|
@ -899,14 +882,18 @@ async def test_router_session_mapping_repeats_across_pods(
|
|||
|
||||
monkeypatch.setattr(loop, "getaddrinfo", getaddrinfo)
|
||||
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
|
||||
_patch_router_pod_discovery(monkeypatch, clock=lambda: 0.0)
|
||||
|
||||
session_ids: Final = tuple(f"session-{index}" for index in range(30))
|
||||
with respx.mock(assert_all_called=True) as respx_mock:
|
||||
route: Final = respx_mock.post(url__regex=_THREE_POD_CHAT_PATTERN).mock(
|
||||
return_value=httpx.Response(200, json=_CHAT_RESPONSE)
|
||||
)
|
||||
router: Final = Router(model_list=[_deployment()])
|
||||
router: Final = Router(
|
||||
model_list=[_deployment()],
|
||||
kubernetes_pod_discovery=KubernetesPodDiscovery(
|
||||
clock=lambda: 0.0,
|
||||
proxy_environment=_empty_proxy_environment,
|
||||
),
|
||||
)
|
||||
client: Final = _cache_async_client(router)
|
||||
for session_id in session_ids:
|
||||
await _router_acompletion(router, metadata={"session_id": session_id})
|
||||
|
|
@ -942,13 +929,17 @@ async def test_generated_session_id_uses_round_robin_pod_selection(
|
|||
|
||||
monkeypatch.setattr(loop, "getaddrinfo", getaddrinfo)
|
||||
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
|
||||
_patch_router_pod_discovery(monkeypatch, clock=lambda: 0.0)
|
||||
|
||||
with respx.mock(assert_all_called=True) as respx_mock:
|
||||
route: Final = respx_mock.post(url__regex=_THREE_POD_CHAT_PATTERN).mock(
|
||||
return_value=httpx.Response(200, json=_CHAT_RESPONSE)
|
||||
)
|
||||
router: Final = Router(model_list=[_deployment()])
|
||||
router: Final = Router(
|
||||
model_list=[_deployment()],
|
||||
kubernetes_pod_discovery=KubernetesPodDiscovery(
|
||||
clock=lambda: 0.0,
|
||||
proxy_environment=_empty_proxy_environment,
|
||||
),
|
||||
)
|
||||
client: Final = _cache_async_client(router)
|
||||
for _ in range(6):
|
||||
await _router_acompletion(
|
||||
|
|
@ -982,13 +973,17 @@ async def test_empty_session_id_uses_round_robin_pod_selection(
|
|||
|
||||
monkeypatch.setattr(loop, "getaddrinfo", getaddrinfo)
|
||||
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
|
||||
_patch_router_pod_discovery(monkeypatch, clock=lambda: 0.0)
|
||||
|
||||
with respx.mock(assert_all_called=True) as respx_mock:
|
||||
route: Final = respx_mock.post(url__regex=_THREE_POD_CHAT_PATTERN).mock(
|
||||
return_value=httpx.Response(200, json=_CHAT_RESPONSE)
|
||||
)
|
||||
router: Final = Router(model_list=[_deployment()])
|
||||
router: Final = Router(
|
||||
model_list=[_deployment()],
|
||||
kubernetes_pod_discovery=KubernetesPodDiscovery(
|
||||
clock=lambda: 0.0,
|
||||
proxy_environment=_empty_proxy_environment,
|
||||
),
|
||||
)
|
||||
client: Final = _cache_async_client(router)
|
||||
for _ in range(6):
|
||||
await _router_acompletion(router, metadata={"session_id": ""})
|
||||
|
|
@ -1016,13 +1011,17 @@ async def test_session_requests_do_not_advance_round_robin_cursor(
|
|||
|
||||
monkeypatch.setattr(loop, "getaddrinfo", getaddrinfo)
|
||||
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
|
||||
_patch_router_pod_discovery(monkeypatch, clock=lambda: 0.0)
|
||||
|
||||
with respx.mock(assert_all_called=True) as respx_mock:
|
||||
route: Final = respx_mock.post(url__regex=_THREE_POD_CHAT_PATTERN).mock(
|
||||
return_value=httpx.Response(200, json=_CHAT_RESPONSE)
|
||||
)
|
||||
router: Final = Router(model_list=[_deployment()])
|
||||
router: Final = Router(
|
||||
model_list=[_deployment()],
|
||||
kubernetes_pod_discovery=KubernetesPodDiscovery(
|
||||
clock=lambda: 0.0,
|
||||
proxy_environment=_empty_proxy_environment,
|
||||
),
|
||||
)
|
||||
client: Final = _cache_async_client(router)
|
||||
await _router_acompletion(router)
|
||||
await _router_acompletion(router, metadata={"session_id": "s1"})
|
||||
|
|
@ -1063,16 +1062,18 @@ async def test_session_pod_membership_change_remaps_only_affected_sessions(
|
|||
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
|
||||
|
||||
session_ids: Final = tuple(f"session-{index}" for index in range(60))
|
||||
_patch_router_pod_discovery(
|
||||
monkeypatch,
|
||||
refresh_interval_seconds=10,
|
||||
clock=_clock((0.0,) * 62 + (10.0,) * 10),
|
||||
)
|
||||
with respx.mock(assert_all_called=True) as respx_mock:
|
||||
route: Final = respx_mock.post(url__regex=_THREE_POD_CHAT_PATTERN).mock(
|
||||
return_value=httpx.Response(200, json=_CHAT_RESPONSE)
|
||||
)
|
||||
router: Final = Router(model_list=[_deployment()])
|
||||
router: Final = Router(
|
||||
model_list=[_deployment()],
|
||||
kubernetes_pod_discovery=KubernetesPodDiscovery(
|
||||
refresh_interval_seconds=10,
|
||||
clock=_clock((0.0,) * 62 + (10.0,) * 10),
|
||||
proxy_environment=_empty_proxy_environment,
|
||||
),
|
||||
)
|
||||
client: Final = _cache_async_client(router)
|
||||
for session_id in session_ids:
|
||||
await _router_acompletion(router, metadata={"session_id": session_id})
|
||||
|
|
@ -1132,14 +1133,24 @@ async def test_custom_routing_strategy_keeps_sync_and_async_sessions_on_one_pod(
|
|||
monkeypatch.setattr(loop, "getaddrinfo", async_getaddrinfo)
|
||||
monkeypatch.setattr(socket, "getaddrinfo", sync_getaddrinfo)
|
||||
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
|
||||
_patch_router_pod_discovery(monkeypatch, clock=lambda: 0.0)
|
||||
|
||||
with respx.mock(assert_all_called=True) as respx_mock:
|
||||
route: Final = respx_mock.post(url__regex=_THREE_POD_CHAT_PATTERN).mock(
|
||||
return_value=httpx.Response(200, json=_CHAT_RESPONSE)
|
||||
)
|
||||
async_router: Final = Router(model_list=[_deployment()])
|
||||
sync_router: Final = Router(model_list=[_deployment()])
|
||||
async_router: Final = Router(
|
||||
model_list=[_deployment()],
|
||||
kubernetes_pod_discovery=KubernetesPodDiscovery(
|
||||
clock=lambda: 0.0,
|
||||
proxy_environment=_empty_proxy_environment,
|
||||
),
|
||||
)
|
||||
sync_router: Final = Router(
|
||||
model_list=[_deployment()],
|
||||
kubernetes_pod_discovery=KubernetesPodDiscovery(
|
||||
clock=lambda: 0.0,
|
||||
proxy_environment=_empty_proxy_environment,
|
||||
),
|
||||
)
|
||||
async_router.set_custom_routing_strategy(_ModelListRoutingStrategy(async_router))
|
||||
sync_router.set_custom_routing_strategy(_ModelListRoutingStrategy(sync_router))
|
||||
async_client: Final = _cache_async_client(async_router)
|
||||
|
|
@ -1176,13 +1187,17 @@ async def test_top_level_litellm_session_id_keeps_requests_on_one_pod(
|
|||
|
||||
monkeypatch.setattr(loop, "getaddrinfo", getaddrinfo)
|
||||
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
|
||||
_patch_router_pod_discovery(monkeypatch, clock=lambda: 0.0)
|
||||
|
||||
with respx.mock(assert_all_called=True) as respx_mock:
|
||||
route: Final = respx_mock.post(url__regex=_THREE_POD_CHAT_PATTERN).mock(
|
||||
return_value=httpx.Response(200, json=_CHAT_RESPONSE)
|
||||
)
|
||||
router: Final = Router(model_list=[_deployment()])
|
||||
router: Final = Router(
|
||||
model_list=[_deployment()],
|
||||
kubernetes_pod_discovery=KubernetesPodDiscovery(
|
||||
clock=lambda: 0.0,
|
||||
proxy_environment=_empty_proxy_environment,
|
||||
),
|
||||
)
|
||||
client: Final = _cache_async_client(router)
|
||||
for _ in range(6):
|
||||
await _router_acompletion(router, litellm_session_id="s1")
|
||||
|
|
@ -1210,13 +1225,17 @@ async def test_generated_metadata_ignores_top_level_litellm_session_id(
|
|||
|
||||
monkeypatch.setattr(loop, "getaddrinfo", getaddrinfo)
|
||||
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
|
||||
_patch_router_pod_discovery(monkeypatch, clock=lambda: 0.0)
|
||||
|
||||
with respx.mock(assert_all_called=True) as respx_mock:
|
||||
route: Final = respx_mock.post(url__regex=_THREE_POD_CHAT_PATTERN).mock(
|
||||
return_value=httpx.Response(200, json=_CHAT_RESPONSE)
|
||||
)
|
||||
router: Final = Router(model_list=[_deployment()])
|
||||
router: Final = Router(
|
||||
model_list=[_deployment()],
|
||||
kubernetes_pod_discovery=KubernetesPodDiscovery(
|
||||
clock=lambda: 0.0,
|
||||
proxy_environment=_empty_proxy_environment,
|
||||
),
|
||||
)
|
||||
client: Final = _cache_async_client(router)
|
||||
for _ in range(6):
|
||||
await _router_acompletion(
|
||||
|
|
@ -1267,16 +1286,26 @@ async def test_custom_routing_strategy_resolves_pod_hosts_for_sync_and_async_req
|
|||
monkeypatch.setattr(loop, "getaddrinfo", async_getaddrinfo)
|
||||
monkeypatch.setattr(socket, "getaddrinfo", sync_getaddrinfo)
|
||||
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
|
||||
_patch_router_pod_discovery(monkeypatch, clock=clock)
|
||||
|
||||
with respx.mock(assert_all_called=True) as respx_mock:
|
||||
route: Final = respx_mock.post(
|
||||
url__regex=re.compile(
|
||||
r"http://(?:10\.0\.0\.1|10\.0\.0\.2|vllm-headless\.ns\.svc\.cluster\.local):8000/v1/chat/completions"
|
||||
)
|
||||
).mock(return_value=httpx.Response(200, json=_CHAT_RESPONSE))
|
||||
async_router: Final = Router(model_list=[_deployment()])
|
||||
sync_router: Final = Router(model_list=[_deployment()])
|
||||
async_router: Final = Router(
|
||||
model_list=[_deployment()],
|
||||
kubernetes_pod_discovery=KubernetesPodDiscovery(
|
||||
clock=clock,
|
||||
proxy_environment=_empty_proxy_environment,
|
||||
),
|
||||
)
|
||||
sync_router: Final = Router(
|
||||
model_list=[_deployment()],
|
||||
kubernetes_pod_discovery=KubernetesPodDiscovery(
|
||||
clock=clock,
|
||||
proxy_environment=_empty_proxy_environment,
|
||||
),
|
||||
)
|
||||
async_router.set_custom_routing_strategy(_ModelListRoutingStrategy(async_router))
|
||||
sync_router.set_custom_routing_strategy(_ModelListRoutingStrategy(sync_router))
|
||||
for _ in range(4):
|
||||
|
|
@ -1328,13 +1357,18 @@ async def test_router_async_retry_uses_next_discovered_pod(monkeypatch: pytest.M
|
|||
|
||||
monkeypatch.setattr(loop, "getaddrinfo", getaddrinfo)
|
||||
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
|
||||
_patch_router_pod_discovery(monkeypatch, clock=clock)
|
||||
|
||||
with respx.mock(assert_all_called=True) as respx_mock:
|
||||
route: Final = respx_mock.post(
|
||||
url__regex=re.compile(r"http://(?:10\.0\.0\.1|10\.0\.0\.2):8000/v1/chat/completions")
|
||||
).mock(side_effect=response_for)
|
||||
router: Final = Router(model_list=[_deployment()], num_retries=1)
|
||||
router: Final = Router(
|
||||
model_list=[_deployment()],
|
||||
num_retries=1,
|
||||
kubernetes_pod_discovery=KubernetesPodDiscovery(
|
||||
clock=clock,
|
||||
proxy_environment=_empty_proxy_environment,
|
||||
),
|
||||
)
|
||||
|
||||
response: Final = await router.acompletion(
|
||||
model="gpu-model",
|
||||
|
|
@ -1368,13 +1402,18 @@ def test_router_sync_retry_uses_next_discovered_pod(monkeypatch: pytest.MonkeyPa
|
|||
|
||||
monkeypatch.setattr(socket, "getaddrinfo", getaddrinfo)
|
||||
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
|
||||
_patch_router_pod_discovery(monkeypatch, clock=clock)
|
||||
|
||||
with respx.mock(assert_all_called=True) as respx_mock:
|
||||
route: Final = respx_mock.post(
|
||||
url__regex=re.compile(r"http://(?:10\.0\.0\.1|10\.0\.0\.2):8000/v1/chat/completions")
|
||||
).mock(side_effect=response_for)
|
||||
router: Final = Router(model_list=[_deployment()], num_retries=1)
|
||||
router: Final = Router(
|
||||
model_list=[_deployment()],
|
||||
num_retries=1,
|
||||
kubernetes_pod_discovery=KubernetesPodDiscovery(
|
||||
clock=clock,
|
||||
proxy_environment=_empty_proxy_environment,
|
||||
),
|
||||
)
|
||||
|
||||
response: Final = router.completion(
|
||||
model="gpu-model",
|
||||
|
|
@ -1443,12 +1482,14 @@ async def test_router_session_retry_uses_next_rendezvous_pod(monkeypatch: pytest
|
|||
|
||||
monkeypatch.setattr(loop, "getaddrinfo", getaddrinfo)
|
||||
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
|
||||
_patch_router_pod_discovery(monkeypatch)
|
||||
|
||||
session_id: Final = "retry-session"
|
||||
with respx.mock(assert_all_called=True) as respx_mock:
|
||||
route: Final = respx_mock.post(url__regex=_THREE_POD_CHAT_PATTERN).mock(side_effect=response_for)
|
||||
router: Final = Router(model_list=[_deployment()], num_retries=1)
|
||||
router: Final = Router(
|
||||
model_list=[_deployment()],
|
||||
num_retries=1,
|
||||
kubernetes_pod_discovery=KubernetesPodDiscovery(proxy_environment=_empty_proxy_environment),
|
||||
)
|
||||
|
||||
response: Final = await router.acompletion(
|
||||
model="gpu-model",
|
||||
|
|
|
|||
|
|
@ -11,7 +11,7 @@ import warnings
|
|||
from collections.abc import AsyncIterable, AsyncIterator, Awaitable, Callable, Mapping
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from types import SimpleNamespace
|
||||
from typing import Final, Literal
|
||||
from typing import Final, Literal, TypedDict
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import httpx
|
||||
|
|
@ -19,6 +19,7 @@ import openai
|
|||
import pytest
|
||||
import respx
|
||||
from fastapi import HTTPException
|
||||
from typing_extensions import ReadOnly
|
||||
|
||||
import litellm
|
||||
from litellm import Router
|
||||
|
|
@ -7895,6 +7896,64 @@ def test_update_kwargs_with_deployment_clears_pod_routing_on_non_discovery_fallb
|
|||
assert kwargs["metadata"][KUBERNETES_POD_ROUTING_KEY] is None
|
||||
|
||||
|
||||
def test_update_kwargs_with_deployment_synchronizes_existing_litellm_metadata_pod_routing():
|
||||
from litellm.constants import KUBERNETES_POD_ROUTING_KEY
|
||||
|
||||
routing_record: Final = {
|
||||
"service_host": "vllm-headless.ns.svc.cluster.local",
|
||||
"pod_ip": "10.0.0.1",
|
||||
"pod_count": 3,
|
||||
"selection": "session_affinity",
|
||||
}
|
||||
spoofed_record: Final = {
|
||||
"service_host": "vllm-headless.ns.svc.cluster.local",
|
||||
"pod_ip": "10.0.0.9",
|
||||
"pod_count": 3,
|
||||
"selection": "round_robin",
|
||||
}
|
||||
deployment: Final = {
|
||||
"model_name": "gpt-4o-mini",
|
||||
"litellm_params": {"model": "openai/gpt-4o-mini", "api_key": "fake"},
|
||||
"model_info": {"id": "discovery-id"},
|
||||
KUBERNETES_POD_ROUTING_KEY: routing_record,
|
||||
}
|
||||
original_litellm_metadata: Final = {KUBERNETES_POD_ROUTING_KEY: spoofed_record}
|
||||
kwargs: Final = {"metadata": {}, "litellm_metadata": original_litellm_metadata}
|
||||
router: Final = litellm.Router(model_list=[deployment])
|
||||
|
||||
router._update_kwargs_with_deployment(deployment=deployment, kwargs=kwargs, function_name="completion")
|
||||
|
||||
assert kwargs["metadata"][KUBERNETES_POD_ROUTING_KEY] == routing_record
|
||||
assert kwargs["litellm_metadata"][KUBERNETES_POD_ROUTING_KEY] == routing_record
|
||||
assert kwargs["litellm_metadata"] is not original_litellm_metadata
|
||||
assert original_litellm_metadata[KUBERNETES_POD_ROUTING_KEY] == spoofed_record
|
||||
|
||||
|
||||
def test_update_kwargs_with_non_discovery_deployment_clears_other_pod_routing_bucket():
|
||||
from litellm.constants import KUBERNETES_POD_ROUTING_KEY
|
||||
|
||||
previous_record: Final = {
|
||||
"service_host": "vllm-headless.ns.svc.cluster.local",
|
||||
"pod_ip": "10.0.0.1",
|
||||
"pod_count": 3,
|
||||
"selection": "session_affinity",
|
||||
}
|
||||
deployment: Final = {
|
||||
"model_name": "gpt-4o-mini",
|
||||
"litellm_params": {"model": "openai/gpt-4o-mini", "api_key": "fake"},
|
||||
"model_info": {"id": "fallback-id"},
|
||||
}
|
||||
original_litellm_metadata: Final = {KUBERNETES_POD_ROUTING_KEY: previous_record}
|
||||
kwargs: Final = {"metadata": {}, "litellm_metadata": original_litellm_metadata}
|
||||
router: Final = litellm.Router(model_list=[deployment])
|
||||
|
||||
router._update_kwargs_with_deployment(deployment=deployment, kwargs=kwargs, function_name="completion")
|
||||
|
||||
assert kwargs["metadata"][KUBERNETES_POD_ROUTING_KEY] is None
|
||||
assert kwargs["litellm_metadata"][KUBERNETES_POD_ROUTING_KEY] is None
|
||||
assert original_litellm_metadata[KUBERNETES_POD_ROUTING_KEY] == previous_record
|
||||
|
||||
|
||||
def test_combine_fallback_usage():
|
||||
"""Test that _combine_fallback_usage merges partial and fallback usage."""
|
||||
from litellm.router import Router
|
||||
|
|
@ -13728,8 +13787,14 @@ def _anthropic_messages_make_wrapper() -> FallbackAwareAnthropicMessagesStream:
|
|||
return FallbackAwareAnthropicMessagesStream(_anthropic_messages_empty_generator(), object())
|
||||
|
||||
|
||||
def _anthropic_messages_make_router(**router_kwargs) -> Router:
|
||||
router_kwargs.setdefault("fallbacks", [{"primary": ["fallback"]}])
|
||||
class _AnthropicMessagesRouterKwargs(TypedDict, total=False):
|
||||
fallbacks: ReadOnly[list[dict[str, list[str]]] | None]
|
||||
content_policy_fallbacks: ReadOnly[list[dict[str, list[str]]] | None]
|
||||
enable_weighted_failover: ReadOnly[bool]
|
||||
|
||||
|
||||
def _anthropic_messages_make_router(router_kwargs: _AnthropicMessagesRouterKwargs) -> Router:
|
||||
fallbacks: Final = router_kwargs.get("fallbacks", [{"primary": ["fallback"]}])
|
||||
return Router(
|
||||
model_list=[
|
||||
{
|
||||
|
|
@ -13746,7 +13811,9 @@ def _anthropic_messages_make_router(**router_kwargs) -> Router:
|
|||
},
|
||||
},
|
||||
],
|
||||
**router_kwargs,
|
||||
fallbacks=fallbacks,
|
||||
content_policy_fallbacks=router_kwargs.get("content_policy_fallbacks"),
|
||||
enable_weighted_failover=router_kwargs.get("enable_weighted_failover", False),
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -14115,8 +14182,12 @@ def _anthropic_messages_two_order_primary_model_list() -> list:
|
|||
pytest.param({"fallbacks": None, "enable_weighted_failover": True}, {"model": "primary"}, True, id="weighted-failover"),
|
||||
],
|
||||
)
|
||||
def test_anthropic_messages_stream_can_fall_back_direct_call(router_kwargs, request_kwargs, expected):
|
||||
router = _anthropic_messages_make_router(**router_kwargs)
|
||||
def test_anthropic_messages_stream_can_fall_back_direct_call(
|
||||
router_kwargs: _AnthropicMessagesRouterKwargs,
|
||||
request_kwargs: dict[str, object],
|
||||
expected: bool,
|
||||
) -> None:
|
||||
router = _anthropic_messages_make_router(router_kwargs)
|
||||
assert router._anthropic_messages_stream_can_fall_back("primary", request_kwargs) is expected
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -359,10 +359,20 @@ class TestNoProviderRetryAmplification:
|
|||
await session.aclose()
|
||||
|
||||
@staticmethod
|
||||
def _router(api_base: str, litellm_params: dict, **router_kwargs) -> Router:
|
||||
def _router(
|
||||
api_base: str,
|
||||
litellm_params: dict,
|
||||
*,
|
||||
num_retries: int,
|
||||
retry_policy: RetryPolicy | None = None,
|
||||
) -> Router:
|
||||
params = {"model": "openai/gpt-4o-mini", "api_base": api_base, "api_key": "sk-fake"}
|
||||
params.update(litellm_params)
|
||||
return Router(model_list=[{"model_name": "mock", "litellm_params": params}], **router_kwargs)
|
||||
return Router(
|
||||
model_list=[{"model_name": "mock", "litellm_params": params}],
|
||||
num_retries=num_retries,
|
||||
retry_policy=retry_policy,
|
||||
)
|
||||
|
||||
async def _call_and_count(self, router: Router, **call_kwargs) -> int:
|
||||
counter = self._install_counting_upstream()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue