From 1d51a8dfc37b9f710de915cf63432b6a78613d71 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Thu, 24 Sep 2026 19:51:02 -0700 Subject: [PATCH 01/10] fix(proxy): authorize key model aliases the same way as team aliases (#43049) --- litellm/proxy/auth/auth_checks.py | 105 ++++- litellm/proxy/auth/user_api_key_auth.py | 2 + .../proxy/common_utils/model_listing_utils.py | 8 +- .../test_key_alias_model_access.py | 72 ++++ tests/test_keys.py | 8 +- .../proxy/auth/test_auth_checks.py | 386 ++++++++++++++++++ 6 files changed, 564 insertions(+), 17 deletions(-) create mode 100644 tests/integration/authorization/test_key_alias_model_access.py diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index 67950e603c0..f19a8055ae6 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -89,6 +89,7 @@ from litellm.proxy.common_utils.http_parsing_utils import ( _safe_get_request_headers, _safe_get_request_query_params, ) +from litellm.proxy.common_utils.model_listing_utils import alias_map from litellm.proxy.common_utils.timezone_utils import get_budget_reset_time from litellm.proxy.common_utils.user_api_key_cache import ( END_USER_RESTRICTED_REGISTRY_OVERFLOW_SENTINEL, @@ -722,6 +723,7 @@ async def _run_project_checks( model=_model, project_object=project_object, llm_router=llm_router, + key_model_aliases=key_model_aliases_for_auth_check(valid_token), ) if not skip_budget_checks: @@ -1018,6 +1020,7 @@ async def common_checks( team_object=team_object, llm_router=llm_router, team_model_aliases=(valid_token.team_model_aliases if valid_token else None), + key_model_aliases=key_model_aliases_for_auth_check(valid_token), ) except ProxyException as team_denial: if team_denial.type != ProxyErrorTypes.team_model_access_denied: @@ -1027,6 +1030,7 @@ async def common_checks( valid_token=valid_token, team_object=team_object, llm_router=llm_router, + key_model_aliases=key_model_aliases_for_auth_check(valid_token), ): raise @@ -1043,6 +1047,7 @@ async def common_checks( proxy_logging_obj=proxy_logging_obj, team_membership=loaded_team_membership, team_membership_loaded=team_membership_loaded, + key_model_aliases=key_model_aliases_for_auth_check(valid_token), ) # Require trace id for agent keys when agent has require_trace_id_on_calls_by_agent @@ -1081,6 +1086,7 @@ async def common_checks( model=_model, llm_router=llm_router, user_object=user_object, + key_model_aliases=key_model_aliases_for_auth_check(valid_token), ) # 1.1 - 2.2 - 3.0.2 - 3.0.3: Project checks (blocked, model access, budget) @@ -4349,6 +4355,7 @@ def _can_object_call_model( models: list[str], team_model_aliases: dict[str, str] | None = None, team_id: str | None = None, + key_model_aliases: Mapping[str, str] | None = None, object_type: Literal["user", "team", "key", "org", "project", "agent"] = "user", fallback_depth: int = 0, ) -> Literal[True]: @@ -4378,6 +4385,7 @@ def _can_object_call_model( models=models, team_model_aliases=team_model_aliases, team_id=team_id, + key_model_aliases=key_model_aliases, object_type=object_type, fallback_depth=fallback_depth + 1, ) @@ -4386,13 +4394,32 @@ def _can_object_call_model( from litellm.router_strategy.complexity_router.context_compaction import native_compaction_parent compaction_parent: Final = native_compaction_parent(model) - potential_models: Final = [model, compaction_parent] if compaction_parent is not None else [model] - if model in litellm.model_alias_map: - potential_models.append(litellm.model_alias_map[model]) - elif llm_router and model in llm_router.model_group_alias: - _model: Final = llm_router._get_model_from_alias(model) - if _model: - potential_models.append(_model) + global_or_router_alias_target: Final = ( + litellm.model_alias_map[model] + if model in litellm.model_alias_map + else ( + llm_router._get_model_from_alias(model) + if llm_router is not None and model in llm_router.model_group_alias + else None + ) + ) + after_team_alias: Final = team_model_aliases.get(model, model) if team_model_aliases else model + after_key_alias: Final = ( + key_model_aliases.get(after_team_alias, after_team_alias) if key_model_aliases else after_team_alias + ) + after_global_alias: Final = litellm.model_alias_map.get(after_key_alias, after_key_alias) + dispatched_model: Final = ( + key_model_aliases.get(after_global_alias, after_global_alias) if key_model_aliases else after_global_alias + ) + key_alias_applied: Final = after_key_alias != after_team_alias or dispatched_model != after_global_alias + potential_models: Final = ( + (dispatched_model,) + if key_alias_applied + else ( + *((model, compaction_parent) if compaction_parent is not None else (model,)), + *((global_or_router_alias_target,) if global_or_router_alias_target else ()), + ) + ) ## check model access for alias + underlying model - allow if either is in allowed models for m in potential_models: @@ -4418,6 +4445,35 @@ def _can_object_call_model( ) +def _resolve_team_alias( + model: str | list[str], + team_model_aliases: dict[str, str] | None, + team_id: str | None, + llm_router: Router | None, +) -> str | list[str]: + if not team_model_aliases: + return model + if isinstance(model, str): + return _live_team_alias_target(model, team_model_aliases, team_id, llm_router) + return [ # mutable-ok: _can_object_call_model takes list[str] + _live_team_alias_target(name, team_model_aliases, team_id, llm_router) for name in model + ] + + +def _live_team_alias_target( + model: str, team_model_aliases: dict[str, str], team_id: str | None, llm_router: Router | None +) -> str: + target: Final = team_model_aliases.get(model) + if target is None: + return model + deleted_team_deployment: Final = ( + llm_router is not None + and target.startswith(f"model_name_{team_id}_") + and target not in llm_router.model_name_to_deployment_indices + ) + return model if deleted_team_deployment else target + + async def _check_agent_access_group_model_access( model: str | list[str] | None, # mutable-ok: _can_object_call_model and the client message helper take list[str] valid_token: UserAPIKeyAuth | None, @@ -4438,12 +4494,14 @@ async def _check_agent_access_group_model_access( param="model", code=status.HTTP_403_FORBIDDEN, ) + dispatched: Final = _resolve_team_alias(model, valid_token.team_model_aliases, valid_token.team_id, llm_router) return _can_object_call_model( - model=model, + model=dispatched, llm_router=llm_router, models=sorted(ceiling.models), team_id=valid_token.team_id, object_type="agent", + key_model_aliases=key_model_aliases_for_auth_check(valid_token), ) @@ -4471,12 +4529,14 @@ async def _check_agent_caller_model_access( if caller_auth is None: return caller_team: Final = await load_team(valid_token) + caller_key_model_aliases: Final = key_model_aliases_for_auth_check(valid_token) if caller_team is not None: await can_team_access_model( model=model, team_object=caller_team, llm_router=llm_router, prisma_client=prisma_client, + key_model_aliases=caller_key_model_aliases, ) await _check_team_member_model_access( model=model, @@ -4486,12 +4546,18 @@ async def _check_agent_caller_model_access( prisma_client=prisma_client, user_api_key_cache=user_api_key_cache, proxy_logging_obj=proxy_logging_obj, + key_model_aliases=caller_key_model_aliases, ) return caller_user: Final = await load_user(valid_token) if caller_user is None: return - await can_user_call_model(model=model, llm_router=llm_router, user_object=caller_user) + await can_user_call_model( + model=model, + llm_router=llm_router, + user_object=caller_user, + key_model_aliases=caller_key_model_aliases, + ) def _model_in_team_aliases(model: str, team_model_aliases: dict[str, str] | None = None) -> bool: @@ -4512,6 +4578,10 @@ def _model_in_team_aliases(model: str, team_model_aliases: dict[str, str] | None return False +def key_model_aliases_for_auth_check(valid_token: UserAPIKeyAuth | None) -> Mapping[str, str] | None: + return alias_map(valid_token.aliases) if valid_token is not None and valid_token.aliases else None + + def _resolve_key_models_for_auth_check(valid_token: UserAPIKeyAuth) -> list[str]: """ Expand key model sentinels before auth checks. @@ -4831,6 +4901,7 @@ async def can_key_call_model( models=key_models, team_model_aliases=valid_token.team_model_aliases, team_id=valid_token.team_id, + key_model_aliases=key_model_aliases_for_auth_check(valid_token), object_type="key", ) except ProxyException: @@ -4848,6 +4919,7 @@ async def can_key_call_model( models=models_from_groups, team_model_aliases=valid_token.team_model_aliases, team_id=valid_token.team_id, + key_model_aliases=key_model_aliases_for_auth_check(valid_token), object_type="key", ) raise @@ -4906,6 +4978,7 @@ async def can_key_call_resolved_model( team_object=team_object, llm_router=llm_router, team_model_aliases=valid_token.team_model_aliases, + key_model_aliases=key_model_aliases_for_auth_check(valid_token), ) except ProxyException as team_denial: if team_denial.type != ProxyErrorTypes.team_model_access_denied: @@ -4915,6 +4988,7 @@ async def can_key_call_resolved_model( valid_token=valid_token, team_object=team_object, llm_router=llm_router, + key_model_aliases=key_model_aliases_for_auth_check(valid_token), ): raise @@ -4927,6 +5001,7 @@ async def can_key_call_resolved_model( prisma_client=prisma_client, user_api_key_cache=user_api_key_cache, proxy_logging_obj=proxy_logging_obj, + key_model_aliases=key_model_aliases_for_auth_check(valid_token), ) if valid_token.project_id is not None: @@ -4941,6 +5016,7 @@ async def can_key_call_resolved_model( model=model, project_object=project_object, llm_router=llm_router, + key_model_aliases=key_model_aliases_for_auth_check(valid_token), ) @@ -4968,6 +5044,7 @@ async def can_team_access_model( team_object: LiteLLM_TeamTable | None, llm_router: Router | None, team_model_aliases: dict[str, str] | None = None, + key_model_aliases: Mapping[str, str] | None = None, prisma_client: DatabaseClient | None = None, ) -> Literal[True]: """ @@ -4983,6 +5060,7 @@ async def can_team_access_model( models=team_object.models if team_object else [], team_model_aliases=team_model_aliases, team_id=team_object.team_id if team_object else None, + key_model_aliases=key_model_aliases, object_type="team", ) except ProxyException: @@ -5000,6 +5078,7 @@ async def can_team_access_model( models=list(dict.fromkeys([*(team_object.models if team_object else []), *models_from_groups])), team_model_aliases=team_model_aliases, team_id=team_object.team_id if team_object else None, + key_model_aliases=key_model_aliases, object_type="team", ) raise @@ -5058,6 +5137,7 @@ async def _key_access_group_grants_model( valid_token: UserAPIKeyAuth | None, team_object: LiteLLM_TeamTable | None, llm_router: Router | None, + key_model_aliases: Mapping[str, str] | None = None, ) -> bool: """ Returns True if the key's `access_group_ids` expand to models that grant @@ -5078,6 +5158,7 @@ async def _key_access_group_grants_model( models=authorized_models, team_model_aliases=valid_token.team_model_aliases if valid_token else None, team_id=valid_token.team_id if valid_token else None, + key_model_aliases=key_model_aliases, object_type="key", ) return True @@ -5089,6 +5170,7 @@ def can_project_access_model( model: str | list[str], project_object: LiteLLM_ProjectTable, llm_router: Router | None, + key_model_aliases: Mapping[str, str] | None = None, ) -> Literal[True]: """ Returns True if the project can access a specific model. @@ -5099,6 +5181,7 @@ def can_project_access_model( model=model, llm_router=llm_router, models=project_object.models if project_object else [], + key_model_aliases=key_model_aliases, object_type="project", ) @@ -5107,6 +5190,7 @@ async def can_user_call_model( model: str | list[str], llm_router: Router | None, user_object: LiteLLM_UserTable | None, + key_model_aliases: Mapping[str, str] | None = None, ) -> Literal[True]: if user_object is None: return True @@ -5128,6 +5212,7 @@ async def can_user_call_model( model=model, llm_router=llm_router, models=user_object.models, + key_model_aliases=key_model_aliases, object_type="user", ) @@ -5682,6 +5767,7 @@ async def _check_team_member_model_access( proxy_logging_obj: ProxyLogging, team_membership: LiteLLM_TeamMembership | None = None, team_membership_loaded: bool = False, + key_model_aliases: Mapping[str, str] | None = None, ) -> None: """ Check if a team member's per-member model scope allows access to the requested model. @@ -5717,6 +5803,7 @@ async def _check_team_member_model_access( models=member_allowed_models, object_type="team", team_id=team_object.team_id, + key_model_aliases=key_model_aliases, ) except ProxyException: internal_message: Final = ( diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index ae95e94dd2d..22c3a248b9d 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -66,6 +66,7 @@ from litellm.proxy.auth.auth_checks import ( get_user_object, is_valid_fallback_model, jwt_key_mapping_cache_key, + key_model_aliases_for_auth_check, resolve_and_validate_end_user_id, resolve_default_end_user_budget, ) @@ -469,6 +470,7 @@ async def _check_key_model_budget_with_fallback( models=valid_token.team_models, team_model_aliases=valid_token.team_model_aliases, team_id=valid_token.team_id, + key_model_aliases=key_model_aliases_for_auth_check(valid_token), object_type="team", ) except ProxyException: diff --git a/litellm/proxy/common_utils/model_listing_utils.py b/litellm/proxy/common_utils/model_listing_utils.py index 3c6555662e6..8958fb20918 100644 --- a/litellm/proxy/common_utils/model_listing_utils.py +++ b/litellm/proxy/common_utils/model_listing_utils.py @@ -180,7 +180,7 @@ def caller_alias_maps( return CallerAliases((team_aliases, key_aliases), (team_aliases, key_aliases, litellm.model_alias_map, key_aliases)) -def _alias_map(aliases: object) -> Mapping[str, str]: +def alias_map(aliases: object) -> Mapping[str, str]: try: entries: Final = _ALIAS_ENTRIES.validate_python(aliases, strict=True) except ValidationError: @@ -204,7 +204,7 @@ def alias_target(model_id: str, aliases: CallerAliases, listed: Container[str] = already `listed` keeps its own row, so it is never rewritten.""" if model_id in listed: return None - return _rewrite(model_id, tuple(_alias_map(alias_map) for alias_map in aliases.rewrite)) + return _rewrite(model_id, tuple(alias_map(raw) for raw in aliases.rewrite)) def alias_listing_entries( @@ -213,8 +213,8 @@ def alias_listing_entries( ) -> tuple[tuple[str, str], ...]: """`entries` plus one `(alias, lookup_id)` row per key or team alias whose target is listed. An alias colliding with a listed id keeps the listed entry.""" - maps: Final = tuple(_alias_map(alias_map) for alias_map in aliases.rewrite) - own: Final = tuple(_alias_map(alias_map) for alias_map in aliases.own) + maps: Final = tuple(alias_map(raw) for raw in aliases.rewrite) + own: Final = tuple(alias_map(raw) for raw in aliases.own) lookup_by_response: Final = MappingProxyType(dict(entries)) lookup_ids: Final = frozenset(lookup_by_response.values()) targets: Final = MappingProxyType( diff --git a/tests/integration/authorization/test_key_alias_model_access.py b/tests/integration/authorization/test_key_alias_model_access.py new file mode 100644 index 00000000000..50fc53bd4a9 --- /dev/null +++ b/tests/integration/authorization/test_key_alias_model_access.py @@ -0,0 +1,72 @@ +import uuid +from typing import Final + +import httpx + +from tests.integration._support.client import Gateway, eventually, object_value, string_value + + +def _listed_model_ids(response: httpx.Response) -> frozenset[str]: + entries: Final = response.json()["data"] + assert isinstance(entries, list), response.text + return frozenset(string_value(object_value(entry)["id"]) for entry in entries) + + +def _listed_and_callable(gateway: Gateway, key: str, model: str, alias: str) -> None: + """Every id /v1/models lists for this key must be callable by the same key.""" + response: Final = eventually( + lambda: gateway.request("GET", "/v1/models", key=key), + lambda value: value.status_code == 200 and _listed_model_ids(value) == frozenset({model, alias}), + return_last_on_timeout=True, + ) + assert response.status_code == 200, response.text + listed: Final = _listed_model_ids(response) + assert listed == frozenset({model, alias}), response.text + for model_id in sorted(listed): + called: Final = gateway.request( + "POST", + "/v1/chat/completions", + {"model": model_id, "messages": [{"role": "user", "content": "ping"}]}, + key=key, + ) + assert called.status_code == 200, f"listed id {model_id} is not callable: {called.status_code} {called.text}" + + +def test_key_alias_listed_by_v1_models_is_callable(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + model: Final = scenario.model() + alias: Final = f"integration-alias-{uuid.uuid4().hex}" + key: Final = scenario.key(models=[model], aliases={alias: model}) + _listed_and_callable(gateway, key, model, alias) + + +def test_team_key_alias_listed_by_v1_models_is_callable(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + model: Final = scenario.model() + team_id: Final = scenario.team(models=[model]) + alias: Final = f"integration-alias-{uuid.uuid4().hex}" + key: Final = scenario.key(team_id=team_id, aliases={alias: model}) + _listed_and_callable(gateway, key, model, alias) + + +def test_key_alias_to_model_outside_key_allowlist_is_hidden_and_denied(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + allowed: Final = scenario.model() + hidden: Final = scenario.model() + alias: Final = f"integration-alias-{uuid.uuid4().hex}" + key: Final = scenario.key(models=[allowed], aliases={alias: hidden}) + response: Final = eventually( + lambda: gateway.request("GET", "/v1/models", key=key), + lambda value: value.status_code == 200 and _listed_model_ids(value) == frozenset({allowed}), + return_last_on_timeout=True, + ) + assert response.status_code == 200, response.text + assert _listed_model_ids(response) == frozenset({allowed}), response.text + called: Final = gateway.request( + "POST", + "/v1/chat/completions", + {"model": alias, "messages": [{"role": "user", "content": "ping"}]}, + key=key, + ) + assert called.status_code == 403, called.text + assert "key_model_access_denied" in called.text, called.text diff --git a/tests/test_keys.py b/tests/test_keys.py index 7a5b2502cfd..c1785b88822 100644 --- a/tests/test_keys.py +++ b/tests/test_keys.py @@ -834,12 +834,12 @@ async def test_key_model_list(model_access, model_access_level, model_endpoint): assert len(model_list["data"]) > 0 if model_access == "gpt-3.5-turbo": if model_endpoint == "/v1/models": - assert ( - len(model_list["data"]) == 1 - ), "model_access={}, model_access_level={}".format( + assert {entry["id"] for entry in model_list["data"]} == { + model_access, + "mistral-7b", + }, "generate_key sets alias mistral-7b -> gpt-3.5-turbo, so /v1/models lists both; model_access={}, model_access_level={}".format( model_access, model_access_level ) - assert model_list["data"][0]["id"] == model_access elif model_endpoint == "/model/info": assert isinstance(model_list["data"], list) assert len(model_list["data"]) == 1 diff --git a/tests/test_litellm/proxy/auth/test_auth_checks.py b/tests/test_litellm/proxy/auth/test_auth_checks.py index 30f5abdbb98..b811d4453ca 100644 --- a/tests/test_litellm/proxy/auth/test_auth_checks.py +++ b/tests/test_litellm/proxy/auth/test_auth_checks.py @@ -1783,6 +1783,336 @@ def test_can_object_call_model_access_via_alias_only(): assert result is True +def test_can_object_call_model_key_alias_to_allowed_target_is_allowed(): + """A key alias whose target is on the key allowlist resolves like a team alias.""" + from litellm.proxy.auth.auth_checks import _can_object_call_model + + result = _can_object_call_model( + model="mistral-7b", + llm_router=None, + models=["gpt-4o-mini"], + key_model_aliases={"mistral-7b": "gpt-4o-mini"}, + object_type="key", + fallback_depth=0, + ) + + assert result is True + + +def test_can_object_call_model_key_alias_to_disallowed_target_is_denied(): + """A key alias whose target is outside the key allowlist stays denied.""" + from litellm.proxy._types import ProxyErrorTypes, ProxyException + from litellm.proxy.auth.auth_checks import _can_object_call_model + + with pytest.raises(ProxyException) as exc_info: + _can_object_call_model( + model="mistral-7b", + llm_router=None, + models=["gpt-4o-mini"], + key_model_aliases={"mistral-7b": "gpt-4"}, + object_type="key", + fallback_depth=0, + ) + + assert exc_info.value.type == ProxyErrorTypes.key_model_access_denied + assert exc_info.value.code == "403" + + +@pytest.mark.asyncio +async def test_can_team_access_model_honors_key_alias(): + """A key on a team can call a model through its own alias when the target is on the team allowlist.""" + from litellm.proxy.auth.auth_checks import can_team_access_model + + team_object = LiteLLM_TeamTable( + team_id="team-123", + models=["gpt-4o-mini"], + ) + + assert ( + await can_team_access_model( + model="mistral-7b", + team_object=team_object, + llm_router=None, + key_model_aliases={"mistral-7b": "gpt-4o-mini"}, + ) + is True + ) + + with pytest.raises(ProxyException) as exc_info: + await can_team_access_model( + model="mistral-7b", + team_object=team_object, + llm_router=None, + ) + + assert exc_info.value.type == ProxyErrorTypes.team_model_access_denied + + +@pytest.mark.asyncio +async def test_can_key_call_model_honors_key_alias(): + """The real key entry point resolves a key alias to its target before the allowlist check.""" + from litellm.proxy.auth.auth_checks import can_key_call_model + + allowed_token = UserAPIKeyAuth( + api_key="sk-test", + models=["gpt-4o-mini"], + aliases={"mistral-7b": "gpt-4o-mini"}, + ) + + assert ( + await can_key_call_model( + model="mistral-7b", + llm_model_list=None, + valid_token=allowed_token, + llm_router=None, + ) + is True + ) + + denied_token = UserAPIKeyAuth( + api_key="sk-test", + models=["gpt-4o-mini"], + aliases={"mistral-7b": "gpt-4"}, + ) + + with pytest.raises(ProxyException) as exc_info: + await can_key_call_model( + model="mistral-7b", + llm_model_list=None, + valid_token=denied_token, + llm_router=None, + ) + + assert exc_info.value.type == ProxyErrorTypes.key_model_access_denied + + +def test_can_object_call_model_key_alias_applies_before_global_alias(monkeypatch): + """The key alias rewrite precedes the global one at dispatch, so the key target is authorized.""" + from litellm.proxy.auth.auth_checks import _can_object_call_model + + monkeypatch.setattr(litellm, "model_alias_map", {"foo": "bar"}) + + assert ( + _can_object_call_model( + model="foo", + llm_router=None, + models=["baz"], + key_model_aliases={"foo": "baz"}, + object_type="key", + fallback_depth=0, + ) + is True + ) + + with pytest.raises(ProxyException) as exc_info: + _can_object_call_model( + model="foo", + llm_router=None, + models=["bar"], + key_model_aliases={"foo": "baz"}, + object_type="key", + fallback_depth=0, + ) + + assert exc_info.value.type == ProxyErrorTypes.key_model_access_denied + + +def test_can_object_call_model_key_alias_matches_global_rewritten_name(monkeypatch): + """A key alias on the globally rewritten name resolves the same way the request chain does.""" + from litellm.proxy.auth.auth_checks import _can_object_call_model + + monkeypatch.setattr(litellm, "model_alias_map", {"foo": "bar"}) + + assert ( + _can_object_call_model( + model="foo", + llm_router=None, + models=["baz"], + key_model_aliases={"bar": "baz"}, + object_type="key", + fallback_depth=0, + ) + is True + ) + + +def test_can_object_call_model_chained_alias_requires_final_target(monkeypatch): + """When a key alias fires on the globally rewritten name, only the final target is dispatched.""" + from litellm.proxy.auth.auth_checks import _can_object_call_model + + monkeypatch.setattr(litellm, "model_alias_map", {"foo": "bar"}) + + with pytest.raises(ProxyException) as exc_info: + _can_object_call_model( + model="foo", + llm_router=None, + models=["bar"], + key_model_aliases={"bar": "baz"}, + object_type="key", + fallback_depth=0, + ) + + assert exc_info.value.type == ProxyErrorTypes.key_model_access_denied + + assert ( + _can_object_call_model( + model="foo", + llm_router=None, + models=["baz"], + key_model_aliases={"bar": "baz"}, + object_type="key", + fallback_depth=0, + ) + is True + ) + + +def test_can_object_call_model_key_alias_name_alone_is_not_enough(): + """A key that may call the alias name but not its target cannot call the alias.""" + from litellm.proxy.auth.auth_checks import _can_object_call_model + + with pytest.raises(ProxyException) as exc_info: + _can_object_call_model( + model="bar", + llm_router=None, + models=["bar"], + key_model_aliases={"bar": "baz"}, + object_type="key", + fallback_depth=0, + ) + + assert exc_info.value.type == ProxyErrorTypes.key_model_access_denied + + assert ( + _can_object_call_model( + model="bar", + llm_router=None, + models=["baz"], + key_model_aliases={"bar": "baz"}, + object_type="key", + fallback_depth=0, + ) + is True + ) + + +def test_can_object_call_model_team_alias_applies_before_key_alias(): + """A key alias on the raw name loses to the team alias that rewrites it first at dispatch.""" + from litellm.proxy.auth.auth_checks import _can_object_call_model + + assert ( + _can_object_call_model( + model="foo", + llm_router=None, + models=["bar"], + team_model_aliases={"foo": "bar"}, + key_model_aliases={"foo": "baz"}, + object_type="key", + fallback_depth=0, + ) + is True + ) + + +def test_can_object_call_model_key_alias_on_team_alias_target(): + """A key alias on the team-rewritten name resolves like the dispatch chain does.""" + from litellm.proxy.auth.auth_checks import _can_object_call_model + + assert ( + _can_object_call_model( + model="foo", + llm_router=None, + models=["baz"], + team_model_aliases={"foo": "bar"}, + key_model_aliases={"bar": "baz"}, + object_type="key", + fallback_depth=0, + ) + is True + ) + + with pytest.raises(ProxyException) as exc_info: + _can_object_call_model( + model="foo", + llm_router=None, + models=["bar"], + team_model_aliases={"foo": "bar"}, + key_model_aliases={"bar": "baz"}, + object_type="key", + fallback_depth=0, + ) + + assert exc_info.value.type == ProxyErrorTypes.key_model_access_denied + + +@pytest.mark.asyncio +async def test_can_user_call_model_honors_key_alias(): + """A personal-scope key alias resolves to its target before the user allowlist check.""" + from litellm.proxy.auth.auth_checks import can_user_call_model + + user_object = LiteLLM_UserTable(user_id="test-user", models=["gpt-4o-mini"]) + + assert ( + await can_user_call_model( + model="mistral-7b", + llm_router=None, + user_object=user_object, + key_model_aliases={"mistral-7b": "gpt-4o-mini"}, + ) + is True + ) + + with pytest.raises(ProxyException) as exc_info: + await can_user_call_model( + model="mistral-7b", + llm_router=None, + user_object=user_object, + ) + + assert exc_info.value.type == ProxyErrorTypes.user_model_access_denied + + +@pytest.mark.asyncio +async def test_check_team_member_model_access_honors_key_alias(): + """A key alias resolves against the member allowlist, not just the raw alias name.""" + from litellm.proxy._types import LiteLLM_TeamMembership + from litellm.proxy.auth.auth_checks import _check_team_member_model_access + + membership = LiteLLM_TeamMembership( + user_id="alice", + team_id="team-a", + litellm_budget_table=LiteLLM_BudgetTable(allowed_models=["gpt-4o-mini"]), + ) + + await _check_team_member_model_access( + model="mistral-7b", + team_object=LiteLLM_TeamTable(team_id="team-a"), + valid_token=UserAPIKeyAuth(token="sk-test", user_id="alice", team_id="team-a"), + llm_router=None, + prisma_client=None, + user_api_key_cache=UserApiKeyCache(), + proxy_logging_obj=MagicMock(), + team_membership=membership, + team_membership_loaded=True, + key_model_aliases={"mistral-7b": "gpt-4o-mini"}, + ) + + with pytest.raises(ProxyException) as exc_info: + await _check_team_member_model_access( + model="mistral-7b", + team_object=LiteLLM_TeamTable(team_id="team-a"), + valid_token=UserAPIKeyAuth(token="sk-test", user_id="alice", team_id="team-a"), + llm_router=None, + prisma_client=None, + user_api_key_cache=UserApiKeyCache(), + proxy_logging_obj=MagicMock(), + team_membership=membership, + team_membership_loaded=True, + ) + + assert exc_info.value.type == ProxyErrorTypes.team_model_access_denied + + def test_can_object_call_model_access_via_underlying_model_only(): """ Test that a key can access a model via underlying model even when using an alias. @@ -9139,6 +9469,50 @@ async def test_agent_access_groups_cap_models_even_when_key_allows_them(): assert asked == ["agent-1", "agent-1"] +@pytest.mark.asyncio +async def test_agent_access_group_ceiling_admits_the_key_alias_target(): + agent_key: Final = UserAPIKeyAuth(token="agent-token", agent_id="agent-1", models=["gpt-5"]) + agent_key.aliases = {"fast": "gpt-5"} + resolve, asked = _agent_model_ceiling_resolver(frozenset({"gpt-5"})) + + assert await _check_agent_access_group_model_access("fast", agent_key, None, resolve) is True + assert asked == ["agent-1"] + + +@pytest.mark.asyncio +async def test_agent_access_group_ceiling_checks_the_team_alias_target(): + agent_key: Final = UserAPIKeyAuth(token="agent-token", agent_id="agent-1", team_id="team-1", models=["gpt-5"]) + agent_key.team_model_aliases = {"foo": "gpt-5"} + resolve, asked = _agent_model_ceiling_resolver(frozenset({"gpt-5"})) + + assert await _check_agent_access_group_model_access("foo", agent_key, None, resolve) is True + assert asked == ["agent-1"] + + +@pytest.mark.asyncio +async def test_agent_access_group_ceiling_denies_a_team_alias_outside_the_ceiling(): + agent_key: Final = UserAPIKeyAuth(token="agent-token", agent_id="agent-1", team_id="team-1", models=["gpt-5"]) + agent_key.team_model_aliases = {"foo": "claude-sonnet-4-5"} + resolve, _ = _agent_model_ceiling_resolver(frozenset({"gpt-5"})) + + with pytest.raises(ModelAccessDeniedProxyException) as exc: + await _check_agent_access_group_model_access("foo", agent_key, None, resolve) + assert exc.value.type == ProxyErrorTypes.agent_model_access_denied + + +@pytest.mark.asyncio +async def test_agent_access_group_ceiling_keeps_the_name_for_a_deleted_team_deployment(): + from litellm.router import Router + + agent_key: Final = UserAPIKeyAuth(token="agent-token", agent_id="agent-1", team_id="team-1", models=["foo"]) + agent_key.team_model_aliases = {"foo": "model_name_team-1_deadbeef"} + router: Final = Router(model_list=[]) + resolve, asked = _agent_model_ceiling_resolver(frozenset({"foo"})) + + assert await _check_agent_access_group_model_access("foo", agent_key, router, resolve) is True + assert asked == ["agent-1"] + + @pytest.mark.asyncio async def test_agent_access_groups_naming_no_model_deny_every_model(): agent_key: Final = UserAPIKeyAuth(token="agent-token", agent_id="agent-1", models=[]) @@ -9437,6 +9811,18 @@ async def test_agent_key_acting_for_a_teamless_user_is_capped_at_that_users_mode assert asked == ["team:None", "user:alice", "team:None", "user:alice"] +@pytest.mark.asyncio +async def test_agent_key_alias_resolves_against_the_echoed_teams_models(): + agent_key: Final = _agent_key_acting_for(user_id="alice", team_id="team-a") + agent_key.aliases = {"foo": "bar"} + load_team, load_user, asked = _caller_loaders(LiteLLM_TeamTable(team_id="team-a", models=["bar"]), None) + cache: Final = await _cache_with_membership("alice", "team-a", allowed_models=None) + + await _check_caller_models(agent_key, "foo", load_team, load_user, cache) + + assert asked == ["team:team-a"] + + @pytest.mark.asyncio async def test_agent_key_without_an_echoed_caller_keeps_its_own_models(): agent_key: Final = UserAPIKeyAuth(token="agent-token", agent_id="agent-1", models=["gpt-5", "claude-sonnet"]) From 993a5b9d978d432f0df6c4657b697a0c3cb2757a Mon Sep 17 00:00:00 2001 From: "berriai-litellm-provider-info-sync[bot]" <328147090+berriai-litellm-provider-info-sync[bot]@users.noreply.github.com> Date: Thu, 24 Sep 2026 19:52:16 -0700 Subject: [PATCH 02/10] chore(cost-map): add azure retirement dates for command-r-plus and gpt-4 (#43117) Price-Sync: litellm-providers Co-authored-by: berriai-litellm-provider-info-sync[bot] <328147090+berriai-litellm-provider-info-sync[bot]@users.noreply.github.com> --- litellm/model_prices_and_context_window_backup.json | 4 ++++ model_prices_and_context_window.json | 4 ++++ 2 files changed, 8 insertions(+) diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index b529a225f84..efc0e0e2877 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -3342,6 +3342,7 @@ "supports_vision": true }, "azure/command-r-plus": { + "deprecation_date": "2025-06-30", "input_cost_per_token": 3e-06, "litellm_provider": "azure", "max_input_tokens": 128000, @@ -3349,6 +3350,7 @@ "max_tokens": 4096, "mode": "chat", "output_cost_per_token": 1.5e-05, + "source": "https://learn.microsoft.com/en-us/azure/foundry/openai/concepts/retired-models", "supports_function_calling": true }, "azure_ai/claude-haiku-4-5": { @@ -5194,6 +5196,7 @@ "output_cost_per_token": 2e-06 }, "azure/gpt-4": { + "deprecation_date": "2025-06-06", "input_cost_per_token": 3e-05, "litellm_provider": "azure", "max_input_tokens": 8192, @@ -5201,6 +5204,7 @@ "max_tokens": 4096, "mode": "chat", "output_cost_per_token": 6e-05, + "source": "https://learn.microsoft.com/en-us/azure/foundry/openai/concepts/retired-models", "supports_function_calling": true, "supports_tool_choice": true }, diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index b529a225f84..efc0e0e2877 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -3342,6 +3342,7 @@ "supports_vision": true }, "azure/command-r-plus": { + "deprecation_date": "2025-06-30", "input_cost_per_token": 3e-06, "litellm_provider": "azure", "max_input_tokens": 128000, @@ -3349,6 +3350,7 @@ "max_tokens": 4096, "mode": "chat", "output_cost_per_token": 1.5e-05, + "source": "https://learn.microsoft.com/en-us/azure/foundry/openai/concepts/retired-models", "supports_function_calling": true }, "azure_ai/claude-haiku-4-5": { @@ -5194,6 +5196,7 @@ "output_cost_per_token": 2e-06 }, "azure/gpt-4": { + "deprecation_date": "2025-06-06", "input_cost_per_token": 3e-05, "litellm_provider": "azure", "max_input_tokens": 8192, @@ -5201,6 +5204,7 @@ "max_tokens": 4096, "mode": "chat", "output_cost_per_token": 6e-05, + "source": "https://learn.microsoft.com/en-us/azure/foundry/openai/concepts/retired-models", "supports_function_calling": true, "supports_tool_choice": true }, From d0d3b6a67e0af031cbb2e50b421f74c29d8a47e0 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Thu, 24 Sep 2026 20:04:33 -0700 Subject: [PATCH 03/10] fix(vertex_ai): stop advertising OpenAI platform-only params on Gemma and Llama routes (#43079) * fix(vertex_ai): stop advertising OpenAI platform-only params on Gemma and Llama routes The Anthropic /v1/messages bridge derives prompt_cache_key from Claude Code's session id whenever the provider config advertises it, and every Vertex OpenAI-compatible route (gemma/, openai/, meta/) inherited the full OpenAI list, so the Model Garden vLLM container rejected each turn with a pydantic extra_forbidden 400. Vertex's Llama and Gemma configs now filter one shared list of platform-only params (prompt_cache_key, prompt_cache_retention, safety_identifier, service_tier, store, web_search_options, modalities, prediction, audio, max_retries) out of their supported params, so the bridge no longer derives the key and drop_params drops an explicit one. * fix(vertex_ai): scope the platform-param filter to self-deployed Model Garden endpoints --------- Co-authored-by: mateo-berri <277851410+mateo-berri@users.noreply.github.com> --- litellm/llms/vertex_ai/common_utils.py | 36 ++++++++ .../llama3/transformation.py | 22 +++-- .../vertex_gemma_models/transformation.py | 8 ++ .../vertex_ai/vertex_model_garden/main.py | 21 ++--- ...al_pass_through_adapters_transformation.py | 15 ++++ .../test_vertex_model_garden_openapi.py | 13 ++- ...ai_partner_models_llama3_transformation.py | 79 ++++++++++++++++++ .../test_vertex_gemma_transformation.py | 83 +++++++++++++++++++ 8 files changed, 247 insertions(+), 30 deletions(-) diff --git a/litellm/llms/vertex_ai/common_utils.py b/litellm/llms/vertex_ai/common_utils.py index 14aebcaabaf..6d050d5a856 100644 --- a/litellm/llms/vertex_ai/common_utils.py +++ b/litellm/llms/vertex_ai/common_utils.py @@ -25,6 +25,21 @@ from litellm.types.llms.vertex_ai import ( from litellm.types.utils import TokenCountResponse from litellm.utils import supports_response_schema, supports_system_messages +VERTEX_SELF_DEPLOYED_ENDPOINT_UNSUPPORTED_PARAMS: Final = frozenset( + { + "audio", + "max_retries", + "modalities", + "prediction", + "prompt_cache_key", + "prompt_cache_retention", + "safety_identifier", + "service_tier", + "store", + "web_search_options", + } +) + class VertexAILyriaModelInfo(TypedDict): vertex_ai_audio_api: ReadOnly[Literal["lyria_predict", "lyria_interactions"]] @@ -370,6 +385,27 @@ def get_vertex_base_model_name(model: str) -> str: return model +def vertex_model_garden_model_id_in_json_body(model: str) -> bool: + """ + Vertex catalog / publisher models are addressed as publisher/model (e.g. + xai/grok-4.1-fast-reasoning) on the shared OpenAPI URL, with the id in the JSON body. + + Deployed Model Garden endpoints are typically a single segment (often numeric) + and use .../endpoints/{ENDPOINT_ID}/chat/completions with an empty model field. + """ + return "/" in model + + +def is_vertex_self_deployed_openai_compatible_endpoint(model: str) -> bool: + local_model: Final = model.removeprefix("vertex_ai/") + route: Final = get_vertex_ai_model_route(local_model) + if route == VertexAIModelRoute.GEMMA: + return True + return route == VertexAIModelRoute.MODEL_GARDEN and not vertex_model_garden_model_id_in_json_body( + get_vertex_base_model_name(local_model) + ) + + def get_vertex_ai_fine_tuned_endpoint_id(model: str) -> str | None: """ Fine-tuned Gemini deployments are addressed by a numeric endpoint id, diff --git a/litellm/llms/vertex_ai/vertex_ai_partner_models/llama3/transformation.py b/litellm/llms/vertex_ai/vertex_ai_partner_models/llama3/transformation.py index f2d2c0896d2..ca0bcb74906 100644 --- a/litellm/llms/vertex_ai/vertex_ai_partner_models/llama3/transformation.py +++ b/litellm/llms/vertex_ai/vertex_ai_partner_models/llama3/transformation.py @@ -18,7 +18,11 @@ from litellm.types.utils import ( Usage, ) -from ...common_utils import VertexAIError +from ...common_utils import ( + VERTEX_SELF_DEPLOYED_ENDPOINT_UNSUPPORTED_PARAMS, + VertexAIError, + is_vertex_self_deployed_openai_compatible_endpoint, +) if TYPE_CHECKING: from litellm.litellm_core_utils.tokenizer import Encoding as Tokenizer @@ -66,13 +70,15 @@ class VertexAILlama3Config(OpenAIGPTConfig): and v is not None } - def get_supported_openai_params(self, model: str): - supported_params: Final = super().get_supported_openai_params(model=model) - try: - supported_params.remove("max_retries") - except KeyError: - pass - return supported_params + def get_supported_openai_params(self, model: str) -> list[str]: + unsupported_params: Final = ( + VERTEX_SELF_DEPLOYED_ENDPOINT_UNSUPPORTED_PARAMS + if is_vertex_self_deployed_openai_compatible_endpoint(model) + else frozenset({"max_retries"}) + ) + return [ # mutable-ok: get_optional_params extends the returned list with allowed_openai_params + param for param in super().get_supported_openai_params(model=model) if param not in unsupported_params + ] def map_openai_params( self, diff --git a/litellm/llms/vertex_ai/vertex_gemma_models/transformation.py b/litellm/llms/vertex_ai/vertex_gemma_models/transformation.py index a774dba6cf2..ea97f0a0a9a 100644 --- a/litellm/llms/vertex_ai/vertex_gemma_models/transformation.py +++ b/litellm/llms/vertex_ai/vertex_gemma_models/transformation.py @@ -21,6 +21,7 @@ from litellm.llms.custom_httpx.http_handler import ( _get_httpx_client, ) from litellm.llms.openai.chat.gpt_transformation import OpenAIGPTConfig +from litellm.llms.vertex_ai.common_utils import VERTEX_SELF_DEPLOYED_ENDPOINT_UNSUPPORTED_PARAMS from litellm.types.llms.openai import AllMessageValues from litellm.types.llms.vertex_ai_gemma import VertexGemmaContainerError from litellm.types.utils import ModelResponse @@ -49,6 +50,13 @@ class VertexGemmaConfig(OpenAIGPTConfig): def __init__(self) -> None: super().__init__() + def get_supported_openai_params(self, model: str) -> list[str]: + return [ # mutable-ok: get_optional_params extends the returned list with allowed_openai_params + param + for param in super().get_supported_openai_params(model=model) + if param not in VERTEX_SELF_DEPLOYED_ENDPOINT_UNSUPPORTED_PARAMS + ] + def should_fake_stream( self, model: str | None, diff --git a/litellm/llms/vertex_ai/vertex_model_garden/main.py b/litellm/llms/vertex_ai/vertex_model_garden/main.py index f5c9ac623a1..84907f01685 100644 --- a/litellm/llms/vertex_ai/vertex_model_garden/main.py +++ b/litellm/llms/vertex_ai/vertex_model_garden/main.py @@ -24,21 +24,14 @@ import httpx from litellm.llms.vertex_ai.common_utils import get_vertex_base_url from litellm.utils import ModelResponse -from ..common_utils import VertexAIError, get_vertex_base_model_name +from ..common_utils import ( + VertexAIError, + get_vertex_base_model_name, + vertex_model_garden_model_id_in_json_body, +) from ..vertex_llm_base import VertexBase -def _vertex_model_garden_model_id_in_json_body(model: str) -> bool: - """ - Vertex catalog / publisher models are addressed as publisher/model (e.g. - xai/grok-4.1-fast-reasoning) on the shared OpenAPI URL, with the id in the JSON body. - - Deployed Model Garden endpoints are typically a single segment (often numeric) - and use .../endpoints/{ENDPOINT_ID}/chat/completions with an empty model field. - """ - return "/" in model - - def create_vertex_url( vertex_location: str, vertex_project: str, @@ -48,7 +41,7 @@ def create_vertex_url( ) -> str: """Return the api base for vertex model garden (without /chat/completions).""" base_url: Final = get_vertex_base_url(vertex_location) - if _vertex_model_garden_model_id_in_json_body(model): + if vertex_model_garden_model_id_in_json_body(model): return f"{base_url}/v1/projects/{vertex_project}/locations/{vertex_location}/endpoints/openapi" return f"{base_url}/v1beta1/projects/{vertex_project}/locations/{vertex_location}/endpoints/{model}" @@ -124,7 +117,7 @@ class VertexAIModelGardenModels(VertexBase): ) # Publisher/catalog models: model id must be sent in the JSON body (OpenAPI route). # Single-segment endpoint ids: model is encoded in the URL path; body model stays empty. - if not _vertex_model_garden_model_id_in_json_body(model): + if not vertex_model_garden_model_id_in_json_body(model): model = "" return openai_like_chat_completions.completion( model=model, diff --git a/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_anthropic_experimental_pass_through_adapters_transformation.py b/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_anthropic_experimental_pass_through_adapters_transformation.py index d03174bc2c6..2fe22ba2620 100644 --- a/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_anthropic_experimental_pass_through_adapters_transformation.py +++ b/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_anthropic_experimental_pass_through_adapters_transformation.py @@ -1046,11 +1046,26 @@ def test_translate_anthropic_to_openai_sets_prompt_cache_key_for_azure(model: st assert openai_request["prompt_cache_key"] == "session-abc" +@pytest.mark.parametrize( + "model", + [ + "vertex_ai/meta/llama-4-maverick-17b-128e-instruct-maas", + "vertex_ai/moonshotai/kimi-k2-thinking-maas", + "vertex_ai/xai/grok-4.1-fast-non-reasoning", + ], +) +def test_translate_anthropic_to_openai_sets_prompt_cache_key_for_vertex_maas_models(model: str): + openai_request = _translate_with_metadata(model, {"user_id": CLAUDE_CODE_USER_ID}, "vertex_ai") + assert openai_request["prompt_cache_key"] == "session-abc" + + @pytest.mark.parametrize( "model, custom_llm_provider", [ ("gemini/gemini-2.5-pro", "gemini"), ("vertex_ai/gemini-2.5-pro", "vertex_ai"), + ("vertex_ai/gemma/gemma-2-2b-it", "vertex_ai"), + ("vertex_ai/openai/mg-endpoint-lit8592", "vertex_ai"), ("anthropic/claude-sonnet-4-5", "anthropic"), ("bedrock/anthropic.claude-3-5-sonnet-20240620-v1:0", "bedrock"), ("no-such-model-lit5875", "no-such-provider-lit5875"), diff --git a/tests/test_litellm/llms/vertex_ai/test_vertex_model_garden_openapi.py b/tests/test_litellm/llms/vertex_ai/test_vertex_model_garden_openapi.py index 0dcaa4c72c2..b80d4714253 100644 --- a/tests/test_litellm/llms/vertex_ai/test_vertex_model_garden_openapi.py +++ b/tests/test_litellm/llms/vertex_ai/test_vertex_model_garden_openapi.py @@ -8,10 +8,10 @@ from unittest.mock import AsyncMock, MagicMock, patch import pytest import litellm -from litellm.llms.vertex_ai.vertex_model_garden.main import ( - _vertex_model_garden_model_id_in_json_body, - create_vertex_url, +from litellm.llms.vertex_ai.common_utils import ( + vertex_model_garden_model_id_in_json_body, ) +from litellm.llms.vertex_ai.vertex_model_garden.main import create_vertex_url @pytest.mark.parametrize( @@ -43,11 +43,8 @@ def test_create_vertex_url_openapi_vs_deployed_endpoint( def test_model_id_in_json_body_heuristic() -> None: - assert ( - _vertex_model_garden_model_id_in_json_body("xai/grok-4.1-fast-reasoning") - is True - ) - assert _vertex_model_garden_model_id_in_json_body("5464397967697903616") is False + assert vertex_model_garden_model_id_in_json_body("xai/grok-4.1-fast-reasoning") is True + assert vertex_model_garden_model_id_in_json_body("5464397967697903616") is False @pytest.fixture diff --git a/tests/unit/llms/vertex_ai/vertex_ai_partner_models/llama3/test_vertex_ai_partner_models_llama3_transformation.py b/tests/unit/llms/vertex_ai/vertex_ai_partner_models/llama3/test_vertex_ai_partner_models_llama3_transformation.py index 3bca51ec6b3..05e4e36edd7 100644 --- a/tests/unit/llms/vertex_ai/vertex_ai_partner_models/llama3/test_vertex_ai_partner_models_llama3_transformation.py +++ b/tests/unit/llms/vertex_ai/vertex_ai_partner_models/llama3/test_vertex_ai_partner_models_llama3_transformation.py @@ -11,6 +11,39 @@ from litellm.llms.vertex_ai.vertex_ai_partner_models.llama3.transformation impor ) +OPENAI_PLATFORM_PARAMS = ( + "prompt_cache_key", + "prompt_cache_retention", + "safety_identifier", + "service_tier", + "store", + "web_search_options", + "modalities", + "prediction", + "audio", +) + +SELF_DEPLOYED_ENDPOINT_MODELS = ( + "gemma/gemma-2-2b-it", + "vertex_ai/gemma/gemma-2-2b-it", + "openai/mg-endpoint-lit8592", + "vertex_ai/openai/mg-endpoint-lit8592", + "openai/5464397967697903616", +) + +MAAS_MODELS = ( + "meta/llama-4-maverick-17b-128e-instruct-maas", + "vertex_ai/meta/llama-4-maverick-17b-128e-instruct-maas", + "moonshotai/kimi-k2-thinking-maas", + "qwen/qwen3-next-80b-a3b-instruct-maas", + "google/gemma-4-26b-a4b-it-maas", + "xai/grok-4.1-fast-non-reasoning", + "openai/xai/grok-4.1-fast-reasoning", + "1984786713414729728", + "llama3", +) + + class TestVertexAILlama3Config: def test_transform_choices(self): """ @@ -56,6 +89,52 @@ class TestVertexAILlama3Config: assert response[0].message.tool_calls is not None assert response[0].finish_reason == "tool_calls" + @pytest.mark.parametrize("model", SELF_DEPLOYED_ENDPOINT_MODELS) + @pytest.mark.parametrize("param", OPENAI_PLATFORM_PARAMS) + def test_get_supported_openai_params_omits_platform_params_for_self_deployed_endpoints( + self, model: str, param: str + ): + assert param not in VertexAILlama3Config().get_supported_openai_params(model=model) + + @pytest.mark.parametrize("model", MAAS_MODELS) + @pytest.mark.parametrize("param", OPENAI_PLATFORM_PARAMS) + def test_get_supported_openai_params_keeps_platform_params_for_maas_models(self, model: str, param: str): + assert param in VertexAILlama3Config().get_supported_openai_params(model=model) + + @pytest.mark.parametrize("model", [*SELF_DEPLOYED_ENDPOINT_MODELS, *MAAS_MODELS]) + def test_get_supported_openai_params_never_lists_max_retries(self, model: str): + assert "max_retries" not in VertexAILlama3Config().get_supported_openai_params(model=model) + + @pytest.mark.parametrize("model", [*SELF_DEPLOYED_ENDPOINT_MODELS, *MAAS_MODELS]) + @pytest.mark.parametrize( + "param", + ["max_completion_tokens", "tools", "tool_choice", "response_format", "seed", "logprobs", "parallel_tool_calls"], + ) + def test_get_supported_openai_params_keeps_params_every_vertex_openai_endpoint_accepts( + self, model: str, param: str + ): + assert param in VertexAILlama3Config().get_supported_openai_params(model=model) + + @pytest.mark.parametrize("model", SELF_DEPLOYED_ENDPOINT_MODELS) + def test_map_openai_params_drops_prompt_cache_key_for_self_deployed_endpoints(self, model: str): + mapped = VertexAILlama3Config().map_openai_params( + {"prompt_cache_key": "session-lit8592", "max_completion_tokens": 10}, + {}, + model, + drop_params=True, + ) + assert mapped == {"max_tokens": 10} + + @pytest.mark.parametrize("model", MAAS_MODELS) + def test_map_openai_params_forwards_prompt_cache_key_for_maas_models(self, model: str): + mapped = VertexAILlama3Config().map_openai_params( + {"prompt_cache_key": "session-lit8592", "max_completion_tokens": 10}, + {}, + model, + drop_params=True, + ) + assert mapped == {"prompt_cache_key": "session-lit8592", "max_tokens": 10} + class TestVertexAILlama3StreamingHandler: def test_first_chunk_has_role_assistant_when_missing(self): diff --git a/tests/unit/llms/vertex_ai/vertex_gemma_models/test_vertex_gemma_transformation.py b/tests/unit/llms/vertex_ai/vertex_gemma_models/test_vertex_gemma_transformation.py index e9ae5234094..e5ca31833ce 100644 --- a/tests/unit/llms/vertex_ai/vertex_gemma_models/test_vertex_gemma_transformation.py +++ b/tests/unit/llms/vertex_ai/vertex_gemma_models/test_vertex_gemma_transformation.py @@ -694,6 +694,89 @@ class TestVertexGemmaCompletion: assert instance["@requestFormat"] == "chatCompletions" assert "messages" in instance + @pytest.mark.parametrize( + "param", + [ + "prompt_cache_key", + "prompt_cache_retention", + "safety_identifier", + "service_tier", + "store", + "web_search_options", + "modalities", + "prediction", + "audio", + "max_retries", + ], + ) + def test_get_supported_openai_params_omits_params_the_predict_endpoint_rejects(self, param: str): + from litellm.llms.vertex_ai.vertex_gemma_models.transformation import ( + VertexGemmaConfig, + ) + + assert param not in VertexGemmaConfig().get_supported_openai_params(model="gemma-2-2b-it") + + @pytest.mark.asyncio + async def test_acompletion_drops_prompt_cache_key_when_drop_params_is_set(self): + with ( + patch("litellm.llms.custom_httpx.http_handler.get_async_httpx_client") as mock_get_client, + patch( + "litellm.llms.vertex_ai.vertex_gemma_models.main.VertexAIGemmaModels._ensure_access_token", + return_value=("fake-access-token", "PROJECT_ID"), + ), + ): + mock_client = Mock() + mock_response = Mock() + mock_response.status_code = 200 + mock_response.json.return_value = _make_gemma_vertex_response() + mock_client.post = AsyncMock(return_value=mock_response) + mock_get_client.return_value = mock_client + + await litellm.acompletion( + model="vertex_ai/gemma/gemma-2-2b-it", + messages=[{"role": "user", "content": "Test"}], + prompt_cache_key="session-lit8592", + service_tier="default", + max_completion_tokens=16, + drop_params=True, + api_base="https://test.us-central1-project.prediction.vertexai.goog/v1/projects/PROJECT_ID/locations/us-central1/endpoints/ENDPOINT_ID:predict", + vertex_project="PROJECT_ID", + vertex_location="us-central1", + ) + + instance = mock_client.post.call_args.kwargs["json"]["instances"][0] + assert "prompt_cache_key" not in instance + assert "service_tier" not in instance + assert instance["max_tokens"] == 16 + assert instance["messages"] == [{"role": "user", "content": "Test"}] + + @pytest.mark.asyncio + async def test_acompletion_rejects_prompt_cache_key_before_calling_vertex(self, monkeypatch: pytest.MonkeyPatch): + monkeypatch.setattr(litellm, "drop_params", False) + with ( + patch("litellm.llms.custom_httpx.http_handler.get_async_httpx_client") as mock_get_client, + patch( + "litellm.llms.vertex_ai.vertex_gemma_models.main.VertexAIGemmaModels._ensure_access_token", + return_value=("fake-access-token", "PROJECT_ID"), + ), + ): + mock_client = Mock() + mock_client.post = AsyncMock() + mock_get_client.return_value = mock_client + + with pytest.raises(litellm.UnsupportedParamsError, match="prompt_cache_key"): + await litellm.acompletion( + model="vertex_ai/gemma/gemma-2-2b-it", + messages=[{"role": "user", "content": "Test"}], + prompt_cache_key="session-lit8592", + drop_params=False, + api_base="https://test.us-central1-project.prediction.vertexai.goog/v1/projects/PROJECT_ID/locations/us-central1/endpoints/ENDPOINT_ID:predict", + vertex_project="PROJECT_ID", + vertex_location="us-central1", + ) + + mock_client.post.assert_not_called() + def test_transform_request_strips_context_management(self): """ Direct unit test for VertexGemmaConfig.transform_request: verify that From f3cf1cdfefa51059cbdd6a3d86e0fec6468d189f Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Thu, 24 Sep 2026 20:06:53 -0700 Subject: [PATCH 04/10] fix(router): honor disable_fallbacks on mid-stream fallback (#43111) * fix(router): honor disable_fallbacks on mid-stream fallback The mid-stream fallback hop on chat, Responses, and Messages hardcoded disable_fallbacks=False, so a request or key that opted out of fallbacks still got a fallback deployment's answer when the primary's stream died before its first chunk. Each hop now reads the opt-out from the request the way the pre-stream path does, and the sync chat stream re-raises the primary's error instead of re-entering the fallback chain * test(router): prove disable_fallbacks reaches the mid-stream hop through the public entrypoints --------- Co-authored-by: mateo-berri <277851410+mateo-berri@users.noreply.github.com> --- litellm/router.py | 11 +- tests/test_litellm/test_router.py | 189 ++++++++++++++++++++++++++++++ 2 files changed, 195 insertions(+), 5 deletions(-) diff --git a/litellm/router.py b/litellm/router.py index 9bf8c410bcb..ee77aa45656 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -2940,7 +2940,7 @@ class Router: self._update_kwargs_before_fallbacks(model=model_group, kwargs=initial_kwargs) fallback_response = await self.async_function_with_fallbacks_common_utils( e=e, - disable_fallbacks=False, + disable_fallbacks=fallbacks_disabled_for_request(initial_kwargs), fallbacks=fallbacks, context_window_fallbacks=context_window_fallbacks, content_policy_fallbacks=content_policy_fallbacks, @@ -3384,7 +3384,7 @@ class Router: ) fallback_response = await self.async_function_with_fallbacks_common_utils( e=fallback_trigger, - disable_fallbacks=False, + disable_fallbacks=fallbacks_disabled_for_request(initial_kwargs), fallbacks=fallbacks, context_window_fallbacks=context_window_fallbacks, content_policy_fallbacks=content_policy_fallbacks, @@ -3475,8 +3475,9 @@ class Router: for item in model_response: yield item except MidStreamFallbackError as e: - if not e.is_pre_first_chunk and ( - e.generated_content or _stream_chunks_have_generated_content(model_response.chunks) + if fallbacks_disabled_for_request(initial_kwargs) or ( + not e.is_pre_first_chunk + and (e.generated_content or _stream_chunks_have_generated_content(model_response.chunks)) ): if e.original_exception is not None: raise e.original_exception from e @@ -5611,7 +5612,7 @@ class Router: ) fallback_response = await self.async_function_with_fallbacks_common_utils( # rebind-ok: set on success e=fallback_trigger, - disable_fallbacks=False, + disable_fallbacks=fallbacks_disabled_for_request(initial_kwargs), fallbacks=fallbacks, context_window_fallbacks=context_window_fallbacks, content_policy_fallbacks=content_policy_fallbacks, diff --git a/tests/test_litellm/test_router.py b/tests/test_litellm/test_router.py index 82122da15dc..80131534183 100644 --- a/tests/test_litellm/test_router.py +++ b/tests/test_litellm/test_router.py @@ -25,6 +25,7 @@ from litellm import Router from litellm.caching.caching import DualCache from litellm.caching.redis_cache import _redis_circuit_breaker_guard from litellm.exceptions import MidStreamFallbackError +from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper from litellm.integrations.custom_guardrail import CustomGuardrail from litellm.integrations.custom_logger import CustomLogger from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging @@ -49,6 +50,7 @@ from litellm.router import ( from litellm.router_strategy import simple_shuffle from litellm.router_utils.client_initalization_utils import MaxParallelRequestsLimit from litellm.router_utils.cooldown_handlers import _async_get_cooldown_deployments +from litellm.router_utils.fallback_event_handlers import DISABLE_FALLBACKS_METADATA_KEY from litellm.router_utils.router_callbacks.track_deployment_metrics import get_deployment_successes_for_current_minute from litellm.types.llms.openai import ChatCompletionRequest from litellm.types.router import Deployment, DeploymentTypedDict, LiteLLM_Params, ModelInfo, PreRoutingHookResponse, RetryPolicy @@ -14392,6 +14394,193 @@ async def test_anthropic_messages_fallback_also_catches_raised_midstream_error() assert mock_fallback.await_args.kwargs["e"] is raised_error +_MID_STREAM_OPT_OUT_SHAPES: Final = ( + pytest.param({"disable_fallbacks": True}, id="raw-kwarg"), + pytest.param({"metadata": {DISABLE_FALLBACKS_METADATA_KEY: True}}, id="metadata-stamp"), + pytest.param({"litellm_metadata": {DISABLE_FALLBACKS_METADATA_KEY: True}}, id="litellm_metadata-stamp"), +) + + +def _mid_stream_opt_out_router() -> Router: + return Router( + model_list=[ + {"model_name": "primary", "litellm_params": {"model": "openai/gpt-5.4", "api_key": "k1"}}, + {"model_name": "fallback", "litellm_params": {"model": "openai/gpt-5.4-mini", "api_key": "k2"}}, + ], + fallbacks=[{"primary": ["fallback"]}], + ) + + +def _mid_stream_opt_out_primary_error() -> litellm.InternalServerError: + return litellm.InternalServerError(message="primary failed at stream start", llm_provider="openai", model="primary") + + +def _mid_stream_opt_out_trigger(primary_error: Exception) -> MidStreamFallbackError: + return MidStreamFallbackError( + message=str(primary_error), + model="primary", + llm_provider="openai", + original_exception=primary_error, + is_pre_first_chunk=True, + ) + + +class _MidStreamOptOutChatStream(CustomStreamWrapper): + """A chat deployment stream, as the router sees one, that dies before its first chunk.""" + + def __init__(self, error: Exception, model: str = "primary") -> None: + super().__init__(completion_stream=object(), model=model, custom_llm_provider="openai", logging_obj=MagicMock()) + self._error: Final = error + + def __aiter__(self): + return self + + async def __anext__(self) -> object: + raise self._error + + def __iter__(self): + return self + + def __next__(self) -> object: + raise self._error + + +@pytest.mark.asyncio +@pytest.mark.parametrize("opt_out", _MID_STREAM_OPT_OUT_SHAPES) +async def test_acompletion_streaming_iterator_honors_disable_fallbacks(opt_out): + """A chat stream that fails before its first chunk on a request that opted out of fallbacks + surfaces the primary's own error and never tries the fallback deployment.""" + router = _mid_stream_opt_out_router() + primary_error = _mid_stream_opt_out_primary_error() + source = _MidStreamOptOutChatStream(_mid_stream_opt_out_trigger(primary_error)) + + with patch.object(router, "_acompletion", new=AsyncMock(return_value=_AsyncList([]))) as fallback_attempt: + wrapped = await router._acompletion_streaming_iterator( + model_response=source, + messages=[{"role": "user", "content": "Hi"}], + initial_kwargs={"model": "primary", "stream": True, **copy.deepcopy(opt_out)}, + ) + with pytest.raises(litellm.InternalServerError) as raised: + [chunk async for chunk in wrapped] + + assert raised.value is primary_error + fallback_attempt.assert_not_awaited() + + +@pytest.mark.parametrize("opt_out", _MID_STREAM_OPT_OUT_SHAPES) +def test_completion_streaming_iterator_honors_disable_fallbacks(opt_out): + """Sync counterpart of test_acompletion_streaming_iterator_honors_disable_fallbacks.""" + router = _mid_stream_opt_out_router() + primary_error = _mid_stream_opt_out_primary_error() + source = _MidStreamOptOutChatStream(_mid_stream_opt_out_trigger(primary_error)) + + with patch.object(router, "_completion", new=MagicMock(return_value=iter([]))) as fallback_attempt: + wrapped = router._completion_streaming_iterator( + model_response=source, + messages=[{"role": "user", "content": "Hi"}], + initial_kwargs={"model": "primary", "stream": True, **copy.deepcopy(opt_out)}, + ) + with pytest.raises(litellm.InternalServerError) as raised: + list(wrapped) + + assert raised.value is primary_error + fallback_attempt.assert_not_called() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("opt_out", _MID_STREAM_OPT_OUT_SHAPES) +async def test_aresponses_streaming_iterator_honors_disable_fallbacks(opt_out): + """Same opt-out contract on the Responses API mid-stream fallback path.""" + router = _mid_stream_opt_out_router() + primary_error = _mid_stream_opt_out_primary_error() + source = _make_responses_iterator(error=_mid_stream_opt_out_trigger(primary_error), model="primary") + + with patch.object( + router, + "_ageneric_api_call_with_fallbacks_responses_attempt", + new=AsyncMock(return_value=_AsyncList([])), + ) as fallback_attempt: + wrapped = await router._aresponses_streaming_iterator( + response=source, + initial_kwargs={ + "model": "primary", + "stream": True, + "input": "Hi", + "original_generic_function": litellm.aresponses, + **copy.deepcopy(opt_out), + }, + ) + with pytest.raises(litellm.InternalServerError) as raised: + [chunk async for chunk in wrapped] + + assert raised.value is primary_error + fallback_attempt.assert_not_awaited() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("opt_out", _MID_STREAM_OPT_OUT_SHAPES) +async def test_anthropic_messages_streaming_iterator_honors_disable_fallbacks(opt_out): + """Same opt-out contract on the Anthropic Messages mid-stream fallback path.""" + router = _mid_stream_opt_out_router() + primary_error = _mid_stream_opt_out_primary_error() + source = _AnthropicMessagesRaisingByteStream([], _mid_stream_opt_out_trigger(primary_error)) + + with patch.object( + router, + "_ageneric_api_call_with_fallbacks_anthropic_messages_attempt", + new=AsyncMock(return_value=_AnthropicMessagesFallbackByteStream([])), + ) as fallback_attempt: + wrapped = await router._aanthropic_messages_streaming_iterator( + response=source, + initial_kwargs={"model": "primary", "stream": True, **copy.deepcopy(opt_out)}, + ) + with pytest.raises(litellm.InternalServerError) as raised: + [chunk async for chunk in wrapped] + + assert raised.value is primary_error + fallback_attempt.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_acompletion_disable_fallbacks_reaches_the_mid_stream_hop(): + """`disable_fallbacks=True` sent to the public entrypoint survives the fallback wrapper's handoff + into the stream: the primary's own error surfaces and no fallback deployment is ever called.""" + router = _mid_stream_opt_out_router() + primary_error = _mid_stream_opt_out_primary_error() + + async def primary_stream(**kwargs): + return _MidStreamOptOutChatStream(_mid_stream_opt_out_trigger(primary_error), model=kwargs["model"]) + + with patch("litellm.acompletion", side_effect=primary_stream) as provider_calls: + response = await router.acompletion( + model="primary", messages=[{"role": "user", "content": "Hi"}], stream=True, disable_fallbacks=True + ) + with pytest.raises(litellm.InternalServerError) as raised: + [chunk async for chunk in response] + + assert raised.value is primary_error + assert [call.kwargs["metadata"]["model_group"] for call in provider_calls.call_args_list] == ["primary"] + + +def test_completion_disable_fallbacks_reaches_the_mid_stream_hop(): + """Sync counterpart of test_acompletion_disable_fallbacks_reaches_the_mid_stream_hop.""" + router = _mid_stream_opt_out_router() + primary_error = _mid_stream_opt_out_primary_error() + + def primary_stream(**kwargs): + return _MidStreamOptOutChatStream(_mid_stream_opt_out_trigger(primary_error), model=kwargs["model"]) + + with patch("litellm.completion", side_effect=primary_stream) as provider_calls: + response = router.completion( + model="primary", messages=[{"role": "user", "content": "Hi"}], stream=True, disable_fallbacks=True + ) + with pytest.raises(litellm.InternalServerError) as raised: + list(response) + + assert raised.value is primary_error + assert [call.kwargs["metadata"]["model_group"] for call in provider_calls.call_args_list] == ["primary"] + + @pytest.mark.asyncio @pytest.mark.parametrize( "raised_error", From f61b3c3f38e7eacc6438ba956d8c072dae9132ca Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Thu, 24 Sep 2026 20:18:41 -0700 Subject: [PATCH 05/10] refactor(types): declare litellm-owned kwargs as typed objects and derive the lists from their fields (#42843) * refactor(types): declare litellm-owned params in one registry * refactor(types): re-export registry constants without redundant aliases * refactor(types): satisfy type-discipline rules in registry projections and tests * refactor(types): classify every registry entry and check groups against typed config models * test(types): pin load-bearing names and exact projections in registry tests * style(types): keep agentic projection comment within ruff format * refactor(types): declare litellm-owned params as typed objects and derive the lists from their fields * refactor(types): fields of the typed objects become the registry; tests use a hand-written inventory * refactor(types): split traversal into wire_names and owned_wire_names, move rust to kwarg artifacts rust is a module-level switch (litellm.rust) that nothing reads from a call's kwargs, so it joins self, use_client and model_config as a registered artifact instead of a DispatchOptions field. The field constants now import from litellm.types.litellm_params directly instead of through a re-export in litellm.types.utils. metadata and litellm_metadata are MutableMapping because their readers mutate them in place, and client accepts raw httpx clients * refactor(types): own max_agentic_loops as an option and walk only nested leaves Move max_agentic_loops from AgenticLoopState to a new AgenticLoopOptions leaf under LiteLLMOptions, since the interception handlers read it as a deployment ceiling rather than stamping it. Drop the owned_wire_names fallback that treated an unresolved annotation as a direct field, which under postponed annotations silently shrank the registry. Re-export TRUSTED_CALLBACK_VARS_FIELD and ADDRESSED_RESPONSE_ID_FIELD from types.utils so that import path keeps working. Tests use hand-written inventories for the callback and pricing names * refactor(types): move data_residency to call state and drop aliased re-exports data_residency is stamped by get_litellm_params and responses.main during the call, so it lives on CallState, not CostOptions. mock_response also accepts a float sequence, which main.py reads for mock embeddings. The types/utils.py re-exports become one plain import with an exact F401 suppression instead of two X as X aliases that pushed PLC0414 over its strict-gate ceiling. Redundant leaf docstrings and the structural artifact test are gone; the re-exported FIELD constants are checked by identity instead * refactor(types): project owned kwarg names once and keep pass-through extraction in request order * refactor(types): type caching_groups from its cache reader and hoist the pass-through ownership set caching_groups is a sequence of flat model-group sequences, which is what Cache._get_caching_group iterates. A regression test drives the public cache key path so two groups in one caching group share a key and a third does not. The pass-through endpoint builds its frozenset of owned names once at import instead of per request, reads the two metadata carriers from the extracted mapping instead of popping them, and its extraction mappings are read-only. Concatenation tests assert the whole derived list and tuple, docstrings drop reader claims that nothing in the module backs * refactor(types): read owned names live in pass-through and pin tests to literal inventories The pass-through endpoint checks body keys against the public all_litellm_params list at request time again, as the base does, instead of a frozenset taken at import, so a name registered after import is still extracted. A test drives that path with a name added after import, and another sends both metadata carriers interleaved with provider keys and asserts the whole merged result. retry_policy accepts the mapping form its router reader builds a RetryPolicy from. The pricing inventory in the typed tests is a literal tuple checked against the model's fields, the agentic compatibility test asserts type, length and set instead of declaration order, and the typed-model overlap tests assert the exact intersection. * refactor(types): move model_alias_map to CallState and read the owned registry in registry order in pass-through * refactor(types): drop restating docstrings, keep FIELD importers on types.utils, pin pass-through registry order * fix(types): satisfy strict lint for public FIELD re-exports * fix(types): restore clean parameter re-exports * fix(tests): compare pass-through extraction order to registry body keys * refactor(types): type owned request parameter leaves * refactor(types): share routing strategy literal and tighten leaf tests * fix(proxy): drop client-supplied proxy-stamped names from pass-through litellm_params * refactor(proxy): name pass-through litellm key split for what it holds * refactor(types): drop TODO markers on the kept readerless fields * fix(types): keep deployment tag_regex and max_file_size_mb out of provider requests * fix(types): include every routing strategy the router accepts --------- Co-authored-by: shrey kharbanda --- litellm/main.py | 13 +- .../pass_through_endpoints.py | 21 +- litellm/router.py | 11 +- litellm/types/integrations/custom_logger.py | 6 +- litellm/types/litellm_params.py | 364 ++++++++++ litellm/types/router.py | 8 +- litellm/types/utils.py | 216 +----- .../test_pass_through_endpoints.py | 158 ++++- tests/test_litellm/test_utils.py | 17 + tests/unit/types/test_litellm_params.py | 655 ++++++++++++++++++ 10 files changed, 1236 insertions(+), 233 deletions(-) create mode 100644 litellm/types/litellm_params.py create mode 100644 tests/unit/types/test_litellm_params.py diff --git a/litellm/main.py b/litellm/main.py index 98bb5126a90..72c9afad36c 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -126,6 +126,7 @@ from litellm.types.completion import ( _CompletionDispatchContext, _CompletionDispatchResult, ) +from litellm.types.litellm_params import RetryStrategy from litellm.types.router import GenericLiteLLMParams from litellm.types.utils import ( CustomPricingLiteLLMParams, @@ -6026,9 +6027,7 @@ def completion_with_retries(*args, **kwargs): # reset retries in .completion() kwargs["max_retries"] = 0 kwargs["num_retries"] = 0 - retry_strategy: Final[Literal["exponential_backoff_retry", "constant_retry"]] = kwargs.pop( - "retry_strategy", "constant_retry" - ) + retry_strategy: Final[RetryStrategy] = kwargs.pop("retry_strategy", "constant_retry") original_function: Final = kwargs.pop("original_function", completion) if retry_strategy == "exponential_backoff_retry": retryer = tenacity.Retrying( @@ -6054,7 +6053,7 @@ async def acompletion_with_retries(*args, **kwargs): num_retries: Final = kwargs.pop("num_retries", 3) kwargs["max_retries"] = 0 kwargs["num_retries"] = 0 - retry_strategy: Final = kwargs.pop("retry_strategy", "constant_retry") + retry_strategy: Final[RetryStrategy] = kwargs.pop("retry_strategy", "constant_retry") original_function: Final = kwargs.pop("original_function", completion) if retry_strategy == "exponential_backoff_retry": retryer = tenacity.AsyncRetrying( @@ -6082,9 +6081,7 @@ def responses_with_retries(*args, **kwargs): # reset retries in .responses() kwargs["max_retries"] = 0 kwargs["num_retries"] = 0 - retry_strategy: Final[Literal["exponential_backoff_retry", "constant_retry"]] = kwargs.pop( - "retry_strategy", "constant_retry" - ) + retry_strategy: Final[RetryStrategy] = kwargs.pop("retry_strategy", "constant_retry") original_function: Final = kwargs.pop("original_function", responses) if retry_strategy == "exponential_backoff_retry": retryer = tenacity.Retrying( @@ -6111,7 +6108,7 @@ async def aresponses_with_retries(*args, **kwargs): num_retries: Final = kwargs.pop("num_retries", 3) kwargs["max_retries"] = 0 kwargs["num_retries"] = 0 - retry_strategy: Final = kwargs.pop("retry_strategy", "constant_retry") + retry_strategy: Final[RetryStrategy] = kwargs.pop("retry_strategy", "constant_retry") original_function: Final = kwargs.pop("original_function", aresponses) if retry_strategy == "exponential_backoff_retry": retryer = tenacity.AsyncRetrying( diff --git a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py index 7d6db30e3e3..e0a4184291e 100644 --- a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py @@ -105,6 +105,8 @@ from litellm.proxy.route_llm_request import ProxyModelNotFoundError from litellm.proxy.utils import normalize_route_for_root_path from litellm.repositories.team_repository import TeamRepository from litellm.secret_managers.main import get_secret_str +from litellm.types import utils as types_utils +from litellm.types.litellm_params import ProxyRequestState, wire_names from litellm.types.llms.custom_http import httpxSpecialProvider from litellm.types.passthrough_endpoints.pass_through_endpoints import ( LITELLM_PASS_THROUGH_CUSTOM_BODY_STATE_KEY, @@ -133,6 +135,9 @@ router: Final = APIRouter() pass_through_endpoint_logging: Final = PassThroughEndpointLogging() +_METADATA_KEYS: Final = frozenset(("litellm_metadata", "metadata")) +_KEPT_OUT_OF_LITELLM_PARAMS: Final = _METADATA_KEYS | frozenset(wire_names(ProxyRequestState)) + # Global registry to track registered pass-through routes and prevent memory leaks _registered_pass_through_routes: Final[dict[str, dict[str, str | bool | list[str] | Mapping[str, object]]]] = {} @@ -578,21 +583,21 @@ class HttpPassThroughEndpointHelpers(BasePassthroughUtils): """ Filter out litellm params from the request body """ - from litellm.types.utils import all_litellm_params - _parsed_body = _parsed_body or {} - litellm_params_in_body: Final = {} - for k in all_litellm_params: - if k in _parsed_body: - litellm_params_in_body[k] = _parsed_body.pop(k, None) + litellm_keys_in_body: Final = MappingProxyType( + {k: _parsed_body.pop(k) for k in types_utils.all_litellm_params if k in _parsed_body} + ) + litellm_params_in_body: Final = MappingProxyType( + {k: v for k, v in litellm_keys_in_body.items() if k not in _KEPT_OUT_OF_LITELLM_PARAMS} + ) _metadata = dict( LiteLLMProxyRequestSetup.get_sanitized_user_information_from_key(user_api_key_dict=user_api_key_dict) ) - litellm_metadata: Final = litellm_params_in_body.pop("litellm_metadata", None) - metadata: Final = litellm_params_in_body.pop("metadata", None) + litellm_metadata: Final = litellm_keys_in_body.get("litellm_metadata") + metadata: Final = litellm_keys_in_body.get("metadata") if litellm_metadata: _metadata.update(litellm_metadata) if metadata: diff --git a/litellm/router.py b/litellm/router.py index ee77aa45656..023b99cd64e 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -259,6 +259,7 @@ from litellm.router_utils.routing_groups import ( validate_routing_strategy, ) from litellm.scheduler import FlowItem, Scheduler +from litellm.types.litellm_params import RoutingStrategyName from litellm.types.llms.openai import ( AllMessageValues, ChatCompletionToolParam, @@ -796,15 +797,7 @@ class Router: allowed_fails_policy: AllowedFailsPolicy | None = None, # set custom allowed fails policy cooldown_time: float | None = None, # (seconds) time to cooldown a deployment after failure disable_cooldowns: bool | None = None, - routing_strategy: Literal[ - "simple-shuffle", - "least-busy", - "usage-based-routing", - "latency-based-routing", - "cost-based-routing", - "usage-based-routing-v2", - "lar1", - ] = "simple-shuffle", + routing_strategy: RoutingStrategyName = "simple-shuffle", optional_pre_call_checks: OptionalPreCallChecks | None = None, routing_strategy_args: dict = {}, # just for latency-based routing_groups: list[RoutingGroup | dict] | None = None, diff --git a/litellm/types/integrations/custom_logger.py b/litellm/types/integrations/custom_logger.py index 5de58a20242..9a9f3ae34ce 100644 --- a/litellm/types/integrations/custom_logger.py +++ b/litellm/types/integrations/custom_logger.py @@ -3,8 +3,10 @@ from typing import Any, Final from pydantic import BaseModel, Field -CHAT_COMPLETION_AGENTIC_SURFACE: Final = "chat_completions" -RESPONSES_AGENTIC_SURFACE: Final = "responses" +from litellm.types.litellm_params import AgenticSurface + +CHAT_COMPLETION_AGENTIC_SURFACE: Final[AgenticSurface] = "chat_completions" +RESPONSES_AGENTIC_SURFACE: Final[AgenticSurface] = "responses" CODE_INTERPRETER_INTERCEPTION_PREFIX: Final = "_code_interpreter_interception" HEADROOM_INTERCEPTION_PREFIX: Final = "_headroom_interception" HEADROOM_CONVERTED_STREAM_KEY: Final = f"{HEADROOM_INTERCEPTION_PREFIX}_converted_stream" diff --git a/litellm/types/litellm_params.py b/litellm/types/litellm_params.py new file mode 100644 index 00000000000..83a42c235f9 --- /dev/null +++ b/litellm/types/litellm_params.py @@ -0,0 +1,364 @@ +"""LiteLLM-owned request kwargs declared as typed fields; types/utils.py splices these with the callback and pricing +models and KWARG_ARTIFACTS into all_litellm_params.""" + +from collections.abc import Callable, Iterator, Mapping, MutableMapping, Sequence +from dataclasses import dataclass, field, fields, is_dataclass +from types import MappingProxyType +from typing import TYPE_CHECKING, Final, Literal, TypeAlias + +if TYPE_CHECKING: + import httpx + from aiohttp import ClientSession + from openai import AsyncAzureOpenAI, AsyncOpenAI, AzureOpenAI, OpenAI + + from litellm.litellm_core_utils.litellm_logging import Logging + from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler + from litellm.router_strategy.complexity_router.context_compaction import CompactionState + from litellm.router_utils.fallback_event_handlers import AttemptedFallbackTargets + from litellm.types.caching import DynamicCacheControl + from litellm.types.llms.openai import ChatCompletionAssistantMessage, ChatCompletionUserMessage + from litellm.types.proxy.litellm_pre_call_utils import SecretFields + from litellm.types.router import ConfigurableClientsideParamsCustomAuth, DeploymentTypedDict, RetryPolicy + from litellm.types.router_weights import RouterWeights + from litellm.types.utils import ModelResponse, ModelResponseStream, ProviderSpecificHeader + + ProviderClient: TypeAlias = ( + OpenAI + | AsyncOpenAI + | AzureOpenAI + | AsyncAzureOpenAI + | HTTPHandler + | AsyncHTTPHandler + | httpx.Client + | httpx.AsyncClient + ) + MockResponse: TypeAlias = ( + str | Exception | Mapping[str, object] | Sequence[float] | ModelResponse | ModelResponseStream + ) + +RetryStrategy: TypeAlias = Literal["constant_retry", "exponential_backoff_retry"] +AgenticSurface: TypeAlias = Literal["chat_completions", "responses"] +RoutingStrategyName: TypeAlias = Literal[ + "simple-shuffle", + "least-busy", + "usage-based-routing", + "latency-based-routing", + "cost-based-routing", + "usage-based-routing-v2", + "lar1", +] + +TRUSTED_CALLBACK_VARS_FIELD: Final = "litellm_trusted_callback_vars" +ADDRESSED_RESPONSE_ID_FIELD: Final = "_litellm_addressed_response_id" + +WIRE_NAME: Final = "wire_name" + + +def wire(name: str) -> Mapping[str, str]: + return MappingProxyType({WIRE_NAME: name}) + + +@dataclass(frozen=True, slots=True, kw_only=True) +class ProviderConnection: + api_key: str | None = None + api_base: str | None = None + api_version: str | None = None + region_name: str | None = None + headers: Mapping[str, str] | None = None + provider_specific_header: "ProviderSpecificHeader | Sequence[ProviderSpecificHeader] | None" = None + client: "ProviderClient | None" = None + shared_session: "ClientSession | None" = None + ssl_verify: bool | str | None = None + request_timeout: float | None = None + force_timeout: float | None = None + stream_timeout: float | str | None = None + max_retries: int | None = None + tenant_id: str | None = None + client_id: str | None = None + client_secret: str | None = None + azure_username: str | None = None + azure_password: str | None = None + azure_scope: str | None = None + azure_ad_token_provider: Callable[[], str] | None = None + litellm_credential_name: str | None = None + configurable_clientside_auth_params: "Sequence[str | ConfigurableClientsideParamsCustomAuth] | None" = None + use_xai_oauth: bool | None = None + + +@dataclass(frozen=True, slots=True, kw_only=True) +class BedrockBatchConnection: + # Bedrock rejects these names in request bodies, so register them as LiteLLM-owned + aws_batch_role_arn: str | None = None + s3_bucket_name: str | None = None + s3_region_name: str | None = None + s3_endpoint_url: str | None = None + s3_output_bucket_name: str | None = None + s3_bucket_owner: str | None = None + s3_access_key_id: str | None = None + s3_secret_access_key: str | None = None + s3_encryption_key_id: str | None = None + bedrock_tags: Sequence[Mapping[str, str]] | None = None + + +@dataclass(frozen=True, slots=True, kw_only=True) +class ConnectionSettings: + provider: ProviderConnection + bedrock_batch: BedrockBatchConnection + + +@dataclass(frozen=True, slots=True, kw_only=True) +class DispatchOptions: + custom_llm_provider: str | None = None + azure: bool | None = None + use_litellm_proxy: bool | None = None + use_chat_completions_api: bool | None = None + use_in_pass_through: bool | None = None + allowed_openai_params: Sequence[str] | None = None + + +@dataclass(frozen=True, slots=True, kw_only=True) +class RoutingOptions: + fallbacks: Sequence[str | Mapping[str, object]] | None = None + context_window_fallback_dict: Mapping[str, str] | None = None + num_retries: int | None = None + retry_policy: "RetryPolicy | Mapping[str, object] | None" = None + retry_strategy: RetryStrategy | None = None + routing_strategy: RoutingStrategyName | None = None + cooldown_time: float | None = None + allowed_model_region: str | None = None + enable_tag_filtering: bool | None = None + fastest_response: bool | None = None + provider_affinity_header: str | None = None + search_tool_name: str | None = None + model_list: "Sequence[DeploymentTypedDict] | None" = None + + +@dataclass(frozen=True, slots=True, kw_only=True) +class DeploymentOptions: + model_info: Mapping[str, object] | None = None + rpm: int | None = None + tpm: int | None = None + itpm: int | None = None + otpm: int | None = None + default_api_key_rpm_limit: int | None = None + default_api_key_tpm_limit: int | None = None + max_parallel_requests: int | None = None + weight: int | None = None + order: int | None = None + tag_regex: Sequence[str] | None = None + max_file_size_mb: float | None = None + + +@dataclass(frozen=True, slots=True, kw_only=True) +class SpecializedRouterOptions: + auto_router_config_path: str | None = None + auto_router_config: str | None = None + auto_router_default_model: str | None = None + auto_router_embedding_model: str | None = None + auto_router_max_input_chars: int | None = None + auto_router_routing_compression: str | None = None + auto_router_model_compression: str | None = None + complexity_router_config: Mapping[str, object] | None = None + complexity_router_default_model: str | None = None + adaptive_router_config: Mapping[str, object] | None = None + adaptive_router_default_model: str | None = None + quality_router_config: Mapping[str, object] | None = None + quality_router_default_model: str | None = None + + +@dataclass(frozen=True, slots=True, kw_only=True) +class CachingOptions: + caching: bool | None = None + cache: "DynamicCacheControl | None" = None + ttl: float | None = None + enable_prompt_caching: bool | None = None + caching_groups: Sequence[Sequence[str]] | None = None + + +@dataclass(frozen=True, slots=True, kw_only=True) +class CostOptions: + cost_per_query: float | None = None + base_model: str | None = None + max_budget: float | None = None + budget_duration: str | None = None + + +@dataclass(frozen=True, slots=True, kw_only=True) +class ObservabilityOptions: + id: str | None = None + metadata: MutableMapping[str, object] | None = None # mutable-ok: the router and logging write keys into it + litellm_metadata: MutableMapping[str, object] | None = None # mutable-ok: the proxy writes keys into it + tags: Sequence[str] | None = None + litellm_trace_id: str | None = None + litellm_session_id: str | None = None + litellm_request_debug: bool | None = None + logger_fn: Callable[[Mapping[str, object]], None] | None = None + verbose: bool | None = None + no_log: bool | None = field(default=None, metadata=wire("no-log")) + + +@dataclass(frozen=True, slots=True, kw_only=True) +class AgenticLoopOptions: + max_agentic_loops: int | None = None + + +@dataclass(frozen=True, slots=True, kw_only=True) +class GuardrailOptions: + guardrails: Sequence[str] | None = None + + +@dataclass(frozen=True, slots=True, kw_only=True) +class PromptOptions: + prompt_id: str | None = None + prompt_variables: Mapping[str, object] | None = None + prompt_version: str | None = None + prompt_environment: str | None = None + prompt_label: str | None = None + litellm_system_prompt: str | None = None + custom_prompt_dict: Mapping[str, object] | None = None + roles: Mapping[str, object] | None = None + final_prompt_value: str | None = None + bos_token: str | None = None + eos_token: str | None = None + hf_model_name: str | None = None + supports_system_message: bool | None = None + ensure_alternating_roles: bool | None = None + user_continue_message: "ChatCompletionUserMessage | None" = None + assistant_continue_message: "ChatCompletionAssistantMessage | None" = None + disable_add_transform_inline_image_block: bool | None = None + + +@dataclass(frozen=True, slots=True, kw_only=True) +class ResponseOptions: + merge_reasoning_content_in_choices: bool | None = None + enable_json_schema_validation: bool | None = None + complete_response: bool | None = None + stream_chunk_size: int | None = None + keepalive_seconds: float | None = None + allow_client_keepalive_override: bool | None = None + + +@dataclass(frozen=True, slots=True, kw_only=True) +class MockOptions: + mock_response: "MockResponse | None" = None + mock_timeout: bool | None = None + + +@dataclass(frozen=True, slots=True, kw_only=True) +class LiteLLMOptions: + dispatch: DispatchOptions + routing: RoutingOptions + deployment: DeploymentOptions + specialized_routers: SpecializedRouterOptions + caching: CachingOptions + cost: CostOptions + observability: ObservabilityOptions + agentic_loop: AgenticLoopOptions + guardrails: GuardrailOptions + prompt: PromptOptions + response: ResponseOptions + mock: MockOptions + + +@dataclass(frozen=True, slots=True, kw_only=True) +class CallState: + litellm_call_id: str | None = None + completion_call_id: str | None = None + model_alias_map: Mapping[str, str] | None = None + data_residency: str | None = None + litellm_logging_obj: "Logging | None" = None + preset_cache_key: str | None = None + cache_key: str | None = None + stream_response: "Mapping[str, ModelResponse] | None" = None + context_compaction_state: "CompactionState | None" = field(default=None, metadata=wire("_context_compaction_state")) + + +@dataclass(frozen=True, slots=True, kw_only=True) +class AgenticLoopState: + depth: int | None = field(default=None, metadata=wire("_agentic_loop_depth")) + fingerprints: Sequence[str] | None = field(default=None, metadata=wire("_agentic_loop_fingerprints")) + api_surface: Literal["chat_completions", "responses"] | None = field( + default=None, metadata=wire("_agentic_loop_api_surface") + ) + code_interpreter_active: bool | None = field(default=None, metadata=wire("_code_interpreter_interception_active")) + code_interpreter_sandbox_key: str | None = field( + default=None, metadata=wire("_code_interpreter_interception_sandbox_key") + ) + code_interpreter_session_scoped: bool | None = field( + default=None, metadata=wire("_code_interpreter_interception_session_scoped") + ) + code_interpreter_converted_stream: bool | None = field( + default=None, metadata=wire("_code_interpreter_interception_converted_stream") + ) + websearch_emit_native_blocks: bool | None = field( + default=None, metadata=wire("_websearch_interception_emit_native_blocks") + ) + websearch_converted_stream: bool | None = field( + default=None, metadata=wire("_websearch_interception_converted_stream") + ) + headroom_converted_stream: bool | None = field( + default=None, metadata=wire("_headroom_interception_converted_stream") + ) + + +@dataclass(frozen=True, slots=True, kw_only=True) +class RouterState: + weights: "RouterWeights | None" = field(default=None, metadata=wire("_router_weights")) + fallback_depth: int | None = None + max_fallbacks: int | None = None + attempted_targets: "AttemptedFallbackTargets | None" = None + + +@dataclass(frozen=True, slots=True, kw_only=True) +class ProxyRequestState: + proxy_server_request: Mapping[str, object] | None = None + secret_fields: "SecretFields | None" = None + trusted_callback_vars: Mapping[str, str] | None = field(default=None, metadata=wire(TRUSTED_CALLBACK_VARS_FIELD)) + addressed_response_id: str | None = field(default=None, metadata=wire(ADDRESSED_RESPONSE_ID_FIELD)) + strip_stream_usage: bool | None = field(default=None, metadata=wire("_litellm_strip_stream_usage")) + client_side_timeout: bool | None = None + model_file_id_mapping: Mapping[str, Mapping[str, str]] | None = None + + +@dataclass(frozen=True, slots=True, kw_only=True) +class EntrypointState: + acompletion: bool | None = None + aembedding: bool | None = None + aimg_generation: bool | None = None + atext_completion: bool | None = None + text_completion: bool | None = None + allm_passthrough_route: bool | None = None + async_call: bool | None = None + + +@dataclass(frozen=True, slots=True, kw_only=True) +class InternalState: + call: CallState + agentic_loop: AgenticLoopState + router: RouterState + proxy: ProxyRequestState + entrypoint: EntrypointState + + +KWARG_ARTIFACTS: Final[tuple[str, ...]] = ("self", "use_client", "model_config", "rust") + +LITELLM_OWNED_ROOTS: Final = (ConnectionSettings, LiteLLMOptions, InternalState) + + +def wire_names(owner: type) -> tuple[str, ...]: + return tuple(owned.metadata.get(WIRE_NAME, owned.name) for owned in fields(owner)) + + +def owned_wire_names(root: type) -> tuple[str, ...]: + def names() -> Iterator[str]: + for leaf in fields(root): + if not is_dataclass(leaf.type): + raise TypeError(f"{root.__name__}.{leaf.name} is not a dataclass leaf") + yield from wire_names(leaf.type) # pyright: ignore[reportArgumentType] # Field.type admits str + + return tuple(names()) + + +OWNED_KWARG_NAMES: Final = tuple(name for root in LITELLM_OWNED_ROOTS for name in owned_wire_names(root)) +AGENTIC_LOOP_KWARG_NAMES: Final = (*wire_names(AgenticLoopState), *wire_names(AgenticLoopOptions)) +BEDROCK_BATCH_KWARG_NAMES: Final = wire_names(BedrockBatchConnection) diff --git a/litellm/types/router.py b/litellm/types/router.py index c0f724584fd..b72809f625f 100644 --- a/litellm/types/router.py +++ b/litellm/types/router.py @@ -24,6 +24,7 @@ if TYPE_CHECKING: from .completion import CompletionRequest from .embedding import EmbeddingRequest +from .litellm_params import RoutingStrategyName from .llms.bedrock import AwsSessionTag from .llms.openai import OpenAIFileObject from .search import SearchProvider @@ -104,12 +105,7 @@ class RouterConfig(BaseModel): context_window_fallbacks: list | None = [] model_group_alias: dict[str, list[str]] | None = {} retry_after: int | None = 0 - routing_strategy: Literal[ - "simple-shuffle", - "least-busy", - "usage-based-routing", - "latency-based-routing", - ] = "simple-shuffle" + routing_strategy: RoutingStrategyName = "simple-shuffle" routing_groups: list[RoutingGroup] | None = None model_config = ConfigDict(protected_namespaces=()) diff --git a/litellm/types/utils.py b/litellm/types/utils.py index caf88e5d517..7aaf11faa5d 100644 --- a/litellm/types/utils.py +++ b/litellm/types/utils.py @@ -56,8 +56,15 @@ from litellm.types.llms.base import ( from litellm.types.mcp import MCPServerCostInfo from ..litellm_core_utils.core_helpers import map_finish_reason, process_response_headers +from . import litellm_params as _litellm_params from .agents import LiteLLMSendMessageResponse from .guardrails import GuardrailEventHooks +from .litellm_params import ( + AGENTIC_LOOP_KWARG_NAMES, + BEDROCK_BATCH_KWARG_NAMES, + KWARG_ARTIFACTS, + OWNED_KWARG_NAMES, +) from .llms.anthropic_messages.anthropic_response import AnthropicMessagesResponse from .llms.base import HiddenParams from .llms.openai import ( @@ -3901,205 +3908,20 @@ def pricing_override_fields(*sources: Mapping[str, object]) -> tuple[str, ...]: ) -# Server-controlled fields that bound or drive an interceptor's agentic loop -# (depth, cycle fingerprints, ceiling, code-interpreter sandbox state). Listed -# in all_litellm_params so they are treated as LiteLLM-level and excluded from -# get_non_default_completion_params; otherwise the OpenAI param builder sweeps -# any unrecognized top-level key into extra_body and leaks them to the provider. -# This is what lets the loop carry state across rerun calls without a provider -# scrubber. -agentic_loop_internal_litellm_params: Final = [ - "_agentic_loop_depth", - "_agentic_loop_fingerprints", - "_agentic_loop_api_surface", - "max_agentic_loops", - "_code_interpreter_interception_active", - "_code_interpreter_interception_sandbox_key", - "_code_interpreter_interception_session_scoped", - "_code_interpreter_interception_converted_stream", - "_websearch_interception_emit_native_blocks", - "_websearch_interception_converted_stream", - "_headroom_interception_converted_stream", +agentic_loop_internal_litellm_params: Final = list(AGENTIC_LOOP_KWARG_NAMES) # mutable-ok: public type stays a list + +bedrock_batch_litellm_params: Final = BEDROCK_BATCH_KWARG_NAMES + +TRUSTED_CALLBACK_VARS_FIELD: Final = _litellm_params.TRUSTED_CALLBACK_VARS_FIELD +ADDRESSED_RESPONSE_ID_FIELD: Final = _litellm_params.ADDRESSED_RESPONSE_ID_FIELD + +all_litellm_params = [ # rebind-ok: two star imports in litellm/__init__.py re-bind it # mutable-ok: callers concat + *OWNED_KWARG_NAMES, + *KWARG_ARTIFACTS, + *StandardCallbackDynamicParams.__annotations__, + *CustomPricingLiteLLMParams.model_fields, ] -# Proxy-owned callback credentials, stamped from admin-configured team/key callback -# settings. Listed in all_litellm_params for the same reason as the agentic-loop -# fields above: an unrecognized top-level key is swept into extra_body and sent to -# the provider. -TRUSTED_CALLBACK_VARS_FIELD: Final = "litellm_trusted_callback_vars" - -ADDRESSED_RESPONSE_ID_FIELD: Final = "_litellm_addressed_response_id" - -# Bedrock managed-batch deployment config, read from litellm_params by the batch and -# files transformations. Listed for the same reason as the fields above: these sit on -# a deployment that also serves chat, so leaking them into extra_body makes Bedrock -# reject every non-batch request to that deployment. -bedrock_batch_litellm_params: Final = ( - "aws_batch_role_arn", - "s3_bucket_name", - "s3_region_name", - "s3_endpoint_url", - "s3_output_bucket_name", - "s3_bucket_owner", - "s3_access_key_id", - "s3_secret_access_key", - "s3_encryption_key_id", - "bedrock_tags", -) - -all_litellm_params = ( - agentic_loop_internal_litellm_params - + [TRUSTED_CALLBACK_VARS_FIELD, ADDRESSED_RESPONSE_ID_FIELD, *bedrock_batch_litellm_params] - + [ - "_context_compaction_state", - "metadata", - "litellm_metadata", - "keepalive_seconds", - "allow_client_keepalive_override", - "litellm_trace_id", - "litellm_request_debug", - "guardrails", - "tags", - "acompletion", - "aimg_generation", - "atext_completion", - "text_completion", - "caching", - "mock_response", - "mock_timeout", - "disable_add_transform_inline_image_block", - "api_key", - "api_version", - "prompt_id", - "prompt_variables", - "litellm_system_prompt", - "provider_specific_header", - "prompt_version", - "prompt_environment", - "api_base", - "force_timeout", - "logger_fn", - "verbose", - "custom_llm_provider", - "model_file_id_mapping", - "litellm_logging_obj", - "litellm_call_id", - "completion_call_id", - "model_alias_map", - "custom_prompt_dict", - "stream_response", - "cost_per_query", - "ssl_verify", - "data_residency", - "async_call", - "aembedding", - "allm_passthrough_route", - "_litellm_strip_stream_usage", - "use_client", - "id", - "fallbacks", - "routing_strategy", - "_router_weights", - "azure", - "headers", - "model_list", - "num_retries", - "context_window_fallback_dict", - "retry_policy", - "retry_strategy", - "roles", - "final_prompt_value", - "bos_token", - "eos_token", - "request_timeout", - "client_side_timeout", - "complete_response", - "self", - "client", - "rpm", - "tpm", - "default_api_key_rpm_limit", - "default_api_key_tpm_limit", - "itpm", - "otpm", - "max_parallel_requests", - "input_cost_per_token", - "output_cost_per_token", - "input_cost_per_second", - "output_cost_per_second", - "hf_model_name", - "model_info", - "proxy_server_request", - "secret_fields", - "preset_cache_key", - "caching_groups", - "ttl", - "cache", - "enable_prompt_caching", - "no-log", - "base_model", - "stream_timeout", - "stream_chunk_size", - "supports_system_message", - "region_name", - "allowed_model_region", - "model_config", - "fastest_response", - "cooldown_time", - "cache_key", - "max_retries", - "azure_ad_token_provider", - "tenant_id", - "client_id", - "azure_username", - "azure_password", - "azure_scope", - "client_secret", - "user_continue_message", - "configurable_clientside_auth_params", - "weight", - "ensure_alternating_roles", - "assistant_continue_message", - "user_continue_message", - "fallback_depth", - "max_fallbacks", - "attempted_targets", - "max_budget", - "budget_duration", - "use_in_pass_through", - "merge_reasoning_content_in_choices", - "litellm_credential_name", - "allowed_openai_params", - "litellm_session_id", - "provider_affinity_header", - "use_litellm_proxy", - "use_chat_completions_api", - "rust", - "prompt_label", - "shared_session", - "search_tool_name", - "order", - "enable_tag_filtering", - "enable_json_schema_validation", - "use_xai_oauth", - "auto_router_config_path", - "auto_router_config", - "auto_router_default_model", - "auto_router_embedding_model", - "auto_router_max_input_chars", - "auto_router_routing_compression", - "auto_router_model_compression", - "complexity_router_config", - "complexity_router_default_model", - "adaptive_router_config", - "adaptive_router_default_model", - "quality_router_config", - "quality_router_default_model", - ] - + list(StandardCallbackDynamicParams.__annotations__.keys()) - + list(CustomPricingLiteLLMParams.model_fields.keys()) -) - class KeyGenerationConfig(TypedDict, total=False): required_params: list[str] # specify params that must be present in the key generation request diff --git a/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py b/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py index 2d929a832a5..a40741c8fdb 100644 --- a/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py +++ b/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py @@ -5,10 +5,11 @@ import logging import os import sys import zlib -from collections.abc import Callable +from collections.abc import Callable, Mapping from contextlib import ExitStack, contextmanager +from dataclasses import dataclass from io import BytesIO -from types import SimpleNamespace +from types import MappingProxyType, SimpleNamespace from typing import Final from unittest.mock import AsyncMock, MagicMock, patch @@ -16,7 +17,7 @@ import httpx import pytest from fastapi import HTTPException, Request, Response, UploadFile from fastapi.responses import StreamingResponse -from pydantic import ValidationError +from pydantic import TypeAdapter, ValidationError from starlette.datastructures import FormData, Headers, QueryParams from starlette.datastructures import UploadFile as StarletteUploadFile @@ -45,6 +46,7 @@ from litellm.proxy.pass_through_endpoints.success_handler import ( PassThroughEndpointLogging, ) from litellm.proxy.route_llm_request import ProxyModelNotFoundError +from litellm.types import utils as types_utils from litellm.types.passthrough_endpoints.pass_through_endpoints import ( LITELLM_PASS_THROUGH_DEPLOYMENT_MODEL_INFO_STATE_KEY, LITELLM_PASS_THROUGH_RAW_BODY_STATE_KEY, @@ -7305,6 +7307,156 @@ def test_passthrough_logs_the_resolved_deployment_model_info_over_the_request_bo assert kwargs["litellm_params"]["metadata"]["model_info"] == {"id": "vertex-gemini-38-flash-dep"} +@dataclass(frozen=True, slots=True, kw_only=True) +class _PassThroughSplit: + litellm_params: Mapping[str, object] + forwarded_body: Mapping[str, object] + + +_LITELLM_PARAMS: Final = TypeAdapter(dict[str, object]) +_PROXY_SERVER_REQUEST: Final = TypeAdapter(dict[str, object]) + + +def _split_pass_through_body(body: str) -> _PassThroughSplit: + mock_request: Final = MagicMock(spec=Request) + mock_request.method = "POST" + mock_request.url = "http://0.0.0.0:4000/gemini/v1beta/models/gemini-2.5-flash:generateContent" + mock_request.headers = Headers() + mock_request.scope = MappingProxyType({}) + + init_kwargs_for_pass_through_endpoint: Final = HttpPassThroughEndpointHelpers._init_kwargs_for_pass_through_endpoint # pyright: ignore[reportUnknownVariableType, reportUnknownMemberType] # untyped legacy helper + kwargs: Final = init_kwargs_for_pass_through_endpoint( + request=mock_request, + user_api_key_dict=UserAPIKeyAuth(api_key="hashed-key"), + passthrough_logging_payload=MagicMock(), + logging_obj=MagicMock(), + _parsed_body=json.loads(body), + litellm_call_id="lit-owned-keys-call-id", + ) + validate_litellm_params: Final = _LITELLM_PARAMS.validate_python # pyright: ignore[reportUnknownArgumentType] # untyped legacy helper + litellm_params: Final = validate_litellm_params(kwargs["litellm_params"]) + return _PassThroughSplit( + litellm_params=MappingProxyType(litellm_params), + forwarded_body=MappingProxyType( + _LITELLM_PARAMS.validate_python( + _PROXY_SERVER_REQUEST.validate_python(litellm_params["proxy_server_request"])["body"] + ) + ), + ) + + +GEMINI_BODY: Final = '{"contents": [{"parts": [{"text": "hi"}]}], "generationConfig": {"temperature": 0}}' + + +def _metadata_of(split: _PassThroughSplit) -> Mapping[str, object]: + return MappingProxyType(_LITELLM_PARAMS.validate_python(split.litellm_params["metadata"])) + + +def test_passthrough_moves_every_litellm_owned_key_from_the_forwarded_body_into_litellm_params() -> None: + split: Final = _split_pass_through_body( + '{"ttl": 30, "contents": [{"parts": [{"text": "hi"}]}], "num_retries": 2,' + ' "generationConfig": {"temperature": 0}, "litellm_trace_id": "trace-a"}' + ) + + assert frozenset(split.litellm_params) == frozenset( + ("ttl", "num_retries", "litellm_trace_id", "metadata", "proxy_server_request") + ) + assert tuple(split.litellm_params[k] for k in ("ttl", "num_retries", "litellm_trace_id")) == (30, 2, "trace-a") + assert split.forwarded_body == json.loads(GEMINI_BODY) + + +PROXY_STAMPED_NAMES: Final = frozenset( + ( + "proxy_server_request", + "secret_fields", + "litellm_trusted_callback_vars", + "_litellm_addressed_response_id", + "_litellm_strip_stream_usage", + "client_side_timeout", + "model_file_id_mapping", + ) +) + + +@pytest.mark.parametrize( + "name", + sorted(frozenset(litellm.all_litellm_params) - frozenset(("metadata", "litellm_metadata")) - PROXY_STAMPED_NAMES), +) +def test_passthrough_keeps_each_registered_litellm_owned_name_out_of_the_forwarded_body(name: str) -> None: + split: Final = _split_pass_through_body(json.dumps({name: "owned", **json.loads(GEMINI_BODY)})) + + assert frozenset(split.litellm_params) == frozenset((name, "metadata", "proxy_server_request")) + assert split.litellm_params[name] == "owned" + assert split.forwarded_body == json.loads(GEMINI_BODY) + + +@pytest.mark.parametrize("name", sorted(PROXY_STAMPED_NAMES)) +def test_passthrough_drops_a_client_supplied_proxy_stamped_name(name: str) -> None: + split: Final = _split_pass_through_body(json.dumps({name: {"forged": "by-client"}, **json.loads(GEMINI_BODY)})) + + assert frozenset(split.litellm_params) == frozenset(("metadata", "proxy_server_request")) + assert split.litellm_params["proxy_server_request"] != {"forged": "by-client"}, split.litellm_params + assert split.forwarded_body == json.loads(GEMINI_BODY) + + +def test_passthrough_merges_both_metadata_carriers_from_the_body_into_one_metadata_key() -> None: + split: Final = _split_pass_through_body( + '{"metadata": {"client_tag": "a"}, "contents": [{"parts": [{"text": "hi"}]}], "ttl": 30,' + ' "litellm_metadata": {"lm": "b"}, "generationConfig": {"temperature": 0}}' + ) + + assert frozenset(split.litellm_params) == frozenset(("ttl", "metadata", "proxy_server_request")) + assert _metadata_of(split) == {**_metadata_of(_split_pass_through_body(GEMINI_BODY)), "client_tag": "a", "lm": "b"} + assert split.forwarded_body == json.loads(GEMINI_BODY) + + +def test_passthrough_lets_metadata_win_over_litellm_metadata_on_a_shared_key() -> None: + split: Final = _split_pass_through_body( + '{"litellm_metadata": {"shared": "from-litellm-metadata", "lm": "b"},' + ' "metadata": {"shared": "from-metadata", "client_tag": "a"}, "contents": []}' + ) + + assert _metadata_of(split) == { + **_metadata_of(_split_pass_through_body('{"contents": []}')), + "shared": "from-metadata", + "lm": "b", + "client_tag": "a", + } + + +def test_passthrough_orders_extracted_litellm_params_by_the_registry() -> None: + body: Final = json.dumps({"ttl": 30, "tags": ["team-a"], "num_retries": 2, "contents": []}) + split: Final = _split_pass_through_body(body) + body_keys: Final = frozenset(json.loads(body)) + + assert tuple(k for k in split.litellm_params if k in body_keys) == tuple( + k for k in types_utils.all_litellm_params if k in body_keys + ) + + +LATE_REGISTERED_BODY: Final = '{"registered_later": 1, "contents": [{"parts": [{"text": "hi"}]}]}' + + +def test_passthrough_sees_a_name_appended_to_the_public_list_after_import() -> None: + litellm.all_litellm_params.append("registered_later") + try: + split: Final = _split_pass_through_body(LATE_REGISTERED_BODY) + finally: + litellm.all_litellm_params.remove("registered_later") + + assert frozenset(split.litellm_params) == frozenset(("registered_later", "metadata", "proxy_server_request")) + assert split.forwarded_body == {"contents": [{"parts": [{"text": "hi"}]}]} + + +def test_passthrough_sees_the_public_list_rebound_after_import(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr(types_utils, "all_litellm_params", (*litellm.all_litellm_params, "registered_later")) + + split: Final = _split_pass_through_body(LATE_REGISTERED_BODY) + + assert frozenset(split.litellm_params) == frozenset(("registered_later", "metadata", "proxy_server_request")) + assert split.forwarded_body == {"contents": [{"parts": [{"text": "hi"}]}]} + + @pytest.mark.asyncio async def test_chat_completion_pass_through_endpoint_answers_an_openai_typed_error_for_an_unknown_model( monkeypatch: pytest.MonkeyPatch, diff --git a/tests/test_litellm/test_utils.py b/tests/test_litellm/test_utils.py index 285188c9c09..768d8955b8e 100644 --- a/tests/test_litellm/test_utils.py +++ b/tests/test_litellm/test_utils.py @@ -3730,6 +3730,23 @@ def test_scoped_weights_are_excluded_from_provider_params(filter_name: str) -> N assert filtered == {"provider_option": "kept"} +@pytest.mark.parametrize( + "provider_filter", + [ + litellm.utils.get_non_default_completion_params, + litellm.utils.get_non_default_transcription_params, + litellm.utils.filter_out_litellm_params, + ], +) +@pytest.mark.parametrize("setting", [("tag_regex", ["^team-a$"]), ("max_file_size_mb", 5)]) +def test_deployment_only_settings_copied_by_the_router_stay_out_of_provider_params( + provider_filter: Callable[[dict[str, object]], Mapping[str, object]], setting: tuple[str, object] +) -> None: + name, value = setting + filtered: Final = provider_filter({"provider_option": "kept", name: value}) + assert filtered == {"provider_option": "kept"}, filtered + + class TestGetOptionalParamsTencent: """Tests that tencent provider uses TencentChatConfig for parameter mapping.""" diff --git a/tests/unit/types/test_litellm_params.py b/tests/unit/types/test_litellm_params.py new file mode 100644 index 00000000000..e421321aaaa --- /dev/null +++ b/tests/unit/types/test_litellm_params.py @@ -0,0 +1,655 @@ +import inspect +from collections.abc import Callable, Mapping, Sequence +from dataclasses import dataclass, field, fields +from operator import attrgetter +from types import MappingProxyType +from typing import Final, TypeAlias, cast, get_type_hints + +import httpx +import pytest +from aiohttp import ClientSession +from openai import AsyncAzureOpenAI, AsyncOpenAI, AzureOpenAI, OpenAI +from pydantic import BaseModel, ConfigDict, TypeAdapter, ValidationError + +import litellm +from litellm.caching.caching import Cache +from litellm.litellm_core_utils.get_litellm_params import ( + get_litellm_params, # pyright: ignore[reportUnknownVariableType] # untyped legacy carrier +) +from litellm.litellm_core_utils.litellm_logging import Logging +from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler +from litellm.router_strategy.complexity_router.context_compaction import CompactionState +from litellm.router_utils.fallback_event_handlers import AttemptedFallbackTargets +from litellm.types import litellm_params +from litellm.types import utils as types_utils +from litellm.types.caching import DynamicCacheControl +from litellm.types.litellm_params import ( + ADDRESSED_RESPONSE_ID_FIELD, + LITELLM_OWNED_ROOTS, + TRUSTED_CALLBACK_VARS_FIELD, + CachingOptions, + owned_wire_names, + wire, + wire_names, +) +from litellm.types.llms.openai import ChatCompletionAssistantMessage, ChatCompletionUserMessage +from litellm.types.proxy.litellm_pre_call_utils import SecretFields +from litellm.types.router import ( + ConfigurableClientsideParamsCustomAuth, + CredentialLiteLLMParams, + DeploymentTypedDict, + RetryPolicy, + RouterConfig, + UpdateRouterConfig, +) +from litellm.types.router_weights import RouterWeights +from litellm.types.utils import ( + CustomPricingLiteLLMParams, + ModelResponse, + ModelResponseStream, + ProviderSpecificHeader, + StandardCallbackDynamicParams, + agentic_loop_internal_litellm_params, + all_litellm_params, + bedrock_batch_litellm_params, +) +from litellm.utils import ( + filter_out_litellm_params, # pyright: ignore[reportUnknownVariableType] # untyped legacy classifier + get_non_default_completion_params, # pyright: ignore[reportUnknownVariableType] # untyped legacy classifier + get_non_default_transcription_params, # pyright: ignore[reportUnknownVariableType] # untyped legacy classifier +) + +PROVIDER_KNOB: Final = "registry_test_provider_only_knob" + +CONNECTION_NAMES: Final = ( + "api_key", + "api_base", + "api_version", + "region_name", + "headers", + "provider_specific_header", + "client", + "shared_session", + "ssl_verify", + "request_timeout", + "force_timeout", + "stream_timeout", + "max_retries", + "tenant_id", + "client_id", + "client_secret", + "azure_username", + "azure_password", + "azure_scope", + "azure_ad_token_provider", + "litellm_credential_name", + "configurable_clientside_auth_params", + "use_xai_oauth", + "aws_batch_role_arn", + "s3_bucket_name", + "s3_region_name", + "s3_endpoint_url", + "s3_output_bucket_name", + "s3_bucket_owner", + "s3_access_key_id", + "s3_secret_access_key", + "s3_encryption_key_id", + "bedrock_tags", +) + +OPTION_NAMES: Final = ( + "custom_llm_provider", + "azure", + "use_litellm_proxy", + "use_chat_completions_api", + "use_in_pass_through", + "allowed_openai_params", + "fallbacks", + "context_window_fallback_dict", + "num_retries", + "retry_policy", + "retry_strategy", + "routing_strategy", + "cooldown_time", + "allowed_model_region", + "enable_tag_filtering", + "fastest_response", + "provider_affinity_header", + "search_tool_name", + "model_list", + "model_info", + "rpm", + "tpm", + "itpm", + "otpm", + "default_api_key_rpm_limit", + "default_api_key_tpm_limit", + "max_parallel_requests", + "weight", + "order", + "tag_regex", + "max_file_size_mb", + "auto_router_config_path", + "auto_router_config", + "auto_router_default_model", + "auto_router_embedding_model", + "auto_router_max_input_chars", + "auto_router_routing_compression", + "auto_router_model_compression", + "complexity_router_config", + "complexity_router_default_model", + "adaptive_router_config", + "adaptive_router_default_model", + "quality_router_config", + "quality_router_default_model", + "caching", + "cache", + "ttl", + "enable_prompt_caching", + "caching_groups", + "cost_per_query", + "base_model", + "max_budget", + "budget_duration", + "id", + "metadata", + "litellm_metadata", + "tags", + "litellm_trace_id", + "litellm_session_id", + "litellm_request_debug", + "logger_fn", + "verbose", + "no-log", + "max_agentic_loops", + "guardrails", + "prompt_id", + "prompt_variables", + "prompt_version", + "prompt_environment", + "prompt_label", + "litellm_system_prompt", + "custom_prompt_dict", + "roles", + "final_prompt_value", + "bos_token", + "eos_token", + "hf_model_name", + "supports_system_message", + "ensure_alternating_roles", + "user_continue_message", + "assistant_continue_message", + "disable_add_transform_inline_image_block", + "merge_reasoning_content_in_choices", + "enable_json_schema_validation", + "complete_response", + "stream_chunk_size", + "keepalive_seconds", + "allow_client_keepalive_override", + "mock_response", + "mock_timeout", +) + +AGENTIC_LOOP_STATE_NAMES: Final = ( + "_agentic_loop_depth", + "_agentic_loop_fingerprints", + "_agentic_loop_api_surface", + "_code_interpreter_interception_active", + "_code_interpreter_interception_sandbox_key", + "_code_interpreter_interception_session_scoped", + "_code_interpreter_interception_converted_stream", + "_websearch_interception_emit_native_blocks", + "_websearch_interception_converted_stream", + "_headroom_interception_converted_stream", +) + +INTERNAL_STATE_NAMES: Final = ( + "litellm_call_id", + "completion_call_id", + "model_alias_map", + "data_residency", + "litellm_logging_obj", + "preset_cache_key", + "cache_key", + "stream_response", + "_context_compaction_state", + *AGENTIC_LOOP_STATE_NAMES, + "_router_weights", + "fallback_depth", + "max_fallbacks", + "attempted_targets", + "proxy_server_request", + "secret_fields", + "litellm_trusted_callback_vars", + "_litellm_addressed_response_id", + "_litellm_strip_stream_usage", + "client_side_timeout", + "model_file_id_mapping", + "acompletion", + "aembedding", + "aimg_generation", + "atext_completion", + "text_completion", + "allm_passthrough_route", + "async_call", +) + +BEDROCK_BATCH_NAMES: Final = ( + "aws_batch_role_arn", + "s3_bucket_name", + "s3_region_name", + "s3_endpoint_url", + "s3_output_bucket_name", + "s3_bucket_owner", + "s3_access_key_id", + "s3_secret_access_key", + "s3_encryption_key_id", + "bedrock_tags", +) + +ARTIFACT_NAMES: Final = ("self", "use_client", "model_config", "rust") + +CALLBACK_VAR_NAMES: Final = tuple(StandardCallbackDynamicParams.__annotations__) + +PRICING_NAMES: Final = tuple(CustomPricingLiteLLMParams.model_fields) + +OWNED_NAMES: Final = ( + *CONNECTION_NAMES, + *OPTION_NAMES, + *INTERNAL_STATE_NAMES, + *ARTIFACT_NAMES, + *CALLBACK_VAR_NAMES, + *PRICING_NAMES, +) + +Classifier: TypeAlias = Callable[[dict[str, object]], dict[str, object]] # mutable-ok: classifiers use dict + +CLASSIFIERS: Final[Mapping[str, Classifier]] = MappingProxyType( + { # pyright: ignore[reportUnknownArgumentType] # untyped legacy classifiers + "completion": get_non_default_completion_params, + "transcription": get_non_default_transcription_params, + "filter_out": filter_out_litellm_params, + } +) + + +@pytest.mark.parametrize("classifier_name", CLASSIFIERS) +@pytest.mark.parametrize("name", OWNED_NAMES) +def test_owned_name_is_kept_out_of_provider_params(name: str, classifier_name: str) -> None: + provider_value: Final = object() + classify: Final = CLASSIFIERS[classifier_name] + + result: Final = classify({name: object(), PROVIDER_KNOB: provider_value}) # mutable-ok: classifiers take a dict + + assert result == MappingProxyType({PROVIDER_KNOB: provider_value}) + assert result[PROVIDER_KNOB] is provider_value + + +def test_a_name_no_object_declares_reaches_the_provider() -> None: + result: Final = CLASSIFIERS["completion"]({PROVIDER_KNOB: 1}) # mutable-ok: classifier input type + + assert result == MappingProxyType({PROVIDER_KNOB: 1}) + + +def _cache_key_for_model_group(cache: Cache, model_group: str, options: CachingOptions) -> str: + return cache.get_cache_key( # pyright: ignore[reportUnknownMemberType] # untyped legacy key builder + model=model_group, + messages=(MappingProxyType({"role": "user", "content": "shared prompt"}),), + metadata=MappingProxyType({"caching_groups": options.caching_groups, "model_group": model_group}), + ) + + +def test_caching_groups_is_a_flat_sequence_of_model_groups_that_share_one_cache_key( + monkeypatch: pytest.MonkeyPatch, +) -> None: + for callback_list in ("input_callback", "success_callback", "_async_success_callback"): + monkeypatch.setattr(litellm, callback_list, []) # mutable-ok: Cache() appends "cache" to these lists + options: Final = CachingOptions(caching_groups=(("gpt-4", "gpt-4o"), ("claude-3",))) + cache: Final = Cache() + + keys: Final = tuple(_cache_key_for_model_group(cache, group, options) for group in ("gpt-4", "gpt-4o", "claude-3")) + + assert (keys[0] == keys[1], keys[0] == keys[2]) == (True, False) + + +def test_all_litellm_params_is_exactly_the_owned_inventory() -> None: + assert frozenset(all_litellm_params) == frozenset(OWNED_NAMES) + assert frozenset(ARTIFACT_NAMES).isdisjoint(DECLARED_NAMES) + + +def test_every_owned_name_has_exactly_one_owner() -> None: + duplicated: Final = tuple(name for name in dict.fromkeys(all_litellm_params) if all_litellm_params.count(name) > 1) + + assert duplicated == () + + +@pytest.mark.parametrize( + ("exported", "declared"), + ( + pytest.param( + types_utils.TRUSTED_CALLBACK_VARS_FIELD, + litellm_params.TRUSTED_CALLBACK_VARS_FIELD, + id="TRUSTED_CALLBACK_VARS_FIELD", + ), + pytest.param( + types_utils.ADDRESSED_RESPONSE_ID_FIELD, + litellm_params.ADDRESSED_RESPONSE_ID_FIELD, + id="ADDRESSED_RESPONSE_ID_FIELD", + ), + ), +) +def test_types_utils_still_exports_the_field_constant(exported: str, declared: str) -> None: + assert exported == declared + + +@dataclass(frozen=True, slots=True, kw_only=True) +class _Leaf: + plain: int | None = None + renamed: int | None = field(default=None, metadata=wire("wire-name")) + + +@dataclass(frozen=True, slots=True, kw_only=True) +class _OtherLeaf: + plain: int | None = None + trailing: int | None = None + + +@dataclass(frozen=True, slots=True, kw_only=True) +class _Root: + first: _Leaf + second: _OtherLeaf + + +@dataclass(frozen=True, slots=True, kw_only=True) +class _RootDeclaringAKwargDirectly: + first: _Leaf + stray: int | None = None + + +def test_wire_names_are_the_field_names_in_declaration_order_unless_wire_renames_them() -> None: + assert wire_names(_Leaf) == ("plain", "wire-name") + + +def test_owned_wire_names_walk_leaves_in_declaration_order_and_keep_every_occurrence() -> None: + assert owned_wire_names(_Root) == ("plain", "wire-name", "plain", "trailing") + + +def test_owned_wire_names_refuse_a_root_that_declares_a_kwarg_outside_a_leaf() -> None: + with pytest.raises(TypeError): + owned_wire_names(_RootDeclaringAKwargDirectly) + + +def test_agentic_loop_names_concatenate_as_a_list() -> None: + extended: Final = agentic_loop_internal_litellm_params + ["caller_added"] # mutable-ok: list contract under test + + assert (type(extended), len(extended), frozenset(extended)) == ( + list, + len(AGENTIC_LOOP_STATE_NAMES) + 2, + frozenset((*AGENTIC_LOOP_STATE_NAMES, "max_agentic_loops", "caller_added")), + ) + + +def test_bedrock_batch_names_concatenate_as_a_tuple() -> None: + extended: Final = bedrock_batch_litellm_params + ("caller_added",) + + assert extended == (*BEDROCK_BATCH_NAMES, "caller_added") + + +def test_proxy_stamped_fields_keep_their_wire_names() -> None: + assert (TRUSTED_CALLBACK_VARS_FIELD, ADDRESSED_RESPONSE_ID_FIELD) == ( + "litellm_trusted_callback_vars", + "_litellm_addressed_response_id", + ) + + +def test_all_litellm_params_concatenates_with_a_list_like_the_completion_entrypoint_does() -> None: + extended: Final = ["aembedding", "extra_headers"] + all_litellm_params # mutable-ok: list contract under test + + assert (type(extended), frozenset(extended)) == (list, frozenset(("aembedding", "extra_headers", *OWNED_NAMES))) + + +CARRIED_AND_FORWARDED: Final = frozenset(("drop_params", "hugging_face", "no_log", "replicate", "together_ai")) + +CARRIER_SIGNATURE: Final = inspect.signature(get_litellm_params) # pyright: ignore[reportUnknownArgumentType] # legacy + +CARRIED_PARAMS: Final = tuple( + name for name in CARRIER_SIGNATURE.parameters if name != "kwargs" and name not in CARRIED_AND_FORWARDED +) + + +@pytest.mark.parametrize("name", CARRIED_PARAMS) +def test_every_param_get_litellm_params_carries_is_kept_out_of_provider_params(name: str) -> None: + provider_value: Final = object() + + result: Final = CLASSIFIERS["completion"]( + {name: object(), PROVIDER_KNOB: provider_value} # mutable-ok: classifier input type + ) + + assert result == MappingProxyType({PROVIDER_KNOB: provider_value}) + + +TYPED_CONFIG_MODELS: Final[Mapping[str, tuple[type[BaseModel], ...]]] = MappingProxyType( + { + "credentials": (CredentialLiteLLMParams,), + "router": (RouterConfig, UpdateRouterConfig), + } +) + +DECLARED_NAMES: Final = frozenset(name for root in LITELLM_OWNED_ROOTS for name in owned_wire_names(root)) + +ProviderClient: TypeAlias = ( + OpenAI + | AsyncOpenAI + | AzureOpenAI + | AsyncAzureOpenAI + | HTTPHandler + | AsyncHTTPHandler + | httpx.Client + | httpx.AsyncClient +) +MockResponse: TypeAlias = str | Exception | Mapping[str, object] | Sequence[float] | ModelResponse | ModelResponseStream + +TYPE_HINT_NAMESPACE: Final[Mapping[str, object]] = { + "ProviderClient": ProviderClient, + "ProviderSpecificHeader": ProviderSpecificHeader, + "ClientSession": ClientSession, + "AsyncAzureOpenAI": AsyncAzureOpenAI, + "AsyncOpenAI": AsyncOpenAI, + "AzureOpenAI": AzureOpenAI, + "OpenAI": OpenAI, + "AsyncHTTPHandler": AsyncHTTPHandler, + "HTTPHandler": HTTPHandler, + "ConfigurableClientsideParamsCustomAuth": ConfigurableClientsideParamsCustomAuth, + "RetryPolicy": RetryPolicy, + "DeploymentTypedDict": DeploymentTypedDict, + "DynamicCacheControl": DynamicCacheControl, + "ChatCompletionUserMessage": ChatCompletionUserMessage, + "ChatCompletionAssistantMessage": ChatCompletionAssistantMessage, + "MockResponse": MockResponse, + "ModelResponse": ModelResponse, + "ModelResponseStream": ModelResponseStream, + "Logging": Logging, + "SecretFields": SecretFields, + "CompactionState": CompactionState, + "RouterWeights": RouterWeights, + "AttemptedFallbackTargets": AttemptedFallbackTargets, +} + +LEAF_SAMPLES: Final[Mapping[type, Mapping[str, object]]] = { + litellm_params.ProviderConnection: {"api_key": "k", "request_timeout": 1.5}, + litellm_params.BedrockBatchConnection: {"aws_batch_role_arn": "arn", "bedrock_tags": ({"k": "v"},)}, + litellm_params.DispatchOptions: {"custom_llm_provider": "openai"}, + litellm_params.RoutingOptions: { + "fallbacks": [{"model": "gpt-4o", "api_key": "k", "temperature": 0}], + "num_retries": 2, + "retry_strategy": "constant_retry", + "routing_strategy": "simple-shuffle", + }, + litellm_params.DeploymentOptions: {"model_info": {"region": "us"}, "rpm": 2}, + litellm_params.SpecializedRouterOptions: {"adaptive_router_default_model": "gpt-4o"}, + litellm_params.CachingOptions: {"ttl": 30.0, "caching_groups": (("gpt-4o", "gpt-4o-mini"),)}, + litellm_params.CostOptions: {"max_budget": 10.0}, + litellm_params.ObservabilityOptions: {"metadata": {"request": "test"}, "no_log": True}, + litellm_params.AgenticLoopOptions: {"max_agentic_loops": 2}, + litellm_params.GuardrailOptions: {"guardrails": ("default",)}, + litellm_params.PromptOptions: {"prompt_id": "prompt", "prompt_variables": {"name": "value"}}, + litellm_params.ResponseOptions: {"stream_chunk_size": 64}, + litellm_params.MockOptions: {"mock_timeout": True}, + litellm_params.CallState: { + "completion_call_id": "call", + "model_alias_map": {"alias": "gpt-4o"}, + "data_residency": "us", + }, + litellm_params.AgenticLoopState: {"api_surface": "chat_completions", "depth": 1}, + litellm_params.RouterState: {"fallback_depth": 1}, + litellm_params.ProxyRequestState: { + "proxy_server_request": {"path": "/chat/completions"}, + "trusted_callback_vars": {"dd_api_key": "k"}, + }, + litellm_params.EntrypointState: {"acompletion": True}, +} + +LEAF_BAD_SAMPLES: Final[Mapping[type, Mapping[str, object]]] = { + litellm_params.ProviderConnection: {"api_key": 1}, + litellm_params.BedrockBatchConnection: {"aws_batch_role_arn": 1}, + litellm_params.DispatchOptions: {"custom_llm_provider": 1}, + litellm_params.RoutingOptions: {"num_retries": "2"}, + litellm_params.DeploymentOptions: {"rpm": "2"}, + litellm_params.SpecializedRouterOptions: {"auto_router_max_input_chars": "2"}, + litellm_params.CachingOptions: {"ttl": "30"}, + litellm_params.CostOptions: {"max_budget": "10"}, + litellm_params.ObservabilityOptions: {"verbose": "true"}, + litellm_params.AgenticLoopOptions: {"max_agentic_loops": "2"}, + litellm_params.GuardrailOptions: {"guardrails": (1,)}, + litellm_params.PromptOptions: {"prompt_id": 1}, + litellm_params.ResponseOptions: {"stream_chunk_size": "64"}, + litellm_params.MockOptions: {"mock_timeout": "true"}, + litellm_params.CallState: {"completion_call_id": 1}, + litellm_params.AgenticLoopState: {"depth": "1"}, + litellm_params.RouterState: {"fallback_depth": "1"}, + litellm_params.ProxyRequestState: {"proxy_server_request": "request"}, + litellm_params.EntrypointState: {"acompletion": "true"}, +} + +INVALID_LITERAL_SAMPLES: Final[tuple[tuple[type, Mapping[str, object]], ...]] = ( + (litellm_params.RoutingOptions, {"retry_strategy": "linear"}), + (litellm_params.RoutingOptions, {"routing_strategy": "random"}), + (litellm_params.AgenticLoopState, {"api_surface": "batches"}), +) + + +def _leaf_id(value: object) -> str: + return value.__name__ if isinstance(value, type) else "" + + +def _leaf_instance(leaf: type, sample: Mapping[str, object]) -> object: + constructor: Final = cast(Callable[..., object], leaf) + return constructor(**sample) + + +def _strict_leaf_validation(leaf: type, instance: object) -> object: + hints: Final[Mapping[str, object]] = cast( + Mapping[str, object], get_type_hints(type(instance), localns=TYPE_HINT_NAMESPACE) + ) + for field_info in fields(leaf): + value = cast(Callable[[object], object], attrgetter(field_info.name))(instance) + field_adapter: TypeAdapter[object] = TypeAdapter[object]( + hints[field_info.name], + config=ConfigDict(arbitrary_types_allowed=True), + ) + field_adapter.validate_python(value, strict=True) + return instance + + +@pytest.mark.parametrize("leaf,sample", LEAF_SAMPLES.items(), ids=_leaf_id) +def test_every_owned_leaf_accepts_a_strict_reader_shaped_sample(leaf: type, sample: Mapping[str, object]) -> None: + instance: Final = _leaf_instance(leaf, sample) + result: Final = _strict_leaf_validation(leaf, instance) + + assert result == instance + assert frozenset(sample) <= frozenset(field.name for field in fields(leaf)) + + +@pytest.mark.parametrize("leaf,sample", LEAF_BAD_SAMPLES.items(), ids=_leaf_id) +def test_every_owned_leaf_rejects_a_strict_wrong_typed_sample(leaf: type, sample: Mapping[str, object]) -> None: + instance: Final = _leaf_instance(leaf, sample) + + with pytest.raises(ValidationError): + _strict_leaf_validation(leaf, instance) + + +@pytest.mark.parametrize("leaf,sample", INVALID_LITERAL_SAMPLES, ids=_leaf_id) +def test_owned_leaf_literals_reject_unknown_values(leaf: type, sample: Mapping[str, object]) -> None: + instance: Final = _leaf_instance(leaf, sample) + + with pytest.raises(ValidationError): + _strict_leaf_validation(leaf, instance) + + +@pytest.mark.parametrize( + "strategy", + [ + "simple-shuffle", + "least-busy", + "usage-based-routing", + "latency-based-routing", + "cost-based-routing", + "usage-based-routing-v2", + "lar1", + ], +) +def test_routing_options_accept_every_strategy_the_router_accepts(strategy: str) -> None: + instance: Final = _leaf_instance(litellm_params.RoutingOptions, {"routing_strategy": strategy}) + + assert _strict_leaf_validation(litellm_params.RoutingOptions, instance) is instance + + +NAMES_SHARED_WITH_TYPED_MODELS: Final[Mapping[str, tuple[str, ...]]] = MappingProxyType( + { + "credentials": ( + "api_base", + "api_key", + "api_version", + "aws_batch_role_arn", + "azure_password", + "azure_scope", + "azure_username", + "bedrock_tags", + "client_id", + "client_secret", + "region_name", + "s3_access_key_id", + "s3_bucket_name", + "s3_bucket_owner", + "s3_encryption_key_id", + "s3_endpoint_url", + "s3_output_bucket_name", + "s3_region_name", + "s3_secret_access_key", + "tenant_id", + ), + "router": ( + "caching_groups", + "cooldown_time", + "enable_tag_filtering", + "fallbacks", + "max_retries", + "model_list", + "num_retries", + "retry_policy", + "routing_strategy", + ), + } +) + + +@pytest.mark.parametrize("source", TYPED_CONFIG_MODELS) +def test_names_a_typed_config_model_shares_with_the_owned_inventory_are_exactly_these(source: str) -> None: + model_names: Final = frozenset(name for model in TYPED_CONFIG_MODELS[source] for name in model.model_fields) + + assert DECLARED_NAMES & model_names == frozenset(NAMES_SHARED_WITH_TYPED_MODELS[source]) + + +@pytest.mark.parametrize("name", PRICING_NAMES) +def test_pricing_name_is_owned_by_the_pricing_model_alone(name: str) -> None: + assert name not in DECLARED_NAMES From 3fa02ef9fcc287a8648c6d5bc998d9962915699e Mon Sep 17 00:00:00 2001 From: yuneng-jiang Date: Thu, 24 Sep 2026 20:28:13 -0700 Subject: [PATCH 06/10] bump: litellm-enterprise 0.1.70 -> 0.1.71, litellm-proxy-extras 0.4.101 -> 0.4.102 (#43120) --- enterprise/pyproject.toml | 4 ++-- litellm-proxy-extras/pyproject.toml | 4 ++-- pyproject.toml | 4 ++-- uv.lock | 4 ++-- 4 files changed, 8 insertions(+), 8 deletions(-) diff --git a/enterprise/pyproject.toml b/enterprise/pyproject.toml index 8509600ad96..e5e54a3df2c 100644 --- a/enterprise/pyproject.toml +++ b/enterprise/pyproject.toml @@ -1,6 +1,6 @@ [project] name = "litellm-enterprise" -version = "0.1.70" +version = "0.1.71" description = "Package for LiteLLM Enterprise features" readme = "README.md" requires-python = ">=3.9" @@ -26,7 +26,7 @@ required-version = ">=0.10.9" module-root = "" [tool.commitizen] -version = "0.1.70" +version = "0.1.71" version_files = [ "pyproject.toml:^version", "../pyproject.toml:litellm-enterprise==", diff --git a/litellm-proxy-extras/pyproject.toml b/litellm-proxy-extras/pyproject.toml index e9a4ff90b9e..2835715ef30 100644 --- a/litellm-proxy-extras/pyproject.toml +++ b/litellm-proxy-extras/pyproject.toml @@ -1,6 +1,6 @@ [project] name = "litellm-proxy-extras" -version = "0.4.101" +version = "0.4.102" description = "Additional files for the LiteLLM Proxy. Reduces the size of the main litellm package." readme = "README.md" requires-python = ">=3.9" @@ -26,7 +26,7 @@ required-version = ">=0.10.9" module-root = "" [tool.commitizen] -version = "0.4.101" +version = "0.4.102" version_files = [ "pyproject.toml:^version", "../pyproject.toml:litellm-proxy-extras==", diff --git a/pyproject.toml b/pyproject.toml index ba72378989a..15eb8f0c4f4 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -75,8 +75,8 @@ proxy = [ "mcp>=2.2.0,<3", "httpx2>=2.5.0,<3", "pydantic>=2.12.0,<3", - "litellm-proxy-extras==0.4.101", - "litellm-enterprise==0.1.70", + "litellm-proxy-extras==0.4.102", + "litellm-enterprise==0.1.71", "RestrictedPython>=8.5,<9.0", "rich>=13.9.4,<14.0", "InquirerPy>=0.3.4,<1.0", diff --git a/uv.lock b/uv.lock index 85e2b6d4e52..8f63ca2b564 100644 --- a/uv.lock +++ b/uv.lock @@ -4959,12 +4959,12 @@ proxy-dev = [ [[package]] name = "litellm-enterprise" -version = "0.1.70" +version = "0.1.71" source = { editable = "enterprise" } [[package]] name = "litellm-proxy-extras" -version = "0.4.101" +version = "0.4.102" source = { editable = "litellm-proxy-extras" } [[package]] From 0d47347ad781b78142bba7f99a3ab1ac995041c0 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Thu, 24 Sep 2026 21:03:59 -0700 Subject: [PATCH 07/10] fix(cost): apply a deployment's pricing override to realtime sessions (#43114) * fix(cost): apply a deployment's pricing override to realtime sessions Pass the resolved custom pricing model into the realtime and transcription cost paths so model_info rates and base_model on a realtime deployment are honoured instead of the model the session reported. Adds an integration test that bills a realtime turn at the deployment's configured rates Carries the fix from #36958 Co-authored-by: Marty Sullivan Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(cost): honour audio-only and base_model realtime pricing overrides Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(cost): keep flat per-unit prices from claiming the deployment pricing key Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(cost): keep base_model out of realtime transcription rate overrides Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(cost): type the realtime pricing test parameters Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(cost): try a realtime deployment's base_model ahead of the session model Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: kerry Co-authored-by: Marty Sullivan Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/cost_calculator.py | 103 ++- .../pricing/test_configured_prices.py | 93 +++ .../test_realtime_cached_audio_pricing.py | 4 +- tests/test_litellm/test_cost_calculator.py | 622 ++++++++++++++++++ 4 files changed, 796 insertions(+), 26 deletions(-) diff --git a/litellm/cost_calculator.py b/litellm/cost_calculator.py index 6cc0d9444cd..a279b9f0903 100644 --- a/litellm/cost_calculator.py +++ b/litellm/cost_calculator.py @@ -790,6 +790,16 @@ def _get_hidden_str_for_cost_calc(hidden_params: object, key: str) -> str | None return value if isinstance(value, str) and value else None +_NON_TOKEN_RATE_FIELDS: Final = frozenset({"input_cost_per_second", "input_cost_per_query", "tiered_pricing"}) + + +def _cost_map_entry_prices_anything(entry: Mapping[str, object]) -> bool: + return any( + value is not None and (field in _NON_TOKEN_RATE_FIELDS or ("cost_per" in field and "token" in field)) + for field, value in entry.items() + ) + + def _select_model_name_for_cost_calc( model: str | None, completion_response: object | None, @@ -828,12 +838,7 @@ def _select_model_name_for_cost_calc( if custom_pricing is True: if router_model_id is not None and router_model_id in litellm.model_cost: entry: Final = litellm.model_cost[router_model_id] - if ( - entry.get("input_cost_per_token") is not None - or entry.get("input_cost_per_second") is not None - or entry.get("input_cost_per_query") is not None - or entry.get("tiered_pricing") is not None - ): + if _cost_map_entry_prices_anything(entry): return_model = router_model_id else: return_model = model @@ -1699,6 +1704,8 @@ def completion_cost( litellm_model_name=model, data_residency=data_residency, litellm_logging_obj=litellm_logging_obj, + custom_pricing_model=selected_model if custom_pricing else None, + base_pricing_model=(selected_model if base_model is not None and not custom_pricing else None), ) elif call_type == _MCP_CALL_TYPE: from litellm.proxy._experimental.mcp_server.cost_calculator import ( @@ -2870,14 +2877,20 @@ def _candidate_realtime_token_costs( def _cost_map_entry_declares_pricing(model_name: str, custom_llm_provider: str) -> bool: + """Whether the entry behind ``model_name`` sets any rate of its own, even a zero one. + + The name is resolved the way ``get_model_info`` resolves it before the raw entry is read, + because a deployment-scoped name arrives here already carrying its provider prefix. Two raw + lookups cannot strip that prefix, so a zero-rated override read as declaring nothing, and a + session that should bill nothing fell through to the public rates instead. + """ + resolved: Final = _get_model_info_or_none(model_name, custom_llm_provider) entries: Final = ( + litellm.model_cost.get(resolved.get("key")) if resolved is not None else None, litellm.model_cost.get(model_name), litellm.model_cost.get(f"{custom_llm_provider}/{model_name}"), ) - return any( - entry is not None and any("cost_per" in field and value is not None for field, value in entry.items()) - for entry in entries - ) + return any(entry is not None and _cost_map_entry_prices_anything(entry) for entry in entries) def _first_priced_realtime_token_costs( @@ -2917,6 +2930,8 @@ def handle_realtime_stream_cost_calculation( litellm_model_name: str, data_residency: str | None = None, litellm_logging_obj: LitellmLoggingObject | None = None, + custom_pricing_model: str | None = None, + base_pricing_model: str | None = None, ) -> float: """ Handles the cost calculation for realtime stream responses. @@ -2925,9 +2940,13 @@ def handle_realtime_stream_cost_calculation( Args: results: A list of OpenAIRealtimeStreamBaseObject objects + custom_pricing_model: deployment-scoped pricing key from the deployment's + custom rates, tried ahead of the session-reported model + base_pricing_model: the deployment's resolved base_model, tried ahead of the + session-reported model but after custom rates """ received_model = None - potential_model_names: Final = [] + potential_model_names: Final = [custom_pricing_model, base_pricing_model] for result in results: if result["type"] == "session.created": received_model = cast(OpenAIRealtimeStreamSessionEvents, result)["session"].get("model", None) @@ -2945,6 +2964,7 @@ def handle_realtime_stream_cost_calculation( results=results, custom_llm_provider=custom_llm_provider, litellm_model_name=litellm_model_name, + custom_pricing_model=custom_pricing_model, ) if any(r.get("type") == _TRANSCRIPTION_COMPLETED_EVENT_TYPE for r in results) else 0.0 @@ -2968,6 +2988,7 @@ def handle_realtime_transcription_cost_calculation( results: OpenAIRealtimeStreamList, custom_llm_provider: str, litellm_model_name: str, + custom_pricing_model: str | None = None, ) -> float: """ Cost for realtime transcription sessions (e.g. gpt-realtime-whisper). @@ -2985,15 +3006,15 @@ def handle_realtime_transcription_cost_calculation( return 0.0 model_name: Final = _get_transcription_model_name_from_results(results) or litellm_model_name - try: - model_info = litellm.get_model_info(model=model_name, custom_llm_provider=custom_llm_provider) - except Exception: - model_info = None + model_info: Final = _get_model_info_or_none(model_name, custom_llm_provider) + override_info: Final = ( + _get_model_info_or_none(custom_pricing_model, custom_llm_provider) if custom_pricing_model is not None else None + ) total_cost = 0.0 for event in completed_events: usage = event.get("usage") or {} - total_cost += _transcription_usage_cost(usage, model_info) + total_cost += _transcription_usage_cost(usage, model_info, override_info) return total_cost @@ -3018,23 +3039,57 @@ def _get_transcription_model_name_from_results( return None -def _transcription_usage_cost(usage: dict, model_info: ModelInfo | None) -> float: - if model_info is None: +def _get_model_info_or_none(model: str, custom_llm_provider: str) -> ModelInfo | None: + try: + return litellm.get_model_info(model=model, custom_llm_provider=custom_llm_provider) + except Exception: + return None + + +def _declared_transcription_rate(info: ModelInfo | None, keys: tuple[str, ...]) -> float | None: + """First of ``keys`` this entry prices, read off the raw ``litellm.model_cost`` entry + because ``get_model_info`` synthesizes zero token rates for entries that omit them.""" + if info is None: + return None + declared: Final = litellm.model_cost.get(info.get("key")) + if declared is None: + return None + return next( + (float(value) for key in keys if declared.get(key) is not None and (value := info.get(key)) is not None), + None, + ) + + +def _transcription_rate(keys: tuple[str, ...], override: ModelInfo | None, base: ModelInfo | None) -> float: + rates: Final = (_declared_transcription_rate(info, keys) for info in (override, base)) + return next((rate for rate in rates if rate is not None), 0.0) + + +def _transcription_usage_cost( + usage: dict, + model_info: ModelInfo | None, + override_info: ModelInfo | None = None, +) -> float: + if model_info is None and override_info is None: return 0.0 + usage_type: Final = usage.get("type") if usage_type == "duration": seconds: Final = usage.get("seconds") or 0.0 - per_second: Final = model_info.get("input_cost_per_second") or 0.0 - return float(seconds) * float(per_second) + return float(seconds) * _transcription_rate(("input_cost_per_second",), override_info, model_info) if usage_type == "tokens": input_token_details: Final = usage.get("input_token_details") or {} audio_tokens: Final = input_token_details.get("audio_tokens") or 0 text_tokens: Final = input_token_details.get("text_tokens") or 0 output_tokens: Final = usage.get("output_tokens") or 0 - audio_cost: Final = float(audio_tokens) * float( - model_info.get("input_cost_per_audio_token") or model_info.get("input_cost_per_token") or 0.0 + audio_cost: Final = float(audio_tokens) * _transcription_rate( + ("input_cost_per_audio_token", "input_cost_per_token"), override_info, model_info + ) + text_cost: Final = float(text_tokens) * _transcription_rate( + ("input_cost_per_token",), override_info, model_info + ) + output_cost: Final = float(output_tokens) * _transcription_rate( + ("output_cost_per_token",), override_info, model_info ) - text_cost: Final = float(text_tokens) * float(model_info.get("input_cost_per_token") or 0.0) - output_cost: Final = float(output_tokens) * float(model_info.get("output_cost_per_token") or 0.0) return audio_cost + text_cost + output_cost return 0.0 diff --git a/tests/integration/pricing/test_configured_prices.py b/tests/integration/pricing/test_configured_prices.py index 0e4efea3a15..93290a404fc 100644 --- a/tests/integration/pricing/test_configured_prices.py +++ b/tests/integration/pricing/test_configured_prices.py @@ -1,6 +1,9 @@ +import asyncio import json +import os import uuid from collections.abc import Iterator, Mapping +from hashlib import sha256 from pathlib import Path from typing import Final @@ -12,6 +15,96 @@ from litellm import get_model_info from tests.integration._support.client import Gateway, eventually, object_value, string_value from tests.integration._support.database import read_rows from tests.integration._support.process import owned_proxy +from tests.integration._support.upstream import delete_scenario, register_scenario +from tests.integration.cost_calculation.cost_tracking_case import RealtimeResponse +from tests.integration.pricing.test_realtime_cached_audio_pricing import one_realtime_turn + +REALTIME_MODEL: Final = "gpt-realtime-2" +REALTIME_INPUT_TEXT_TOKENS: Final = 10 +REALTIME_INPUT_AUDIO_TOKENS: Final = 20 +REALTIME_OUTPUT_TEXT_TOKENS: Final = 5 +REALTIME_OUTPUT_AUDIO_TOKENS: Final = 7 + + +def _realtime_response_done() -> RealtimeResponse: + return RealtimeResponse( + content_type="application/x-realtime", + events=( + { + "type": "response.done", + "event_id": "evt_$REQUEST_ID", + "response": { + "id": "resp_$REQUEST_ID", + "object": "realtime.response", + "status": "completed", + "output": [], + "usage": { + "total_tokens": REALTIME_INPUT_TEXT_TOKENS + + REALTIME_INPUT_AUDIO_TOKENS + + REALTIME_OUTPUT_TEXT_TOKENS + + REALTIME_OUTPUT_AUDIO_TOKENS, + "input_tokens": REALTIME_INPUT_TEXT_TOKENS + REALTIME_INPUT_AUDIO_TOKENS, + "output_tokens": REALTIME_OUTPUT_TEXT_TOKENS + REALTIME_OUTPUT_AUDIO_TOKENS, + "input_token_details": { + "text_tokens": REALTIME_INPUT_TEXT_TOKENS, + "audio_tokens": REALTIME_INPUT_AUDIO_TOKENS, + "cached_tokens": 0, + }, + "output_token_details": { + "text_tokens": REALTIME_OUTPUT_TEXT_TOKENS, + "audio_tokens": REALTIME_OUTPUT_AUDIO_TOKENS, + }, + }, + }, + }, + ), + ) + + +@pytest.mark.parametrize( + ("input_text_rate", "input_audio_rate", "output_text_rate", "output_audio_rate"), + ((0.001, 0.002, 0.003, 0.004), (0.0, 0.0, 0.0, 0.0)), + ids=("custom_rates", "zero_rated"), +) +def test_realtime_session_is_charged_at_the_deployment_configured_rates( + gateway: Gateway, + input_text_rate: float, + input_audio_rate: float, + output_text_rate: float, + output_audio_rate: float, +) -> None: + with gateway.scenario() as scenario: + scenario_id: Final = f"realtime-configured-price-{uuid.uuid4().hex[:12]}" + handle: Final = register_scenario(scenario_id, _realtime_response_done()) + scenario.cleanups.callback(delete_scenario, handle) + key: Final = scenario.key() + model: Final = scenario.model( + model=f"openai/{REALTIME_MODEL}", + api_key=scenario_id, + api_base=gateway.upstream_url.rstrip("/"), + input_cost_per_token=input_text_rate, + input_cost_per_audio_token=input_audio_rate, + output_cost_per_token=output_text_rate, + output_cost_per_audio_token=output_audio_rate, + ) + session: Final = asyncio.run(one_realtime_turn(os.environ["INTEGRATION_PROXY_URL"].rstrip("/"), key, model)) + assert session.get("type") == "session.created", session + rows: Final = eventually( + lambda: read_rows( + 'SELECT spend, call_type FROM "LiteLLM_SpendLogs" WHERE api_key = %s', + (sha256(key.encode()).hexdigest(),), + ), + lambda values: len(values) == 1, + seconds=70, + ) + assert rows[0]["call_type"] == "_arealtime", rows + assert float(str(rows[0]["spend"])) == pytest.approx( + REALTIME_INPUT_TEXT_TOKENS * input_text_rate + + REALTIME_INPUT_AUDIO_TOKENS * input_audio_rate + + REALTIME_OUTPUT_TEXT_TOKENS * output_text_rate + + REALTIME_OUTPUT_AUDIO_TOKENS * output_audio_rate, + abs=1e-9, + ), rows @pytest.mark.covers("quota_management.spend_tracking.custom_price.matches_input_rates") diff --git a/tests/integration/pricing/test_realtime_cached_audio_pricing.py b/tests/integration/pricing/test_realtime_cached_audio_pricing.py index 4a7598d0cbf..42e90cf2c3e 100644 --- a/tests/integration/pricing/test_realtime_cached_audio_pricing.py +++ b/tests/integration/pricing/test_realtime_cached_audio_pricing.py @@ -92,7 +92,7 @@ def cached_audio_response_done() -> RealtimeResponse: ) -async def _one_realtime_turn(proxy_url: str, key: str, model: str) -> dict[str, JsonValue]: +async def one_realtime_turn(proxy_url: str, key: str, model: str) -> dict[str, JsonValue]: async with websockets.connect( f"{proxy_url.replace('http://', 'ws://').replace('https://', 'wss://')}/v1/realtime?model={model}", additional_headers={"Authorization": f"Bearer {key}"}, @@ -115,7 +115,7 @@ def test_realtime_cached_audio_tokens_bill_at_audio_cache_read_rate_not_full_aud model: Final = scenario.model( model=f"openai/{MODEL}", api_key=scenario_id, api_base=gateway.upstream_url.rstrip("/") ) - session: Final = asyncio.run(_one_realtime_turn(os.environ["INTEGRATION_PROXY_URL"].rstrip("/"), key, model)) + session: Final = asyncio.run(one_realtime_turn(os.environ["INTEGRATION_PROXY_URL"].rstrip("/"), key, model)) assert session.get("type") == "session.created", session rows: Final = eventually( lambda: read_rows( diff --git a/tests/test_litellm/test_cost_calculator.py b/tests/test_litellm/test_cost_calculator.py index f7d6cfaf079..99dea6366f9 100644 --- a/tests/test_litellm/test_cost_calculator.py +++ b/tests/test_litellm/test_cost_calculator.py @@ -309,6 +309,243 @@ def test_realtime_logging_object_does_not_validate_unknown_event_types(): assert len(dumped["results"]) == len(results) +def test_realtime_transcription_honors_deployment_pricing_override(monkeypatch: pytest.MonkeyPatch) -> None: + """A deployment's pricing override must reach transcription events too. + + Transcription is billed separately from response usage inside the same realtime + session, so a deployment registered at zero rates has to zero both. Resolving + transcription against the public ASR model instead billed a zero-rated + deployment for every .completed event. + """ + monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True") + monkeypatch.setattr(litellm, "model_cost", litellm.get_model_cost_map(url="")) + + deployment_id = "deployment-hash-zero-rated-asr" + litellm.register_model( + model_cost={ + deployment_id: { + "litellm_provider": "openai", + "mode": "realtime", + "input_cost_per_second": 0.0, + "input_cost_per_token": 0.0, + "output_cost_per_token": 0.0, + "input_cost_per_audio_token": 0.0, + } + } + ) + + results: OpenAIRealtimeStreamList = [ + { + "type": "session.created", + "session": { + "type": "transcription", + "audio": {"input": {"transcription": {"model": "gpt-realtime-whisper"}}}, + }, + }, + { + "type": "conversation.item.input_audio_transcription.completed", + "usage": {"type": "duration", "seconds": 120.0}, + }, + ] + + public_rate_cost = 120.0 * litellm.model_cost["gpt-realtime-whisper"]["input_cost_per_second"] + assert public_rate_cost > 0, "the public ASR rate must be non-zero for this test to mean anything" + + without_override = handle_realtime_stream_cost_calculation( + results=results, + combined_usage_object=Usage(), + custom_llm_provider="openai", + litellm_model_name="gpt-realtime-whisper", + ) + assert abs(without_override - public_rate_cost) < 1e-9 + + with_override = handle_realtime_stream_cost_calculation( + results=results, + combined_usage_object=Usage(), + custom_llm_provider="openai", + litellm_model_name="gpt-realtime-whisper", + custom_pricing_model=deployment_id, + ) + assert with_override == 0.0, "the zero-rated deployment must not be billed for transcription" + + +def test_realtime_transcription_partial_override_keeps_unset_rates(monkeypatch: pytest.MonkeyPatch) -> None: + """An override must not blank the rates it does not set. + + A deployment that prices tokens but omits input_cost_per_second would otherwise + bill duration-based transcription at nothing, because the cost helpers read + `.get(key) or 0.0`. Only the fields the operator actually set may win. + """ + monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True") + monkeypatch.setattr(litellm, "model_cost", litellm.get_model_cost_map(url="")) + + deployment_id = "deployment-hash-tokens-only" + litellm.register_model( + model_cost={ + deployment_id: { + "litellm_provider": "openai", + "mode": "realtime", + "input_cost_per_token": 0.0, + "output_cost_per_token": 0.0, + } + } + ) + + results: OpenAIRealtimeStreamList = [ + { + "type": "session.created", + "session": { + "type": "transcription", + "audio": {"input": {"transcription": {"model": "gpt-realtime-whisper"}}}, + }, + }, + { + "type": "conversation.item.input_audio_transcription.completed", + "usage": {"type": "duration", "seconds": 120.0}, + }, + ] + + cost = handle_realtime_stream_cost_calculation( + results=results, + combined_usage_object=Usage(), + custom_llm_provider="openai", + litellm_model_name="gpt-realtime-whisper", + custom_pricing_model=deployment_id, + ) + + expected = 120.0 * litellm.model_cost["gpt-realtime-whisper"]["input_cost_per_second"] + assert expected > 0, "the public ASR per-second rate must be non-zero for this test to mean anything" + assert cost == pytest.approx(expected, rel=1e-9), ( + "duration must keep the ASR per-second rate the override left unset" + ) + + +@pytest.mark.parametrize( + "label,override,expected_audio_rate,expected_per_second", + [ + ("tokens only", {"input_cost_per_token": 0.0}, 0.0, 0.017 / 60), + ("audio zeroed", {"input_cost_per_audio_token": 0.0}, 0.0, 0.017 / 60), + ("per second only", {"input_cost_per_second": 0.001}, 6e-06, 0.001), + ("empty override", {}, 6e-06, 0.017 / 60), + ("no override", None, 6e-06, 0.017 / 60), + ], +) +def test_transcription_rate_precedence( + monkeypatch: pytest.MonkeyPatch, + label: str, + override: dict[str, float] | None, + expected_audio_rate: float, + expected_per_second: float, +) -> None: + """Rates resolve within one entry before moving to the next, and zero is a real value. + + An override that prices only tokens must apply its own token rate to audio rather + than reaching past itself for the public audio rate, a deliberate zero must win + instead of being treated as unset, and a rate the override never mentions must keep + the base entry's value. + """ + from litellm.cost_calculator import handle_realtime_transcription_cost_calculation + + monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True") + monkeypatch.setattr(litellm, "model_cost", litellm.get_model_cost_map(url="")) + + base_model = "asr-precedence-base" + deployment_id = "asr-precedence-deployment" + litellm.register_model( + model_cost={ + base_model: { + "litellm_provider": "openai", + "mode": "audio_transcription", + "input_cost_per_audio_token": 6e-06, + "input_cost_per_token": 2.5e-06, + "input_cost_per_second": 0.017 / 60, + } + } + ) + if override is not None: + litellm.register_model( + model_cost={deployment_id: {"litellm_provider": "openai", "mode": "audio_transcription", **override}} + ) + + def cost_for(usage: dict[str, object]) -> float: + return handle_realtime_transcription_cost_calculation( + results=[ + {"type": "transcription_session.created", "session": {"model": base_model}}, + {"type": "conversation.item.input_audio_transcription.completed", "usage": usage}, + ], + custom_llm_provider="openai", + litellm_model_name=base_model, + custom_pricing_model=deployment_id if override is not None else None, + ) + + audio_cost = cost_for({"type": "tokens", "input_token_details": {"audio_tokens": 100}}) + assert audio_cost == pytest.approx(100 * expected_audio_rate, rel=1e-9), f"{label}: audio rate" + + per_second_cost = cost_for({"type": "duration", "seconds": 120.0}) + assert per_second_cost == pytest.approx(120.0 * expected_per_second, rel=1e-9), ( + f"{label}: an override must never blank a rate it does not set" + ) + + +def test_realtime_transcription_per_second_override_keeps_public_token_rates( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """A per-second override must not zero the token rates ``get_model_info`` synthesizes. + + ``get_model_info`` defaults input_cost_per_token and output_cost_per_token to 0 for entries + that omit them, so a deployment priced only per second looked like it had declared token + rates of 0. Token-shaped transcription then billed nothing instead of falling through to the + public ASR rates, while the per-second rate the operator did set stayed in force. + """ + from litellm.cost_calculator import handle_realtime_transcription_cost_calculation + + monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True") + monkeypatch.setattr(litellm, "model_cost", litellm.get_model_cost_map(url="")) + + asr_model = "gpt-4o-transcribe" + per_second_rate = 0.001 + deployment_id = "deployment-hash-per-second-only" + litellm.register_model( + model_cost={ + deployment_id: { + "litellm_provider": "openai", + "mode": "audio_transcription", + "input_cost_per_second": per_second_rate, + } + } + ) + + public = litellm.model_cost[asr_model] + session_event = {"type": "transcription_session.created", "session": {"model": asr_model}} + + def cost_for(usage: dict[str, object]) -> float: + return handle_realtime_transcription_cost_calculation( + results=[session_event, {"type": "conversation.item.input_audio_transcription.completed", "usage": usage}], + custom_llm_provider="openai", + litellm_model_name=asr_model, + custom_pricing_model=deployment_id, + ) + + token_cost = cost_for( + { + "type": "tokens", + "input_token_details": {"audio_tokens": 400, "text_tokens": 12}, + "output_tokens": 30, + } + ) + expected_token_cost = ( + 400 * public["input_cost_per_audio_token"] + + 12 * public["input_cost_per_token"] + + 30 * public["output_cost_per_token"] + ) + assert expected_token_cost > 0, "the public ASR token rates must be non-zero for this test to mean anything" + assert token_cost == pytest.approx(expected_token_cost, rel=1e-9), ( + "an override that prices only seconds must leave the public token rates in place" + ) + + assert cost_for({"type": "duration", "seconds": 120.0}) == pytest.approx(120.0 * per_second_rate, rel=1e-9) + + def test_realtime_transcription_no_completed_events_is_zero(monkeypatch): """A realtime stream without transcription completed events adds no extra cost.""" monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True") @@ -4635,6 +4872,391 @@ def test_gemini_live_native_audio_limits_and_capabilities_match_vendor_model_car assert info["supports_pdf_input"] is False +def test_realtime_honours_deployment_custom_pricing(monkeypatch: pytest.MonkeyPatch) -> None: + """Regression: a deployment's pricing override never reached realtime costing. + + `model_info` overrides are registered under the deployment's own model_id, and + only `_select_model_name_for_cost_calc` knows to look there. The realtime branch + discarded that result and priced by the model the session reported, so a config + that zeroes a realtime deployment was billed at the public rate anyway. Audio is + the bulk of a voice call, so the gap was most of the cost. + """ + from litellm.types.utils import CompletionTokensDetailsWrapper + + model = "gemini-3.1-flash-live-preview" + deployment_key = "deployment-id-for-a-zero-rated-realtime-group" + paid = litellm.model_cost[model] + monkeypatch.setitem( + litellm.model_cost, + deployment_key, + { + **paid, + "input_cost_per_token": 0.0, + "output_cost_per_token": 0.0, + "input_cost_per_audio_token": 0.0, + "output_cost_per_audio_token": 0.0, + "cache_read_input_token_cost": 0.0, + }, + ) + + results: OpenAIRealtimeStreamList = [ + {"type": "session.created", "session": {"model": model}}, + { + "type": "response.done", + "response": {"usage": {"input_tokens": 10, "output_tokens": 200, "total_tokens": 210}}, + }, + ] + usage = Usage( + prompt_tokens=10, + completion_tokens=200, + total_tokens=210, + prompt_tokens_details=PromptTokensDetailsWrapper(text_tokens=10, cached_tokens=0), + completion_tokens_details=CompletionTokensDetailsWrapper(text_tokens=20, audio_tokens=180), + ) + + paid_cost = handle_realtime_stream_cost_calculation( + results=results, + combined_usage_object=usage, + custom_llm_provider="gemini", + litellm_model_name=model, + ) + expected_paid = ( + 10 * paid["input_cost_per_token"] + + 20 * paid["output_cost_per_token"] + + 180 * paid["output_cost_per_audio_token"] + ) + assert paid_cost == pytest.approx(expected_paid, rel=1e-9) + assert paid_cost > 0 + + zero_rated_cost = handle_realtime_stream_cost_calculation( + results=results, + combined_usage_object=usage, + custom_llm_provider="gemini", + litellm_model_name=model, + custom_pricing_model=deployment_key, + ) + assert zero_rated_cost == 0.0 + + +def test_realtime_honours_a_provider_prefixed_zero_rated_deployment(monkeypatch: pytest.MonkeyPatch) -> None: + """Regression: the override arrived provider-prefixed and was read as pricing nothing. + + `_select_model_name_for_cost_calc` hands back `/`, so the name reaching + the pricing guard carries a prefix the raw cost-map lookups cannot strip. The rates resolved + correctly through `get_model_info`, then the guard rejected them as undeclared and the session + billed the public rates. A zero-rated deployment must stay at zero however its name arrives. + """ + from litellm.types.utils import CompletionTokensDetailsWrapper + + model = "gemini-live-2.5-flash-native-audio" + deployment_key = "deployment-id-for-a-prefixed-zero-rated-realtime-group" + paid = litellm.model_cost[model] + monkeypatch.setitem( + litellm.model_cost, + deployment_key, + { + **paid, + "input_cost_per_token": 0.0, + "output_cost_per_token": 0.0, + "input_cost_per_audio_token": 0.0, + "output_cost_per_audio_token": 0.0, + }, + ) + + results: OpenAIRealtimeStreamList = [ + {"type": "session.created", "session": {"model": model}}, + { + "type": "response.done", + "response": {"usage": {"input_tokens": 219, "output_tokens": 81, "total_tokens": 300}}, + }, + ] + usage = Usage( + prompt_tokens=219, + completion_tokens=81, + total_tokens=300, + prompt_tokens_details=PromptTokensDetailsWrapper(text_tokens=16, audio_tokens=203), + completion_tokens_details=CompletionTokensDetailsWrapper(text_tokens=23, audio_tokens=58), + ) + + paid_cost = handle_realtime_stream_cost_calculation( + results=results, + combined_usage_object=usage, + custom_llm_provider="vertex_ai", + litellm_model_name=model, + ) + assert paid_cost > 0 + + zero_rated_cost = handle_realtime_stream_cost_calculation( + results=results, + combined_usage_object=usage, + custom_llm_provider="vertex_ai", + litellm_model_name=model, + custom_pricing_model=f"vertex_ai/{deployment_key}", + ) + assert zero_rated_cost == 0.0 + + +def test_unpriced_deployment_entry_still_falls_through_to_the_session_model( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """The guard's own purpose must survive: an entry that prices nothing is not an override. + + Deployments are auto-registered under their model_id with no rates at all, and those must + keep billing at the session model's public rates rather than silently costing nothing. + """ + from litellm.types.utils import CompletionTokensDetailsWrapper + + model = "gemini-live-2.5-flash-native-audio" + deployment_key = "deployment-id-with-no-declared-rates" + monkeypatch.setitem( + litellm.model_cost, + deployment_key, + {key: value for key, value in litellm.model_cost[model].items() if "cost_per" not in key}, + ) + + results: OpenAIRealtimeStreamList = [ + {"type": "session.created", "session": {"model": model}}, + { + "type": "response.done", + "response": {"usage": {"input_tokens": 219, "output_tokens": 81, "total_tokens": 300}}, + }, + ] + usage = Usage( + prompt_tokens=219, + completion_tokens=81, + total_tokens=300, + prompt_tokens_details=PromptTokensDetailsWrapper(text_tokens=16, audio_tokens=203), + completion_tokens_details=CompletionTokensDetailsWrapper(text_tokens=23, audio_tokens=58), + ) + + with_unpriced_override = handle_realtime_stream_cost_calculation( + results=results, + combined_usage_object=usage, + custom_llm_provider="vertex_ai", + litellm_model_name=model, + custom_pricing_model=f"vertex_ai/{deployment_key}", + ) + without_override = handle_realtime_stream_cost_calculation( + results=results, + combined_usage_object=usage, + custom_llm_provider="vertex_ai", + litellm_model_name=model, + ) + assert with_unpriced_override == pytest.approx(without_override, rel=1e-9) + assert with_unpriced_override > 0 + + +def test_realtime_audio_only_override_bills_audio_at_the_deployment_rate( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """Regression: an audio-only pricing override was never selected as the pricing key. + + The deployment-selection guard recognised only text, per-second, per-query and + tiered rates, so a deployment that priced just the audio meters was passed over + and the session kept billing the public rates for the exact tokens it priced. + """ + from litellm.types.utils import CompletionTokensDetailsWrapper + + model = "gemini-live-2.5-flash-native-audio" + deployment_key = "deployment-id-for-an-audio-only-realtime-group" + monkeypatch.setitem( + litellm.model_cost, + deployment_key, + { + "litellm_provider": "vertex_ai", + "mode": "realtime", + "input_cost_per_audio_token": 0.0, + "output_cost_per_audio_token": 0.0, + }, + ) + + logging_object = LiteLLMRealtimeStreamLoggingObject( + usage=Usage( + prompt_tokens=203, + completion_tokens=58, + total_tokens=261, + prompt_tokens_details=PromptTokensDetailsWrapper(audio_tokens=203), + completion_tokens_details=CompletionTokensDetailsWrapper(audio_tokens=58), + ), + results=[ + {"type": "session.created", "session": {"model": model}}, + { + "type": "response.done", + "response": {"usage": {"input_tokens": 203, "output_tokens": 58, "total_tokens": 261}}, + }, + ], + ) + + public_cost = completion_cost( + completion_response=logging_object, + model=model, + call_type=CallTypes.arealtime.value, + custom_llm_provider="vertex_ai", + ) + assert public_cost > 0 + + overridden_cost = completion_cost( + completion_response=logging_object, + model=model, + call_type=CallTypes.arealtime.value, + custom_llm_provider="vertex_ai", + custom_pricing=True, + router_model_id=deployment_key, + ) + assert overridden_cost == pytest.approx(0.0) + + +def test_realtime_session_falls_back_to_base_model_pricing(monkeypatch: pytest.MonkeyPatch) -> None: + """Regression: a priced base_model was discarded for realtime sessions. + + The resolved base model only reached the realtime cost path when custom pricing + was on, so a session reporting an alias unmapped in the cost map recorded zero + instead of the base model's published price. + """ + from litellm.types.utils import CompletionTokensDetailsWrapper + + base_model = "gemini-live-2.5-flash-native-audio" + logging_object = LiteLLMRealtimeStreamLoggingObject( + usage=Usage( + prompt_tokens=219, + completion_tokens=81, + total_tokens=300, + prompt_tokens_details=PromptTokensDetailsWrapper(text_tokens=16, audio_tokens=203), + completion_tokens_details=CompletionTokensDetailsWrapper(text_tokens=23, audio_tokens=58), + ), + results=[ + { + "type": "session.created", + "session": {"model": "my-voice-alias"}, + }, + { + "type": "response.done", + "response": {"usage": {"input_tokens": 219, "output_tokens": 81, "total_tokens": 300}}, + }, + ], + ) + + aliased_cost = completion_cost( + completion_response=logging_object, + model="my-voice-alias", + call_type=CallTypes.arealtime.value, + custom_llm_provider="vertex_ai", + base_model=base_model, + ) + base_cost = completion_cost( + completion_response=logging_object, + model=base_model, + call_type=CallTypes.arealtime.value, + custom_llm_provider="vertex_ai", + ) + assert aliased_cost == pytest.approx(base_cost, rel=1e-9) + assert aliased_cost > 0 + + +def test_base_model_does_not_override_transcription_rates(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True") + monkeypatch.setattr(litellm, "model_cost", litellm.get_model_cost_map(url="")) + + base_model = "gpt-realtime-2" + asr_model = "gpt-4o-transcribe" + logging_object = LiteLLMRealtimeStreamLoggingObject( + usage=Usage(), + results=[ + { + "type": "session.created", + "session": { + "model": "my-voice-alias", + "audio": {"input": {"transcription": {"model": asr_model}}}, + }, + }, + { + "type": "conversation.item.input_audio_transcription.completed", + "usage": { + "type": "tokens", + "input_token_details": {"audio_tokens": 400, "text_tokens": 12}, + "output_tokens": 30, + }, + }, + ], + ) + + with_base_model = completion_cost( + completion_response=logging_object, + model="my-voice-alias", + call_type=CallTypes.arealtime.value, + custom_llm_provider="openai", + base_model=base_model, + ) + asr_priced = completion_cost( + completion_response=logging_object, + model="my-voice-alias", + call_type=CallTypes.arealtime.value, + custom_llm_provider="openai", + ) + realtime_card = litellm.model_cost[base_model] + billed_at_realtime = ( + 400 * realtime_card["input_cost_per_audio_token"] + + 12 * realtime_card["input_cost_per_token"] + + 30 * realtime_card["output_cost_per_audio_token"] + ) + assert billed_at_realtime != pytest.approx(asr_priced, rel=1e-9) + assert with_base_model == pytest.approx(asr_priced, rel=1e-9) + assert with_base_model > 0 + + +def test_realtime_base_model_outranks_the_session_reported_model(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True") + monkeypatch.setattr(litellm, "model_cost", litellm.get_model_cost_map(url="")) + + from litellm.types.utils import CompletionTokensDetailsWrapper + + session_model = "gpt-realtime-mini" + base_model = "gpt-realtime-2" + + def logging_object_for(session: str) -> LiteLLMRealtimeStreamLoggingObject: + return LiteLLMRealtimeStreamLoggingObject( + usage=Usage( + prompt_tokens=120, + completion_tokens=60, + total_tokens=180, + prompt_tokens_details=PromptTokensDetailsWrapper(text_tokens=20, audio_tokens=100), + completion_tokens_details=CompletionTokensDetailsWrapper(text_tokens=10, audio_tokens=50), + ), + results=[ + { + "type": "session.created", + "session": {"model": session}, + }, + { + "type": "response.done", + "response": {"usage": {"input_tokens": 120, "output_tokens": 60, "total_tokens": 180}}, + }, + ], + ) + + with_base_model = completion_cost( + completion_response=logging_object_for(session_model), + model=session_model, + call_type=CallTypes.arealtime.value, + custom_llm_provider="openai", + base_model=base_model, + ) + base_priced = completion_cost( + completion_response=logging_object_for(base_model), + model=base_model, + call_type=CallTypes.arealtime.value, + custom_llm_provider="openai", + ) + session_priced = completion_cost( + completion_response=logging_object_for(session_model), + model=session_model, + call_type=CallTypes.arealtime.value, + custom_llm_provider="openai", + ) + assert base_priced != pytest.approx(session_priced, rel=1e-9) + assert with_base_model == pytest.approx(base_priced, rel=1e-9) + + def test_baseten_glm_5_3_fast_is_priced_from_registry(_local_model_cost_map: None) -> None: model: Final = "baseten/zai-org/GLM-5.3-Fast" prompt_tokens: Final = 1000 From 1f8997398eab139f95159d7e2dba8ac88b14ff08 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Fri, 25 Sep 2026 04:20:06 +0000 Subject: [PATCH 08/10] refactor(rust): extract the host coroutine into its own crate (#43129) Move the generic async coroutine out of the host crate into litellm-coroutine, with its requirements documented in the crate's AGENTS.md. RouteMachine becomes CallMachine on top of it, host protocol ops move to host/src/protocol.rs, and the messages, OCR, host-python, python-bridge and legacy callbacks crates adopt the new types. Adds error definition rules to litellm-rust/AGENTS.md. Co-authored-by: Yujong Lee --- litellm-rust/AGENTS.md | 8 + litellm-rust/Cargo.lock | 11 + litellm-rust/Cargo.toml | 1 + .../callbacks-legacy-python/src/call.rs | 14 +- .../crates/core/src/messages/route.rs | 49 +-- litellm-rust/crates/core/src/ocr/document.rs | 6 +- litellm-rust/crates/core/src/ocr/route.rs | 343 +++++---------- litellm-rust/crates/core/src/ocr/types.rs | 9 - litellm-rust/crates/coroutine/AGENTS.md | 31 ++ litellm-rust/crates/coroutine/Cargo.toml | 15 + litellm-rust/crates/coroutine/src/co.rs | 42 ++ .../crates/coroutine/src/coroutine.rs | 94 ++++ litellm-rust/crates/coroutine/src/error.rs | 14 + litellm-rust/crates/coroutine/src/lib.rs | 12 + litellm-rust/crates/coroutine/src/reply.rs | 60 +++ .../crates/coroutine/tests/coroutine.rs | 256 +++++++++++ litellm-rust/crates/host-python/AGENTS.md | 4 +- litellm-rust/crates/host-python/Cargo.toml | 1 + .../crates/host-python/src/adapter.rs | 39 +- litellm-rust/crates/host-python/src/driver.rs | 401 +++++++++--------- .../crates/host-python/src/file_reader.rs | 241 +++++++++++ litellm-rust/crates/host-python/src/lib.rs | 6 +- litellm-rust/crates/host/Cargo.toml | 1 + litellm-rust/crates/host/src/host.rs | 38 +- litellm-rust/crates/host/src/lib.rs | 7 +- litellm-rust/crates/host/src/machine/auth.rs | 27 +- .../crates/host/src/machine/call_machine.rs | 137 ++++++ litellm-rust/crates/host/src/machine/mod.rs | 32 +- .../crates/host/src/machine/route_machine.rs | 199 --------- litellm-rust/crates/host/src/protocol.rs | 17 + litellm-rust/crates/host/src/route.rs | 14 - litellm-rust/crates/host/src/run.rs | 147 ++++--- .../crates/llms/src/base_llm/ocr/error.rs | 1 - .../python-bridge/src/logger/machine.rs | 11 +- .../crates/python-bridge/src/logger/tests.rs | 13 +- .../python-bridge/src/routes/messages/host.rs | 33 +- .../python-bridge/src/routes/messages/mod.rs | 4 +- .../python-bridge/src/routes/ocr/document.rs | 247 +++-------- .../python-bridge/src/routes/ocr/host.rs | 100 ++--- .../python-bridge/src/routes/ocr/mod.rs | 4 +- .../python-bridge/src/routes/ocr/project.rs | 110 +++-- 41 files changed, 1660 insertions(+), 1139 deletions(-) create mode 100644 litellm-rust/crates/coroutine/AGENTS.md create mode 100644 litellm-rust/crates/coroutine/Cargo.toml create mode 100644 litellm-rust/crates/coroutine/src/co.rs create mode 100644 litellm-rust/crates/coroutine/src/coroutine.rs create mode 100644 litellm-rust/crates/coroutine/src/error.rs create mode 100644 litellm-rust/crates/coroutine/src/lib.rs create mode 100644 litellm-rust/crates/coroutine/src/reply.rs create mode 100644 litellm-rust/crates/coroutine/tests/coroutine.rs create mode 100644 litellm-rust/crates/host-python/src/file_reader.rs create mode 100644 litellm-rust/crates/host/src/machine/call_machine.rs delete mode 100644 litellm-rust/crates/host/src/machine/route_machine.rs create mode 100644 litellm-rust/crates/host/src/protocol.rs delete mode 100644 litellm-rust/crates/host/src/route.rs diff --git a/litellm-rust/AGENTS.md b/litellm-rust/AGENTS.md index e5ffcd1c57a..70fcc367905 100644 --- a/litellm-rust/AGENTS.md +++ b/litellm-rust/AGENTS.md @@ -8,3 +8,11 @@ - Split a mixed test file along that line instead of widening visibility to move it - A test for another crate's item belongs in that crate, not in a downstream one - Never set `autotests = false` or hand-list `[[test]]` targets; every file directly under `tests/` is discovered by cargo, and a shared helper goes in `tests//mod.rs` or `tests//support.rs` so it is not picked up as a test crate of its own + +## Error definitions + +- A crate's errors live in `src/error.rs`, defined with `thiserror`, and re-exported from `lib.rs` +- Default to one top-level `Error` enum per crate, with one variant per failure mode and a `#[error(...)]` message on each +- Wrap a lower-level error as a variant with `#[from]` or `#[source]` instead of flattening it to a string +- Exception: split into separate types when different functions fail in disjoint ways, especially when different callers see them. A shared enum would force every caller to match variants its function can never return +- Name a split type after what went wrong (a unit struct is fine for a single failure mode), not after the function that returns it diff --git a/litellm-rust/Cargo.lock b/litellm-rust/Cargo.lock index 0a91f0759c2..c522bf205b4 100644 --- a/litellm-rust/Cargo.lock +++ b/litellm-rust/Cargo.lock @@ -3013,6 +3013,15 @@ dependencies = [ "url", ] +[[package]] +name = "litellm-coroutine" +version = "0.1.0" +dependencies = [ + "rstest", + "thiserror 2.0.19", + "tokio", +] + [[package]] name = "litellm-cost" version = "0.1.0" @@ -3040,6 +3049,7 @@ name = "litellm-host" version = "0.1.0" dependencies = [ "litellm-auth", + "litellm-coroutine", "rstest", "serde_json", "tokio", @@ -3049,6 +3059,7 @@ dependencies = [ name = "litellm-host-python" version = "0.1.0" dependencies = [ + "bytes", "futures-util", "litellm-host", "pyo3", diff --git a/litellm-rust/Cargo.toml b/litellm-rust/Cargo.toml index 9d05c8d2b98..0c7236e807e 100644 --- a/litellm-rust/Cargo.toml +++ b/litellm-rust/Cargo.toml @@ -12,6 +12,7 @@ repository = "https://github.com/BerriAI/litellm" litellm-tracing = { path = "crates/tracing" } tracing = "0.1" litellm-core = { path = "crates/core" } +litellm-coroutine = { path = "crates/coroutine" } litellm-host = { path = "crates/host" } litellm-callbacks-legacy-python = { path = "crates/callbacks-legacy-python" } litellm-framing = { path = "crates/framer" } diff --git a/litellm-rust/crates/callbacks-legacy-python/src/call.rs b/litellm-rust/crates/callbacks-legacy-python/src/call.rs index b37790f60a8..9b921070839 100644 --- a/litellm-rust/crates/callbacks-legacy-python/src/call.rs +++ b/litellm-rust/crates/callbacks-legacy-python/src/call.rs @@ -3,8 +3,8 @@ //! lifetime. No other callback host has that obligation, which is why nothing outside //! this crate holds them. -use litellm_host::{machine::Machine, route::Route}; -use litellm_host_python::{RouteHost, lookup, run_call}; +use litellm_host::{machine::Machine, protocol::Protocol}; +use litellm_host_python::{ProtocolHost, lookup, run_call}; use pyo3::{ gc::{PyTraverseError, PyVisit}, prelude::*, @@ -63,25 +63,25 @@ impl PublicCall { } } -/// Runs one native call under the legacy `Logging` contract: the route host projects from +/// Runs one native call under the legacy `Logging` contract: the protocol host projects from /// the keyword view the contract prepares, and the contract observes the call. pub fn run_legacy_call( py: Python<'_>, surface: LegacySurface, call: PublicCall, machine: M, - route: H, + host: H, asynchronous: bool, ) -> PyResult> where - H: RouteHost + 'static, - M: Machine::Response> + 'static, + H: ProtocolHost + 'static, + M: Machine::Response> + 'static, { let arguments = call.kwargs.clone_ref(py); run_call( py, machine, - route, + host, Box::new(LegacyLogging::new(py, surface, call, asynchronous)), arguments, asynchronous, diff --git a/litellm-rust/crates/core/src/messages/route.rs b/litellm-rust/crates/core/src/messages/route.rs index 8cd3eaf3aa3..fc1a9b63252 100644 --- a/litellm-rust/crates/core/src/messages/route.rs +++ b/litellm-rust/crates/core/src/messages/route.rs @@ -1,4 +1,5 @@ use std::{ + convert::Infallible, sync::{Arc, Mutex}, time::Duration, }; @@ -9,8 +10,8 @@ use litellm_core_utils::get_llm_provider_logic::get_custom_llm_provider; use litellm_host::{ event::{MachineEvent, RawResponse, RequestContext, WireRequest}, host::{Demand, Host}, - machine::{HostChannel, MachineFault, RouteMachine}, - route::Route, + machine::{CallMachine, HostChannel, MachineFault}, + protocol::Protocol, }; use litellm_secrets::source::SecretSource; use litellm_types::{ @@ -28,15 +29,6 @@ use super::{ }; use crate::constants::ANTHROPIC_MESSAGES_PROVIDER; -#[derive(Clone, Copy, Debug, PartialEq, Eq)] -pub enum MessagesOp { - ProjectRequest, -} - -pub enum MessagesOpResult { - Request(Box), -} - /// The caller's request as the host projects it. pub struct MessagesCall { pub model: String, @@ -64,11 +56,11 @@ pub enum MessagesOutput { pub struct Messages; -impl Route for Messages { +impl Protocol for Messages { type Response = MessagesOutput; type Error = Error; - type Op = MessagesOp; - type OpResult = MessagesOpResult; + type Projection = MessagesCall; + type Op = Infallible; type Chunk = Bytes; type StreamHead = (); } @@ -78,13 +70,12 @@ impl From for Error { Self::InvalidRequest(match fault { MachineFault::Abandoned => "messages host driver was abandoned".into(), MachineFault::Protocol(message) => format!("messages {message}"), - MachineFault::Mismatch => "invalid messages host operation result".into(), }) } } pub type MessagesHost = HostChannel; -pub type MessagesMachine = RouteMachine; +pub type MessagesMachine = CallMachine; /// Whether this route serves the request, decided before any callback runs so a host /// can still run its own path. @@ -114,30 +105,28 @@ impl LocalMessagesHost { } impl Host for LocalMessagesHost { - async fn route(&self, op: MessagesOp) -> Result { - match op { - MessagesOp::ProjectRequest => self - .call - .lock() - .unwrap_or_else(|error| error.into_inner()) - .take() - .map(|call| MessagesOpResult::Request(Box::new(call))) - .ok_or_else(|| { - Error::InvalidRequest("messages request was already projected".into()) - }), - } + async fn project(&self) -> Result { + self.call + .lock() + .unwrap_or_else(|error| error.into_inner()) + .take() + .ok_or_else(|| Error::InvalidRequest("messages request was already projected".into())) + } + + async fn custom_op(&self, op: Infallible) -> Result<(), Error> { + match op {} } } pub fn messages_machine(secrets: Arc) -> MessagesMachine { - RouteMachine::new(move |host| Box::pin(execute(host, secrets.clone()))) + CallMachine::new(move |host| Box::pin(execute(host, secrets.clone()))) } async fn execute( host: MessagesHost, secrets: Arc, ) -> Result { - let MessagesOpResult::Request(call) = host.route(MessagesOp::ProjectRequest).await?; + let call = host.project().await?; let stream = call.streams(); let resolved = resolve_provider(&call.model, call.custom_llm_provider.as_deref())?; let secrets = secrets.resolve(resolved.config.secret_names()).await?; diff --git a/litellm-rust/crates/core/src/ocr/document.rs b/litellm-rust/crates/core/src/ocr/document.rs index ffa4f045e8e..b78c09298de 100644 --- a/litellm-rust/crates/core/src/ocr/document.rs +++ b/litellm-rust/crates/core/src/ocr/document.rs @@ -23,9 +23,6 @@ pub fn prepare_document(input: OcrDocumentInput) -> Result { file_name.as_deref(), mime_type.as_deref(), )?), - OcrDocumentInput::HostReader { .. } => Err(Error::InvalidRequest( - "OCR file reader was not read by the host".into(), - )), } } @@ -207,7 +204,7 @@ mod tests { } #[test] - fn byte_documents_are_encoded_and_host_readers_must_be_read_first() { + fn byte_documents_are_encoded() { assert_eq!( prepare_document(OcrDocumentInput::Bytes { bytes: b"abc".as_slice().into(), @@ -217,7 +214,6 @@ mod tests { .unwrap(), document("data:application/pdf;base64,YWJj") ); - assert!(prepare_document(OcrDocumentInput::HostReader { mime_type: None }).is_err()); } #[test] diff --git a/litellm-rust/crates/core/src/ocr/route.rs b/litellm-rust/crates/core/src/ocr/route.rs index 7f83291bdab..e3e57bd2d77 100644 --- a/litellm-rust/crates/core/src/ocr/route.rs +++ b/litellm-rust/crates/core/src/ocr/route.rs @@ -3,102 +3,73 @@ use std::sync::{Arc, Mutex}; use litellm_auth::ResolvedCredential; use litellm_host::{ event::{CallEvent, RequestContext, WireRequest}, - machine::{HostChannel, HostTokenProvider, MachineFault, RouteMachine, TokenRoute}, - route::Route, + host::Reply, + machine::{CallMachine, HostChannel, HostTokenProvider, TokenProtocol}, + protocol::Protocol, }; use litellm_llms::base_llm::ocr::{ error::Error, handler::OcrClient, transformation::LiteLLMOcrResponse, }; use super::handler::perform_ocr_request; -use crate::ocr::types::{LiteLLMOcrRequest, OcrDocumentInput, OcrFileContent, ResolvedOcrRequest}; +use crate::ocr::types::{LiteLLMOcrRequest, OcrDocumentInput, ResolvedOcrRequest}; -#[derive(Clone, Copy, Debug, PartialEq, Eq)] pub enum OcrOp { - ProjectRequest, - ReadDocument, - AcquireAzureAdToken, + AcquireAzureAdToken(Reply), } -pub enum OcrOpResult { - Request { - request: Box>, - caller_token: bool, - }, - Document(OcrFileContent), - AzureAdToken(ResolvedCredential), +/// The caller's request as the host projects it. +pub struct OcrProjection { + pub request: LiteLLMOcrRequest, + /// The caller passed its own Azure AD token provider, which the host keeps. + pub caller_token: bool, } pub struct Ocr; -impl Route for Ocr { +impl Protocol for Ocr { type Response = LiteLLMOcrResponse; type Error = Error; + type Projection = OcrProjection; type Op = OcrOp; - type OpResult = OcrOpResult; type Chunk = std::convert::Infallible; type StreamHead = std::convert::Infallible; } -impl TokenRoute for Ocr { - fn acquire_token_op() -> OcrOp { - OcrOp::AcquireAzureAdToken - } - - fn token_credential(result: OcrOpResult) -> Option { - match result { - OcrOpResult::AzureAdToken(credential) => Some(credential), - _ => None, - } +impl TokenProtocol for Ocr { + fn acquire_token_op(reply: Reply) -> OcrOp { + OcrOp::AcquireAzureAdToken(reply) } } pub type OcrHost = HostChannel; -pub type OcrMachine = RouteMachine; +pub type OcrMachine = CallMachine; -/// The OCR call as a machine: projection, document reading and token acquisition are -/// host operations; everything else runs in Rust. +/// The OCR call as a machine: projection and token acquisition are host operations; +/// everything else runs in Rust. pub fn ocr_machine(client: OcrClient) -> OcrMachine { - RouteMachine::new(move |host| Box::pin(execute(client, host))) + CallMachine::new(move |host| Box::pin(execute(client, host))) } async fn execute(client: OcrClient, host: OcrHost) -> Result { - let OcrOpResult::Request { + let OcrProjection { request, caller_token, - } = host.route(OcrOp::ProjectRequest).await? - else { - return Err(MachineFault::Mismatch.into()); - }; + } = host.project().await?; let request = LiteLLMOcrRequest { azure_ad_token_provider: caller_token .then(|| HostTokenProvider::handle(host.clone())) .or(request.azure_ad_token_provider), - ..*request + ..request }; let caller_document = matches!(request.document, OcrDocumentInput::Document(_)); - let request = prepare_request_document(request, &host).await?; + let request = prepare_request_document(request).await?; perform_ocr_request(&client, request, &host, caller_document).await } async fn prepare_request_document( request: LiteLLMOcrRequest, - host: &OcrHost, ) -> Result { - let request = match &request.document { - OcrDocumentInput::HostReader { mime_type } => { - let mime_type = mime_type.clone(); - let OcrOpResult::Document(content) = host.route(OcrOp::ReadDocument).await? else { - return Err(MachineFault::Mismatch.into()); - }; - request.with_document(OcrDocumentInput::Bytes { - bytes: content.bytes, - file_name: content.file_name, - mime_type, - }) - } - _ => request, - }; if let OcrDocumentInput::Document(_) = &request.document { return request.map_document(super::document::prepare_document); } @@ -107,7 +78,6 @@ async fn prepare_request_document( .map_err(|error| Error::DocumentTask(Arc::new(error)))? } -type Reader = Box Result + Send + Sync>; type BeforeSend = Box Result + Send + Sync>; type Observer = Box; @@ -116,7 +86,6 @@ type Observer = Box; /// projection, and the optional observer sees and may rewrite the wire request. pub struct LocalOcrHost { request: Mutex>>, - reader: Option, before_send: Option, observer: Option, } @@ -125,22 +94,11 @@ impl LocalOcrHost { pub fn new(request: LiteLLMOcrRequest) -> Self { Self { request: Mutex::new(Some(request)), - reader: None, before_send: None, observer: None, } } - pub fn with_reader( - self, - reader: impl Fn() -> Result + Send + Sync + 'static, - ) -> Self { - Self { - reader: Some(Box::new(reader)), - ..self - } - } - pub fn with_before_send( self, before_send: impl Fn(WireRequest, &RequestContext) -> Result @@ -163,25 +121,21 @@ impl LocalOcrHost { } impl litellm_host::host::Host for LocalOcrHost { - async fn route(&self, op: OcrOp) -> Result { + async fn project(&self) -> Result { + self.request + .lock() + .unwrap_or_else(|error| error.into_inner()) + .take() + .map(|request| OcrProjection { + request, + caller_token: false, + }) + .ok_or_else(|| Error::InvalidRequest("OCR request was already projected".into())) + } + + async fn custom_op(&self, op: OcrOp) -> Result<(), Error> { match op { - OcrOp::ProjectRequest => self - .request - .lock() - .unwrap_or_else(|error| error.into_inner()) - .take() - .map(|request| OcrOpResult::Request { - request: Box::new(request), - caller_token: false, - }) - .ok_or_else(|| Error::InvalidRequest("OCR request was already projected".into())), - OcrOp::ReadDocument => self - .reader - .as_ref() - .ok_or_else(|| Error::InvalidRequest("OCR host has no document reader".into())) - .and_then(|reader| reader()) - .map(OcrOpResult::Document), - OcrOp::AcquireAzureAdToken => { + OcrOp::AcquireAzureAdToken(_) => { Err(Error::Auth(litellm_auth::Error::AzureTokenAcquisition( "OCR host has no Azure AD token provider".into(), ))) @@ -2757,7 +2711,7 @@ pub(crate) mod tests { use litellm_auth_gcp::VertexAuth; use litellm_host::{ event::{CallEvent, MachineEvent, WireRequest}, - host::{Host, HostOp, HostResult}, + host::{Host, HostOp}, machine::{HostFailure, Machine, MachineStep}, }; use litellm_http::{ @@ -2776,7 +2730,7 @@ pub(crate) mod tests { use rstest::rstest; use serde_json::{Value, json}; - use crate::ocr::route::{LocalOcrHost, OcrOp, OcrOpResult, ocr_machine}; + use crate::ocr::route::{LocalOcrHost, OcrOp, OcrProjection, ocr_machine}; use crate::ocr::{ test_support::{ MockResponse, mock_server, ocr_client, perform_ocr, perform_ocr_with, wire_request, @@ -3212,42 +3166,42 @@ pub(crate) mod tests { crate::ocr::route::OcrMachine, ) { let mut machine = ocr_machine(client); - let mut result = None; let mut ops = Vec::new(); let outcome = loop { - let op = match machine.resume(result.take()).await { + let op = match machine.resume().await { Ok(MachineStep::Host(op)) => op, Ok(MachineStep::Complete(response)) => break Ok(response), Err(error) => break Err(error), }; let answer = match op { - HostOp::Route(op) => { - ops.push(match op { - OcrOp::ProjectRequest => "ProjectRequest", - OcrOp::ReadDocument => "ReadDocument", - OcrOp::AcquireAzureAdToken => "AcquireAzureAdToken", - }); - host.route(op) + HostOp::Project(reply) => { + ops.push("Project"); + host.project() .await - .map(HostResult::Route) + .map(|projection| reply.send(projection)) .map_err(HostFailure::Error) } - HostOp::BeforeSend { wire, .. } => { - ops.push("BeforeSend"); - intercept(*wire).map(|wire| HostResult::BeforeSend(Box::new(wire))) + HostOp::Custom(op) => { + ops.push(match op { + OcrOp::AcquireAzureAdToken(_) => "AcquireAzureAdToken", + }); + host.custom_op(op).await.map_err(HostFailure::Error) } - HostOp::Emit(event) => { + HostOp::BeforeSend { wire, reply, .. } => { + ops.push("BeforeSend"); + intercept(*wire).map(|wire| reply.send(wire)) + } + HostOp::Emit(event, reply) => { let event = CallEvent::Machine(event); ops.push(event_name(&event)); host.emit(&event) .await - .map(|()| HostResult::Emitted) + .map(|()| reply.send(())) .map_err(HostFailure::Error) } }; - match answer { - Ok(answer) => result = Some(answer), - Err(failure) => break machine.interrupt(failure).await, + if let Err(failure) = answer { + break machine.interrupt(failure).await; } }; (outcome, ops, machine) @@ -3269,8 +3223,8 @@ pub(crate) mod tests { assert!( matches!(outcome, Err(OcrError::InvalidRequest(message)) if message == "before_send failed") ); - assert_eq!(ops, ["ProjectRequest", "BeforeSend"]); - assert!(machine.resume(None).await.is_err()); + assert_eq!(ops, ["Project", "BeforeSend"]); + assert!(machine.resume().await.is_err()); } #[tokio::test] @@ -3306,80 +3260,24 @@ pub(crate) mod tests { server.await.unwrap(); assert_eq!(outcome.unwrap().pages[0].markdown, "native"); assert_eq!(seen.lock().unwrap().len(), 1); - assert_eq!(ops, ["ProjectRequest", "BeforeSend", "response"]); + assert_eq!(ops, ["Project", "BeforeSend", "response"]); assert!(matches!( - machine.resume(None).await, + machine.resume().await, Err(OcrError::InvalidRequest(_)) )); } - async fn drive_native_file_call( - request: crate::ocr::types::LiteLLMOcrRequest, - content: Result, - ) -> (Result, usize) { - let reads = Arc::new(Mutex::new(0)); - let counted = reads.clone(); - let content = Mutex::new(Some(content)); - let host = LocalOcrHost::new(request).with_reader(move || { - *counted.lock().unwrap() += 1; - content.lock().unwrap().take().unwrap() - }); - let outcome = perform_ocr_with(host).await; - let reads = *reads.lock().unwrap(); - (outcome, reads) - } - #[tokio::test] - async fn host_reader_documents_are_read_once_at_the_core_selected_point_and_encoded() { - let (base, seen, server) = mock_server(vec![MockResponse::json(json!({ - "pages":[{"index":0,"markdown":"file"}] - }))]) - .await; - let request = wire_request("mistral/model", &base, json!({})).with_document( - crate::ocr::types::OcrDocumentInput::HostReader { - mime_type: Some("application/pdf".into()), - }, - ); - let (response, reads) = drive_native_file_call( - request, - Ok(crate::ocr::types::OcrFileContent { - bytes: b"abc".as_slice().into(), - file_name: Some("scan.png".into()), - }), - ) - .await; - server.await.unwrap(); - assert_eq!(response.unwrap().pages[0].markdown, "file"); - assert_eq!(reads, 1); - assert!(seen.lock().unwrap()[0].contains("data:application/pdf;base64,YWJj")); - } - - #[tokio::test] - async fn host_reader_failures_and_empty_files_fail_before_the_provider_is_called() { + async fn empty_byte_documents_fail_before_the_provider_is_called() { let (base, seen, _server) = mock_server(vec![]).await; - let request = wire_request("mistral/model", &base, json!({})); - let failure = OcrError::InvalidRequest("reader exploded".into()); - let (response, reads) = drive_native_file_call( - request - .with_document(crate::ocr::types::OcrDocumentInput::HostReader { mime_type: None }), - Err(failure.clone()), - ) - .await; - assert!( - matches!(response.unwrap_err(), OcrError::InvalidRequest(message) if message == "reader exploded") - ); - assert_eq!(reads, 1); - - let request = wire_request("mistral/model", &base, json!({})); - let (response, _) = drive_native_file_call( - request - .with_document(crate::ocr::types::OcrDocumentInput::HostReader { mime_type: None }), - Ok(crate::ocr::types::OcrFileContent { + let request = wire_request("mistral/model", &base, json!({})).with_document( + crate::ocr::types::OcrDocumentInput::Bytes { bytes: Default::default(), file_name: None, - }), - ) - .await; + mime_type: None, + }, + ); + let response = perform_ocr_with(LocalOcrHost::new(request)).await; assert!(matches!(response.unwrap_err(), OcrError::EmptyFile)); assert!(seen.lock().unwrap().is_empty()); } @@ -3400,24 +3298,21 @@ pub(crate) mod tests { mime_type: None, }, ); - let (response, reads) = - drive_native_file_call(request, Err(OcrError::InvalidRequest("unused".into()))).await; + let (response, ops, _) = drive_until(ocr_client(), &LocalOcrHost::new(request), Ok).await; server.await.unwrap(); std::fs::remove_dir_all(&dir).unwrap(); assert_eq!(response.unwrap().pages[0].markdown, "path"); - assert_eq!(reads, 0); + assert_eq!(ops, ["Project", "BeforeSend", "response"]); assert!(seen.lock().unwrap()[0].contains("data:image/png;base64,YWJj")); let (base, seen, _server) = mock_server(vec![]).await; - let request = wire_request("mistral/model", &base, json!({})); - let (response, _) = drive_native_file_call( - request.with_document(crate::ocr::types::OcrDocumentInput::Path { + let request = wire_request("mistral/model", &base, json!({})).with_document( + crate::ocr::types::OcrDocumentInput::Path { path: path.clone(), mime_type: None, - }), - Err(OcrError::InvalidRequest("unused".into())), - ) - .await; + }, + ); + let response = perform_ocr_with(LocalOcrHost::new(request)).await; assert!(matches!( response.unwrap_err(), OcrError::FileRead { path: failed, source } if failed == path && source.kind() == std::io::ErrorKind::NotFound @@ -3441,28 +3336,25 @@ pub(crate) mod tests { assert!( matches!(outcome, Err(OcrError::InvalidRequest(message)) if message == "cancelled") ); - assert_eq!(ops, ["ProjectRequest", "BeforeSend"]); - assert!(machine.resume(Some(HostResult::Emitted)).await.is_err()); + assert_eq!(ops, ["Project", "BeforeSend"]); + assert!(machine.resume().await.is_err()); } #[tokio::test] - async fn missing_host_result_preserves_pending_operation() { + async fn resuming_before_answering_preserves_pending_operation() { let request = wire_request("mistral/model", "http://127.0.0.1:1", json!({})); let mut machine = ocr_machine(ocr_client()); + let Ok(MachineStep::Host(HostOp::Project(reply))) = machine.resume().await else { + panic!("expected the projection op first"); + }; + assert!(machine.resume().await.is_err()); + reply.send(OcrProjection { + request, + caller_token: false, + }); assert!(matches!( - machine.resume(None).await.unwrap(), - MachineStep::Host(HostOp::Route(OcrOp::ProjectRequest)) - )); - assert!(machine.resume(None).await.is_err()); - assert!(matches!( - machine - .resume(Some(HostResult::Route(OcrOpResult::Request { - request: Box::new(request), - caller_token: false, - }))) - .await - .unwrap(), - MachineStep::Host(HostOp::BeforeSend { .. }) + machine.resume().await, + Ok(MachineStep::Host(HostOp::BeforeSend { .. })) )); } @@ -3623,20 +3515,18 @@ pub(crate) mod tests { }; let host = LocalOcrHost::new(request); let mut machine = ocr_machine(ocr_client()); - let mut result = None; tokio::time::timeout(std::time::Duration::from_secs(2), async { loop { tokio::select! { _ = entered.notified() => break, - step = machine.resume(result.take()) => { - result = Some(match step.unwrap() { - MachineStep::Host(HostOp::Route(op)) => HostResult::Route(host.route(op).await.unwrap()), - MachineStep::Host(HostOp::BeforeSend { wire, .. }) => { - HostResult::BeforeSend(wire) - } - MachineStep::Host(HostOp::Emit(_)) => HostResult::Emitted, + step = machine.resume() => { + match step.unwrap() { + MachineStep::Host(HostOp::Project(reply)) => reply.send(host.project().await.unwrap()), + MachineStep::Host(HostOp::Custom(op)) => host.custom_op(op).await.unwrap(), + MachineStep::Host(HostOp::BeforeSend { wire, reply, .. }) => reply.send(*wire), + MachineStep::Host(HostOp::Emit(_, reply)) => reply.send(()), MachineStep::Complete(_) => panic!("pending provider completed"), - }); + } } } } @@ -3661,24 +3551,23 @@ pub(crate) mod tests { } impl Host for CallerTokenHost { - async fn route(&self, op: OcrOp) -> Result { + async fn project(&self) -> Result { + self.trace.lock().unwrap().push("project".into()); + Ok(OcrProjection { + request: self.request.lock().unwrap().take().unwrap(), + caller_token: true, + }) + } + + async fn custom_op(&self, op: OcrOp) -> Result<(), OcrError> { match op { - OcrOp::ProjectRequest => { - self.trace.lock().unwrap().push("project".into()); - Ok(OcrOpResult::Request { - request: Box::new(self.request.lock().unwrap().take().unwrap()), - caller_token: true, - }) - } - OcrOp::AcquireAzureAdToken => { + OcrOp::AcquireAzureAdToken(reply) => { self.trace.lock().unwrap().push("token".into()); - Ok(OcrOpResult::AzureAdToken( - litellm_auth::ResolvedCredential::Static(litellm_auth::SecretValue::new( - "caller-token", - )), - )) + reply.send(litellm_auth::ResolvedCredential::Static( + litellm_auth::SecretValue::new("caller-token"), + )); + Ok(()) } - OcrOp::ReadDocument => Err(OcrError::InvalidRequest("no reader".into())), } } @@ -3761,18 +3650,18 @@ pub(crate) mod tests { }); let host = LocalOcrHost::new(wire_request("mistral/model", &base, json!({}))); let mut machine = ocr_machine(ocr_client()); - let mut result = None; tokio::time::timeout(std::time::Duration::from_secs(2), async { loop { tokio::select! { _ = received.notified() => break, - step = machine.resume(result.take()) => { - result = Some(match step.unwrap() { - MachineStep::Host(HostOp::Route(op)) => HostResult::Route(host.route(op).await.unwrap()), - MachineStep::Host(HostOp::BeforeSend { wire, .. }) => HostResult::BeforeSend(wire), - MachineStep::Host(HostOp::Emit(_)) => HostResult::Emitted, + step = machine.resume() => { + match step.unwrap() { + MachineStep::Host(HostOp::Project(reply)) => reply.send(host.project().await.unwrap()), + MachineStep::Host(HostOp::Custom(op)) => host.custom_op(op).await.unwrap(), + MachineStep::Host(HostOp::BeforeSend { wire, reply, .. }) => reply.send(*wire), + MachineStep::Host(HostOp::Emit(_, reply)) => reply.send(()), MachineStep::Complete(_) => panic!("the stalled provider completed"), - }); + } } } } diff --git a/litellm-rust/crates/core/src/ocr/types.rs b/litellm-rust/crates/core/src/ocr/types.rs index 59c9cec8da9..20a21e43676 100644 --- a/litellm-rust/crates/core/src/ocr/types.rs +++ b/litellm-rust/crates/core/src/ocr/types.rs @@ -25,9 +25,6 @@ pub enum OcrDocumentInput { file_name: Option, mime_type: Option, }, - HostReader { - mime_type: Option, - }, } impl From for OcrDocumentInput { @@ -45,12 +42,6 @@ impl From for OcrDocumentInput { } } -#[derive(Clone, Debug, PartialEq, Eq)] -pub struct OcrFileContent { - pub bytes: Bytes, - pub file_name: Option, -} - /// Caller-supplied connection overrides for a [`LiteLLMOcrRequest`], in the /// shape hosts receive them: JSON-ish headers, optional timeout, optional /// credentials, and per-field provenance in `input_sources`. diff --git a/litellm-rust/crates/coroutine/AGENTS.md b/litellm-rust/crates/coroutine/AGENTS.md new file mode 100644 index 00000000000..fcb4f4df47f --- /dev/null +++ b/litellm-rust/crates/coroutine/AGENTS.md @@ -0,0 +1,31 @@ +# Requirements + +Core must pause mid-call to ask the host for things it cannot do itself (Python callbacks, secret and token reads, `before_send` rewrites, stream demand), then continue where it stopped. Any change to this crate must keep every requirement below; the alternatives section says which one each rejected design breaks + +- R1 Core never calls the host: it names an op and waits for the answer, so it stays free of PyO3 and of any other host runtime +- R2 Async host work is awaited by the host's own driver in the caller's asyncio task (`litellm/rust_bridge/lifecycle.py`), so `contextvars` writes reach the caller; a Rust-side `into_future` would run it in a copied context +- R3 The body awaits real I/O (HTTP, `spawn_blocking`, timers) between yields, so `resume` is itself a future driven by the caller's runtime +- R4 Route code stays straight-line async (`host.route(OcrOp::ReadDocument).await?`) instead of hand-written states +- R5 Each op fixes its answer type at compile time: a host cannot answer `ReadDocument` with a token, and core never matches a result variant it did not ask for +- R6 A yield the body makes while being resumed is returned by that same poll, so the host driver's inline first poll needs no extra event-loop turn per op +- R7 No task is spawned: `cancel`, or dropping the coroutine, drops the body, and nothing waits forever on an answer that cannot come +- R8 Several yields can be pending at once, since route code hands clones of its `Co` to token providers and hooks +- R9 Stable Rust + +# Other implementations and why they do not fit + +- Nightly `std::ops::Coroutine`: breaks R9, and its body cannot await futures between yields (R3) +- `genawaiter`: resumes async bodies only with a noop waker, so the body cannot await real I/O (R3) +- `simple_coro`: typestate `Coro` makes answering before resuming a compile-time rule, but its body cannot await arbitrary futures (R3) and its reply type `R` is fixed per coroutine (R5) +- `corosensei` and other stackful coroutines: sync bodies on their own stack, no async I/O inside (R3) +- A hand-written phase enum with an `advance` match (the old `HostPhase`): every await point becomes a state (R4) +- An injected host trait with `async fn`s: core would call the host itself (R1, R2) +- Sans-IO, where core does no I/O and HTTP becomes one more host op: keeps every requirement and makes `resume` a pure step function, but HTTP, streaming, retries and timeouts would move out of core into every bridge; the one real alternative, not taken +- Temporal's Rust workflow SDK (`WorkflowFuture`, `WfContext`) is the closest precedent: an `async fn` polled in place, commands sent over a channel with a oneshot to unblock them. Roles are inverted there (the language SDK owns the program, core answers), and its workflow body may not do real I/O + +# Tradeoffs accepted + +- A tokio `mpsc` channel plus a `oneshot` per yield instead of compiler-generated states +- Protocol mistakes (resuming before answering, resuming after the end) are runtime `ResumeError`s, not compile errors +- Pending yields come out one per `resume`, in the order they were made, and each reply goes back to the yield that made it (R8) +- An answer sent after its yield stopped waiting (for example the body timed out on it) is discarded, since the body already moved on diff --git a/litellm-rust/crates/coroutine/Cargo.toml b/litellm-rust/crates/coroutine/Cargo.toml new file mode 100644 index 00000000000..3ff79ac5f2c --- /dev/null +++ b/litellm-rust/crates/coroutine/Cargo.toml @@ -0,0 +1,15 @@ +[package] +name = "litellm-coroutine" +version = "0.1.0" +edition.workspace = true +license.workspace = true +repository.workspace = true +description = "Async coroutines on stable Rust whose every yield carries its own typed reply" + +[dependencies] +thiserror.workspace = true +tokio = { workspace = true, features = ["sync"] } + +[dev-dependencies] +rstest.workspace = true +tokio = { workspace = true, features = ["rt", "macros", "time"] } diff --git a/litellm-rust/crates/coroutine/src/co.rs b/litellm-rust/crates/coroutine/src/co.rs new file mode 100644 index 00000000000..d84b041931c --- /dev/null +++ b/litellm-rust/crates/coroutine/src/co.rs @@ -0,0 +1,42 @@ +use std::sync::Weak; + +use tokio::sync::mpsc; + +use crate::{Abandoned, Reply, reply}; + +pub(crate) struct Request { + pub(crate) value: Y, + pub(crate) outstanding: Weak<()>, +} + +/// The body's handle for yielding, `genawaiter`'s `Co`. +pub struct Co { + yields: mpsc::UnboundedSender>, +} + +impl Clone for Co { + fn clone(&self) -> Self { + Self { + yields: self.yields.clone(), + } + } +} + +impl Co { + pub(crate) fn new(yields: mpsc::UnboundedSender>) -> Self { + Self { yields } + } + + /// Yields the value `ask` builds around a fresh [`Reply`] and waits for its answer. + pub async fn yield_(&self, ask: impl FnOnce(Reply) -> Y) -> Result { + let (reply, answer) = reply(); + let outstanding = reply.outstanding(); + self.yields + .send(Request { + value: ask(reply), + outstanding, + }) + .map_err(|_| Abandoned)?; + answer.await + } +} diff --git a/litellm-rust/crates/coroutine/src/coroutine.rs b/litellm-rust/crates/coroutine/src/coroutine.rs new file mode 100644 index 00000000000..fa34f816a37 --- /dev/null +++ b/litellm-rust/crates/coroutine/src/coroutine.rs @@ -0,0 +1,94 @@ +use std::{ + future::{Future, poll_fn}, + pin::Pin, + sync::Weak, + task::{Context, Poll}, +}; + +use tokio::sync::mpsc; + +use crate::{Co, ResumeError, co::Request}; + +/// What one `resume` produced, as in [`std::ops::CoroutineState`]. +#[derive(Debug, PartialEq, Eq)] +pub enum CoroutineState { + Yielded(Y), + Complete(C), +} + +type Body = Pin + Send>>; + +enum Step { + Yielded(Request), + Complete(C), +} + +fn queued( + yields: &mut mpsc::UnboundedReceiver>, + context: &mut Context<'_>, +) -> Option> { + match yields.poll_recv(context) { + Poll::Ready(request) => request, + Poll::Pending => None, + } +} + +pub struct Coroutine { + body: Option>, + yields: mpsc::UnboundedReceiver>, + outstanding: Weak<()>, +} + +impl Coroutine { + /// Builds the body from `producer`. Nothing runs until the first `resume`. + pub fn new(producer: impl FnOnce(Co) -> F) -> Self + where + F: Future + Send + 'static, + { + let (sender, yields) = mpsc::unbounded_channel(); + Self { + body: Some(Box::pin(producer(Co::new(sender)))), + yields, + outstanding: Weak::new(), + } + } + + pub async fn resume(&mut self) -> Result, ResumeError> { + let Some(body) = self.body.as_mut() else { + return Err(ResumeError::Finished); + }; + if self.outstanding.strong_count() > 0 { + return Err(ResumeError::Unanswered); + } + let yields = &mut self.yields; + let step = poll_fn(|context| { + if let Some(request) = queued(yields, context) { + return Poll::Ready(Step::Yielded(request)); + } + if let Poll::Ready(output) = body.as_mut().poll(context) { + return Poll::Ready(Step::Complete(output)); + } + queued(yields, context) + .map_or(Poll::Pending, |request| Poll::Ready(Step::Yielded(request))) + }) + .await; + match step { + Step::Yielded(Request { value, outstanding }) => { + self.outstanding = outstanding; + Ok(CoroutineState::Yielded(value)) + } + Step::Complete(output) => { + self.cancel(); + Ok(CoroutineState::Complete(output)) + } + } + } + + /// Drops the body and fails every yield still waiting, or yet to be made, with + /// [`Abandoned`](crate::Abandoned). + pub fn cancel(&mut self) { + self.body = None; + self.yields.close(); + while self.yields.try_recv().is_ok() {} + } +} diff --git a/litellm-rust/crates/coroutine/src/error.rs b/litellm-rust/crates/coroutine/src/error.rs new file mode 100644 index 00000000000..b8fded23fdb --- /dev/null +++ b/litellm-rust/crates/coroutine/src/error.rs @@ -0,0 +1,14 @@ +/// A `resume` the coroutine refused, leaving it as it was. +#[derive(Clone, Copy, Debug, PartialEq, Eq, thiserror::Error)] +pub enum ResumeError { + #[error("coroutine resumed after it finished")] + Finished, + #[error("coroutine resumed before the reply to its last yield was sent or dropped")] + Unanswered, +} + +/// No answer will come to a yield: its [`Reply`](crate::Reply) was dropped unsent, or the +/// coroutine it was sent to is gone. +#[derive(Clone, Copy, Debug, PartialEq, Eq, thiserror::Error)] +#[error("the yield was abandoned before it was answered")] +pub struct Abandoned; diff --git a/litellm-rust/crates/coroutine/src/lib.rs b/litellm-rust/crates/coroutine/src/lib.rs new file mode 100644 index 00000000000..636aaf1b2b6 --- /dev/null +++ b/litellm-rust/crates/coroutine/src/lib.rs @@ -0,0 +1,12 @@ +//! Async coroutines on stable Rust whose every yield carries its own typed [`Reply`]. +//! See `AGENTS.md` for the requirement, the alternatives and the contracts. + +mod co; +mod coroutine; +mod error; +mod reply; + +pub use co::Co; +pub use coroutine::{Coroutine, CoroutineState}; +pub use error::{Abandoned, ResumeError}; +pub use reply::{Answer, Reply, reply}; diff --git a/litellm-rust/crates/coroutine/src/reply.rs b/litellm-rust/crates/coroutine/src/reply.rs new file mode 100644 index 00000000000..b3cb7da2e9d --- /dev/null +++ b/litellm-rust/crates/coroutine/src/reply.rs @@ -0,0 +1,60 @@ +use std::{ + fmt, + future::Future, + pin::Pin, + sync::{Arc, Weak}, + task::{Context, Poll}, +}; + +use tokio::sync::oneshot; + +use crate::Abandoned; + +/// The one way to answer a yield. Sending or dropping it settles the yield. +pub struct Reply { + slot: oneshot::Sender, + outstanding: Arc<()>, +} + +impl Reply { + /// An answer the yield no longer awaits is discarded. + pub fn send(self, answer: A) { + let _ = self.slot.send(answer); + } + + /// Alive until this reply is sent or dropped. + pub(crate) fn outstanding(&self) -> Weak<()> { + Arc::downgrade(&self.outstanding) + } +} + +impl fmt::Debug for Reply { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter.write_str("Reply") + } +} + +/// The waiting end of a [`Reply`]. +pub struct Answer { + slot: oneshot::Receiver, +} + +impl Future for Answer { + type Output = Result; + + fn poll(mut self: Pin<&mut Self>, context: &mut Context<'_>) -> Poll { + Pin::new(&mut self.slot) + .poll(context) + .map(|answer| answer.map_err(|_| Abandoned)) + } +} + +/// A reply outside any coroutine, for answering a host operation directly. +pub fn reply() -> (Reply, Answer) { + let (slot, answer) = oneshot::channel(); + let reply = Reply { + slot, + outstanding: Arc::new(()), + }; + (reply, Answer { slot: answer }) +} diff --git a/litellm-rust/crates/coroutine/tests/coroutine.rs b/litellm-rust/crates/coroutine/tests/coroutine.rs new file mode 100644 index 00000000000..91d195503df --- /dev/null +++ b/litellm-rust/crates/coroutine/tests/coroutine.rs @@ -0,0 +1,256 @@ +use std::{ + future::Future, + sync::{Arc, Mutex}, + time::Duration, +}; + +use litellm_coroutine::{Abandoned, Co, Coroutine, CoroutineState, Reply, ResumeError, reply}; +use rstest::rstest; +use tokio::time::timeout; + +#[derive(Debug)] +enum Ask { + Name(Reply<&'static str>), + Count(Reply), +} + +type Test = Coroutine; + +fn yielded(state: Result, ResumeError>) -> Ask { + match state { + Ok(CoroutineState::Yielded(ask)) => ask, + Ok(CoroutineState::Complete(_)) => panic!("expected a yield, the body returned"), + Err(error) => panic!("expected a yield, resume failed: {error}"), + } +} + +fn complete(state: Result, ResumeError>) -> C { + match state { + Ok(CoroutineState::Complete(output)) => output, + Ok(CoroutineState::Yielded(ask)) => panic!("expected completion, got {ask:?}"), + Err(error) => panic!("expected completion, resume failed: {error}"), + } +} + +fn name(ask: Ask) -> Reply<&'static str> { + match ask { + Ask::Name(reply) => reply, + other => panic!("expected a name ask, got {other:?}"), + } +} + +fn count(ask: Ask) -> Reply { + match ask { + Ask::Count(reply) => reply, + other => panic!("expected a count ask, got {other:?}"), + } +} + +/// A body parked at one name ask, with nothing else going on. +fn suspended_once() -> Test> { + Coroutine::new(|co| async move { co.yield_(Ask::Name).await }) +} + +#[tokio::test] +async fn each_typed_answer_resumes_the_yield_that_asked_for_it() { + let mut coroutine: Test = Coroutine::new(|co| async move { + let first = co.yield_(Ask::Name).await.unwrap(); + let second = co.yield_(Ask::Count).await.unwrap(); + format!("{first}+{second}") + }); + + name(yielded(coroutine.resume().await)).send("a"); + count(yielded(coroutine.resume().await)).send(2); + + assert_eq!(complete(coroutine.resume().await), "a+2"); +} + +/// A driver that polls `resume` once, inline, sees every yield the body makes during +/// that poll instead of being sent back to its event loop. +#[test] +fn a_yield_made_while_resuming_is_returned_by_that_same_poll() { + let mut coroutine: Test = Coroutine::new(|co| async move { + let first = co.yield_(Ask::Count).await.unwrap(); + let second = co.yield_(Ask::Count).await.unwrap(); + first + second + }); + let mut context = std::task::Context::from_waker(std::task::Waker::noop()); + let mut poll_once = + |coroutine: &mut Test| match std::pin::pin!(coroutine.resume()).poll(&mut context) { + std::task::Poll::Ready(state) => state, + std::task::Poll::Pending => panic!("resume needed a second poll"), + }; + + count(yielded(poll_once(&mut coroutine))).send(1); + count(yielded(poll_once(&mut coroutine))).send(2); + + assert_eq!(complete(poll_once(&mut coroutine)), 3); +} + +#[tokio::test] +async fn the_body_awaits_real_futures_between_yields() { + let mut coroutine: Test = Coroutine::new(|co| async move { + tokio::time::sleep(Duration::from_millis(5)).await; + co.yield_(Ask::Count).await.unwrap() + }); + + count(yielded(coroutine.resume().await)).send(7); + + assert_eq!(complete(coroutine.resume().await), 7); +} + +#[tokio::test] +async fn concurrent_yields_come_out_in_order_and_are_answered_separately() { + let mut coroutine: Test<(&str, u32)> = Coroutine::new(|co| async move { + let (first, second) = tokio::join!(co.yield_(Ask::Name), co.yield_(Ask::Count)); + (first.unwrap(), second.unwrap()) + }); + + name(yielded(coroutine.resume().await)).send("one"); + count(yielded(coroutine.resume().await)).send(2); + + assert_eq!(complete(coroutine.resume().await), ("one", 2)); +} + +#[tokio::test] +async fn resuming_before_the_reply_is_settled_is_refused_and_keeps_the_yield_waiting() { + let mut coroutine = suspended_once(); + let reply = name(yielded(coroutine.resume().await)); + + assert_eq!( + coroutine.resume().await.unwrap_err(), + ResumeError::Unanswered + ); + + reply.send("real"); + assert_eq!(complete(coroutine.resume().await), Ok("real")); +} + +#[tokio::test] +async fn a_dropped_reply_abandons_its_yield() { + let mut coroutine = suspended_once(); + drop(yielded(coroutine.resume().await)); + + assert_eq!(complete(coroutine.resume().await), Err(Abandoned)); +} + +#[tokio::test] +async fn an_answer_the_yield_no_longer_awaits_is_discarded() { + let mut coroutine: Test<&str> = Coroutine::new(|co| async move { + tokio::select! { + biased; + _ = co.yield_(Ask::Name) => unreachable!("the answer comes after the body moved on"), + () = std::future::ready(()) => {} + } + co.yield_(Ask::Name).await.unwrap() + }); + let stale = name(yielded(coroutine.resume().await)); + stale.send("stale"); + + name(yielded(coroutine.resume().await)).send("fresh"); + + assert_eq!(complete(coroutine.resume().await), "fresh"); +} + +#[rstest] +#[case::returned(false)] +#[case::cancelled(true)] +#[tokio::test] +async fn a_finished_coroutine_refuses_to_resume(#[case] cancel: bool) { + let mut coroutine = suspended_once(); + let reply = name(yielded(coroutine.resume().await)); + if cancel { + coroutine.cancel(); + } else { + reply.send("done"); + complete(coroutine.resume().await).unwrap(); + } + + assert_eq!(coroutine.resume().await.unwrap_err(), ResumeError::Finished); +} + +#[tokio::test] +async fn a_dropped_resume_leaves_the_coroutine_resumable() { + let mut coroutine: Test = Coroutine::new(|co| async move { + tokio::time::sleep(Duration::from_millis(20)).await; + co.yield_(Ask::Count).await.unwrap() + }); + assert!( + timeout(Duration::from_millis(1), coroutine.resume()) + .await + .is_err() + ); + + count(yielded(coroutine.resume().await)).send(3); + + assert_eq!(complete(coroutine.resume().await), 3); +} + +struct Dropped(Arc>); + +impl Drop for Dropped { + fn drop(&mut self) { + *self.0.lock().unwrap() = true; + } +} + +#[tokio::test] +async fn cancel_drops_the_body() { + let dropped = Arc::new(Mutex::new(false)); + let guard = Dropped(Arc::clone(&dropped)); + let mut coroutine: Test<()> = Coroutine::new(|co| async move { + let _guard = guard; + co.yield_(Ask::Count).await.unwrap(); + }); + let _reply = yielded(coroutine.resume().await); + + coroutine.cancel(); + + assert!(*dropped.lock().unwrap()); +} + +#[rstest] +#[case::cancelled(true)] +#[case::dropped(false)] +#[tokio::test] +async fn a_co_that_escaped_the_body_is_abandoned_once_the_coroutine_ends(#[case] cancel: bool) { + let escaped: Arc>>> = Arc::default(); + let slot = Arc::clone(&escaped); + let mut coroutine: Test<()> = Coroutine::new(move |co| { + *slot.lock().unwrap() = Some(co.clone()); + async move { + co.yield_(Ask::Count).await.unwrap(); + } + }); + let _reply = yielded(coroutine.resume().await); + let co = escaped.lock().unwrap().take().unwrap(); + let waiting = tokio::spawn(async move { co.yield_(Ask::Name).await }); + tokio::task::yield_now().await; + + if cancel { + coroutine.cancel(); + } else { + drop(coroutine); + } + + let outcome = timeout(Duration::from_secs(1), waiting) + .await + .expect("an escaped yield waits forever") + .unwrap(); + assert_eq!(outcome, Err(Abandoned)); +} + +#[rstest] +#[case::sent(true)] +#[case::dropped(false)] +#[tokio::test] +async fn a_detached_reply_settles_its_answer(#[case] send: bool) { + let (reply, answer) = reply::(); + if send { + reply.send(5); + } else { + drop(reply); + } + + assert_eq!(answer.await, if send { Ok(5) } else { Err(Abandoned) }); +} diff --git a/litellm-rust/crates/host-python/AGENTS.md b/litellm-rust/crates/host-python/AGENTS.md index 5aca13eeb18..7c1919f9f39 100644 --- a/litellm-rust/crates/host-python/AGENTS.md +++ b/litellm-rust/crates/host-python/AGENTS.md @@ -1,8 +1,8 @@ - Target invariants; implementation and runtime validation may lag these rules -- Keep this crate the CPython runtime adapter and nothing more: Serde marshalling, interpreter detachment, tokio/asyncio glue, the `Execution` handle, the call driver and the `PythonLifecycle`/`RouteHost` traits +- Keep this crate the CPython runtime adapter and nothing more: Serde marshalling, interpreter detachment, tokio/asyncio glue, the `Execution` handle, the call driver and the `PythonLifecycle`/`ProtocolHost` traits - No LiteLLM domain dependencies beyond `litellm-host`: no route types, no `Logging` policy, no public API registration, no cdylib build features - The driver emits `Succeeded` or `Failed` exactly once and never dispatches after a cancellation; which Python objects consume those events is the adapter's business - - `RouteHost::invoke` receives the keyword view the adapter's `begin` returned, not the caller's dict; a route host that projects from it inherits that adapter's rewrites (for the legacy adapter: setup, deployment hooks, credential inheritance) + - `ProtocolHost::project` receives the keyword view the adapter's `begin` returned, not the caller's dict; a protocol host that projects from it inherits that adapter's rewrites (for the legacy adapter: setup, deployment hooks, credential inheritance) - A native failure, including one a host op returns as `InvokeError::Native`, is classified exactly once through the route's `classify`; a Python exception raised inside the call, and a failure in `begin` or `after_success`, is raised as is - A failing `classify` is raised with the native error's text as its `__context__`, never swallowed - Use standard PyO3 ownership and conversion APIs diff --git a/litellm-rust/crates/host-python/Cargo.toml b/litellm-rust/crates/host-python/Cargo.toml index fb6379dc35a..c1b35c0f69d 100644 --- a/litellm-rust/crates/host-python/Cargo.toml +++ b/litellm-rust/crates/host-python/Cargo.toml @@ -6,6 +6,7 @@ license.workspace = true repository.workspace = true [dependencies] +bytes.workspace = true futures-util.workspace = true litellm-host.workspace = true pyo3.workspace = true diff --git a/litellm-rust/crates/host-python/src/adapter.rs b/litellm-rust/crates/host-python/src/adapter.rs index 3a4cb49be4d..87481aa89b7 100644 --- a/litellm-rust/crates/host-python/src/adapter.rs +++ b/litellm-rust/crates/host-python/src/adapter.rs @@ -1,5 +1,5 @@ use litellm_host::event::{FailureOrigin, MachineEvent, RequestContext, Timing, WireRequest}; -use litellm_host::route::Route; +use litellm_host::protocol::Protocol; use pyo3::exceptions::PyRuntimeError; use pyo3::gc::{PyTraverseError, PyVisit}; use pyo3::prelude::*; @@ -85,7 +85,7 @@ pub trait PythonLifecycle: Send + Sync { fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError>; } -/// Why a route operation the host answered did not produce a result: the route's own code +/// Why a custom operation the host answered did not produce a result: the route's own code /// rejected it, which the route classifies like any other native failure, or Python code /// raised, which reaches the caller as it was raised. #[derive(Debug)] @@ -100,45 +100,54 @@ impl From for InvokeError { } } -/// The Python side of one route: answers the route's own operations, builds the public +/// The Python side of one protocol: answers its custom operations, builds the public /// response and classifies native failures into public exceptions. -pub trait RouteHost: Send + Sync { - type Route: Route; +pub trait ProtocolHost: Send + Sync { + type Protocol: Protocol; /// The public exception a native failure maps to, kept as a value until the driver /// raises it. type Failure: Into; - /// `arguments` is the keyword view the lifecycle's `begin` produced, not the - /// caller's own dict. A route host that projects from it inherits whatever that - /// adapter rewrote. - fn invoke( + /// Projects the call's request. `arguments` is the keyword view the lifecycle's + /// `begin` produced, not the caller's own dict, so the projection inherits whatever + /// that adapter rewrote. + fn project( &mut self, py: Python<'_>, arguments: &Bound<'_, PyDict>, - op: ::Op, - ) -> Result<::OpResult, InvokeError<::Error>>; + ) -> Result< + ::Projection, + InvokeError<::Error>, + >; + + /// Answers `op` through its reply. + fn invoke( + &mut self, + py: Python<'_>, + op: ::Op, + ) -> Result<(), InvokeError<::Error>>; fn complete( &mut self, py: Python<'_>, - response: ::Response, + response: ::Response, ) -> PyResult>; /// One streamed chunk as the caller receives it. fn chunk( &mut self, py: Python<'_>, - chunk: ::Chunk, + chunk: ::Chunk, ) -> PyResult>; fn classify( &self, py: Python<'_>, - error: ::Error, + error: ::Error, ) -> PyResult; - fn host_error(error: &PyErr) -> ::Error; + fn host_error(error: &PyErr) -> ::Error; fn close(&mut self, py: Python<'_>); diff --git a/litellm-rust/crates/host-python/src/driver.rs b/litellm-rust/crates/host-python/src/driver.rs index 77a294d274b..aaa0752522b 100644 --- a/litellm-rust/crates/host-python/src/driver.rs +++ b/litellm-rust/crates/host-python/src/driver.rs @@ -2,10 +2,11 @@ use std::sync::Arc; use std::task::Poll; use futures_util::future::{AbortHandle, Abortable}; +use litellm_host::event::WireRequest; use litellm_host::event::{FailureOrigin, Timing, epoch_seconds}; -use litellm_host::host::{Demand, HostOp, HostResult, HostStep}; +use litellm_host::host::{Demand, HostOp, HostStep, Reply}; use litellm_host::machine::{HostFailure, Machine, MachineStep}; -use litellm_host::route::Route; +use litellm_host::protocol::Protocol; use pyo3::exceptions::{PyBaseException, PyException, PyRuntimeError}; use pyo3::gc::{PyTraverseError, PyVisit}; use pyo3::prelude::*; @@ -13,21 +14,21 @@ use pyo3::types::PyDict; use tokio::sync::Mutex; use crate::adapter::{ - InvokeError, LifecycleEvent, LifecycleStep, PythonLifecycle, RouteHost, missing_state, + InvokeError, LifecycleEvent, LifecycleStep, ProtocolHost, PythonLifecycle, missing_state, }; use crate::execution::{poll_async_value, run_async_value, run_sync_value}; use crate::handle::{Execution, ExecutionBody, ExecutionStep}; -type RouteOf = ::Route; -type ErrorOf = as Route>::Error; -type ResponseOf = as Route>::Response; -type NativeStep = MachineStep, ResponseOf>; +type ProtocolOf = ::Protocol; +type ErrorOf = as Protocol>::Error; +type ResponseOf = as Protocol>::Response; +type NativeStep = MachineStep, ResponseOf>; type NativeResult = Result, ErrorOf>; -type NativeResume = Option>, HostFailure>>>; +type Interruption = Option>>; type MachineResult = Result< - MachineStep<::Route, ::Complete>, - <::Route as Route>::Error, + MachineStep<::Protocol, ::Complete>, + <::Protocol as Protocol>::Error, >; struct MachineState { @@ -44,12 +45,11 @@ enum Stage { Failed(Py), } -#[derive(Clone, Copy)] enum Expect { Started, Arguments, - Wire, - Emitted, + Wire(Reply), + Emitted(Reply<()>), Response, Terminal, } @@ -58,20 +58,30 @@ enum Pending { Native, Adapter(Expect), /// The stream handed to the caller waits for its next read or its close. - Consumer, + Consumer(Reply), } -enum Next { +/// A route answer as the driver resumes on it: a Python exception interrupts the call as +/// raised, a native rejection resumes the machine with it. +fn answered(answer: Result<(), InvokeError>) -> PyResult> { + match answer { + Ok(()) => Ok(Ok(())), + Err(InvokeError::Native(error)) => Ok(Err(error)), + Err(InvokeError::Python(error)) => Err(error), + } +} + +enum Next { Return(ExecutionStep), Continue(HostStep, Py>), } struct PythonDriver where - H: RouteHost, - M: Machine> + 'static, + H: ProtocolHost, + M: Machine> + 'static, { - route: H, + host: H, adapter: Box, machine: Option>>>, arguments: Option>, @@ -89,17 +99,17 @@ where pub fn run_call( py: Python<'_>, machine: M, - route: H, + host: H, adapter: Box, arguments: Py, asynchronous: bool, ) -> PyResult> where - H: RouteHost + 'static, - M: Machine> + 'static, + H: ProtocolHost + 'static, + M: Machine> + 'static, { let mut driver = PythonDriver { - route, + host, adapter, machine: Some(Arc::new(Mutex::new(MachineState { machine, @@ -141,8 +151,8 @@ fn is_cancellation(py: Python<'_>, error: &PyErr) -> bool { impl PythonDriver where - H: RouteHost, - M: Machine> + 'static, + H: ProtocolHost, + M: Machine> + 'static, { fn timing(&self) -> Timing { Timing { @@ -172,13 +182,13 @@ where self.run_steps(py, HostStep::Ready(result)) } (Some(Pending::Native), Some(Err(error))) => self.interrupt(py, error), - (Some(Pending::Consumer), Some(read)) => { - let demand = if read.is_ok() { + (Some(Pending::Consumer(reply)), Some(read)) => { + reply.send(if read.is_ok() { Demand::More } else { Demand::Detached - }; - self.resume_machine(py, Some(Ok(HostResult::Demand(demand)))) + }); + self.resume_machine(py, None) } (Some(Pending::Adapter(expect)), Some(result)) => { match self.adapter.resume(py, result) { @@ -196,22 +206,24 @@ where step: LifecycleStep, expect: Expect, ) -> PyResult { + if let LifecycleStep::Await(awaitable) = step { + self.pending = Some(Pending::Adapter(expect)); + return Ok(ExecutionStep::Await(awaitable)); + } match (expect, step) { - (_, LifecycleStep::Await(awaitable)) => { - self.pending = Some(Pending::Adapter(expect)); - Ok(ExecutionStep::Await(awaitable)) - } (Expect::Started, LifecycleStep::Done) => self.begin(py), (Expect::Arguments, LifecycleStep::Arguments(arguments)) => { self.arguments = Some(arguments); self.stage = Stage::Call; self.resume_machine(py, None) } - (Expect::Wire, LifecycleStep::Wire(wire)) => { - self.resume_machine(py, Some(Ok(HostResult::BeforeSend(wire)))) + (Expect::Wire(reply), LifecycleStep::Wire(wire)) => { + reply.send(*wire); + self.resume_machine(py, None) } - (Expect::Emitted, LifecycleStep::Done) => { - self.resume_machine(py, Some(Ok(HostResult::Emitted))) + (Expect::Emitted(reply), LifecycleStep::Done) => { + reply.send(()); + self.resume_machine(py, None) } (Expect::Response, LifecycleStep::Response(response)) => self.succeeded(py, response), (Expect::Terminal, LifecycleStep::Done) => match &self.stage { @@ -242,9 +254,9 @@ where fn resume_machine( &mut self, py: Python<'_>, - result: NativeResume, + interruption: Interruption, ) -> PyResult { - let step = self.resume_core(py, result)?; + let step = self.resume_core(py, interruption)?; self.run_steps(py, step) } @@ -277,53 +289,62 @@ where } Err(error) => return self.machine_failed(py, error).map(Next::Return), }; - let answer = match op { - HostOp::Route(op) => { + let answered = match op { + HostOp::Project(reply) => { let arguments = self.arguments.as_ref().ok_or_else(missing_state)?; - match self.route.invoke(py, arguments.bind(py), op) { - Ok(result) => Ok(HostResult::Route(result)), - Err(InvokeError::Native(error)) => { - return self - .resume_core(py, Some(Err(HostFailure::Error(error)))) - .map(Next::Continue); - } - Err(InvokeError::Python(error)) => Err(error), - } + let projected = self.host.project(py, arguments.bind(py)); + answered(projected.map(|projection| reply.send(projection))) } - HostOp::BeforeSend { wire, context } => { - match self.adapter.before_send(py, wire, &context) { - Ok(LifecycleStep::Wire(wire)) => Ok(HostResult::BeforeSend(wire)), + HostOp::Custom(op) => answered(self.host.invoke(py, op)), + HostOp::BeforeSend { + wire, + context, + reply, + } => match self.adapter.before_send(py, wire, &context) { + Ok(LifecycleStep::Wire(wire)) => { + reply.send(*wire); + Ok(Ok(())) + } + Ok(LifecycleStep::Await(awaitable)) => { + self.pending = Some(Pending::Adapter(Expect::Wire(reply))); + return Ok(Next::Return(ExecutionStep::Await(awaitable))); + } + Ok(_) => return Err(missing_state()), + Err(error) => Err(error), + }, + HostOp::Open(_, reply) => return self.opened(py, reply).map(Next::Return), + HostOp::Deliver(chunk, reply) => { + return self.delivered(py, chunk, reply).map(Next::Return); + } + HostOp::Emit(event, reply) => { + match self.adapter.emit(py, LifecycleEvent::Machine(&event)) { + Ok(LifecycleStep::Done) => { + reply.send(()); + Ok(Ok(())) + } Ok(LifecycleStep::Await(awaitable)) => { - self.pending = Some(Pending::Adapter(Expect::Wire)); + self.pending = Some(Pending::Adapter(Expect::Emitted(reply))); return Ok(Next::Return(ExecutionStep::Await(awaitable))); } Ok(_) => return Err(missing_state()), Err(error) => Err(error), } } - HostOp::Open(_) => return self.opened(py).map(Next::Return), - HostOp::Deliver(chunk) => return self.delivered(py, chunk).map(Next::Return), - HostOp::Emit(event) => match self.adapter.emit(py, LifecycleEvent::Machine(&event)) { - Ok(LifecycleStep::Done) => Ok(HostResult::Emitted), - Ok(LifecycleStep::Await(awaitable)) => { - self.pending = Some(Pending::Adapter(Expect::Emitted)); - return Ok(Next::Return(ExecutionStep::Await(awaitable))); - } - Ok(_) => return Err(missing_state()), - Err(error) => Err(error), - }, }; - match answer { - Ok(answer) => self.resume_core(py, Some(Ok(answer))).map(Next::Continue), + match answered { + Ok(Ok(())) => self.resume_core(py, None).map(Next::Continue), + Ok(Err(native)) => self + .resume_core(py, Some(HostFailure::Error(native))) + .map(Next::Continue), Err(error) => self.interrupt(py, error).map(Next::Return), } } - fn opened(&mut self, py: Python<'_>) -> PyResult { + fn opened(&mut self, py: Python<'_>, reply: Reply) -> PyResult { self.stage = Stage::Streaming; match self.adapter.opened(py) { Ok(()) => { - self.pending = Some(Pending::Consumer); + self.pending = Some(Pending::Consumer(reply)); Ok(ExecutionStep::Open) } Err(error) => self.interrupt(py, error), @@ -333,15 +354,16 @@ where fn delivered( &mut self, py: Python<'_>, - chunk: as Route>::Chunk, + chunk: as Protocol>::Chunk, + reply: Reply, ) -> PyResult { - let chunk = match self.route.chunk(py, chunk) { + let chunk = match self.host.chunk(py, chunk) { Ok(chunk) => chunk, Err(error) => return self.interrupt(py, error), }; match self.adapter.delivered(py, &chunk) { Ok(()) => { - self.pending = Some(Pending::Consumer); + self.pending = Some(Pending::Consumer(reply)); Ok(ExecutionStep::Yield(chunk)) } Err(error) => self.interrupt(py, error), @@ -357,25 +379,24 @@ where } else { HostFailure::Error(native) }; - self.resume_machine(py, Some(Err(failure))) + self.resume_machine(py, Some(failure)) } fn resume_core( &mut self, py: Python<'_>, - result: NativeResume, + interruption: Interruption, ) -> PyResult, Py>> { let state = Arc::clone(self.machine.as_ref().ok_or_else(missing_state)?); let future = async move { let mut state = state.lock().await; - let result = match result { - Some(Err(failure)) => state + let result = match interruption { + Some(failure) => state .machine .interrupt(failure) .await .map(MachineStep::Complete), - Some(Ok(result)) => state.machine.resume(Some(result)).await, - None => state.machine.resume(None).await, + None => state.machine.resume().await, }; state.result = Some(result); Ok(()) @@ -414,7 +435,7 @@ where fn completed(&mut self, py: Python<'_>, response: ResponseOf) -> PyResult { self.ended_at = Some(epoch_seconds()); - let public = match self.route.complete(py, response) { + let public = match self.host.complete(py, response) { Ok(public) => public, Err(error) => return self.failure(py, error, FailureOrigin::Call), }; @@ -441,7 +462,7 @@ where /// fails, that failure is raised with the native error's text as its `__context__`. fn classified(&self, py: Python<'_>, error: ErrorOf) -> PyErr { let native = error.to_string(); - let classifier_error = match self.route.classify(py, error) { + let classifier_error = match self.host.classify(py, error) { Ok(failure) => return failure.into(), Err(classifier_error) => classifier_error, }; @@ -486,7 +507,7 @@ where if self.machine.take().is_some() { Python::attach(|py| { self.adapter.close(py); - self.route.close(py); + self.host.close(py); }); } } @@ -494,15 +515,15 @@ where impl ExecutionBody for PythonDriver where - H: RouteHost, - M: Machine> + 'static, + H: ProtocolHost, + M: Machine> + 'static, { fn resume(&mut self, result: Option>>) -> PyResult { Python::attach(|py| self.drive(py, result)) } fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError> { - self.route.traverse(visit)?; + self.host.traverse(visit)?; self.adapter.traverse(visit)?; visit.call(&self.arguments)?; visit.call(&self.interrupted)?; @@ -516,8 +537,8 @@ where impl Drop for PythonDriver where - H: RouteHost, - M: Machine> + 'static, + H: ProtocolHost, + M: Machine> + 'static, { fn drop(&mut self) { self.clear(); @@ -528,8 +549,8 @@ where mod tests { use std::sync::{Arc, Mutex}; - use litellm_host::event::{MachineEvent, RequestContext, WireRequest}; - use litellm_host::machine::{Interrupted, Step}; + use litellm_host::event::{MachineEvent, RawResponse, RequestContext}; + use litellm_host::machine::{CallMachine, MachineFault}; use pyo3::exceptions::{PyBaseException, PyValueError}; use pyo3::types::PyDict; @@ -573,22 +594,21 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri } } - struct Synthetic; - - impl Route for Synthetic { - type Response = String; - type Error = Error; - type Op = &'static str; - type OpResult = String; - type Chunk = std::convert::Infallible; - type StreamHead = std::convert::Infallible; + impl From for Error { + fn from(fault: MachineFault) -> Self { + Self(format!("{fault:?}")) + } } - /// Yields the scripted ops in order, then completes or fails as scripted. - struct ScriptedMachine { - ops: Vec>, - outcome: Option>, - answers: Vec, + struct Synthetic; + + impl Protocol for Synthetic { + type Response = String; + type Error = Error; + type Projection = String; + type Op = (&'static str, Reply); + type Chunk = std::convert::Infallible; + type StreamHead = std::convert::Infallible; } fn wire() -> WireRequest { @@ -609,37 +629,6 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri } } - impl Machine for ScriptedMachine { - type Route = Synthetic; - type Complete = String; - - fn resume(&mut self, result: Option>) -> Step<'_, Self> { - Box::pin(async move { - if let Some(result) = result { - self.answers.push(match result { - HostResult::Route(value) => value, - HostResult::BeforeSend(wire) => wire.url, - HostResult::Emitted => "emitted".into(), - HostResult::Demand(demand) => format!("{demand:?}"), - }); - } - if !self.ops.is_empty() { - return Ok(MachineStep::Host(self.ops.remove(0))); - } - self.outcome - .take() - .ok_or_else(|| Error("resumed after completion".into()))? - .map(MachineStep::Complete) - }) - } - - fn interrupt(&mut self, failure: HostFailure) -> Interrupted<'_, Self> { - self.ops.clear(); - self.outcome = None; - Box::pin(async move { Err(failure.into_error()) }) - } - } - #[derive(Default)] struct Log(Arc>>); @@ -677,22 +666,37 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri } } - impl RouteHost for SyntheticHost { - type Route = Synthetic; + impl SyntheticHost { + fn answer(&self, value: impl FnOnce() -> String) -> Result> { + match self.op { + OpScript::Answer => Ok(value()), + OpScript::RaisePython => Err(PyValueError::new_err("op failed").into()), + OpScript::RejectNatively => Err(InvokeError::Native(Error("op rejected".into()))), + } + } + } + + impl ProtocolHost for SyntheticHost { + type Protocol = Synthetic; type Failure = Classified; + fn project( + &mut self, + _: Python<'_>, + arguments: &Bound<'_, PyDict>, + ) -> Result> { + self.log.push("project"); + self.answer(|| format!("project:{}", arguments.len())) + } + fn invoke( &mut self, _: Python<'_>, - arguments: &Bound<'_, PyDict>, - op: &'static str, - ) -> Result> { - self.log.push(format!("route:{op}")); - match self.op { - OpScript::Answer => Ok(format!("{op}:{}", arguments.len())), - OpScript::RaisePython => Err(PyValueError::new_err("op failed").into()), - OpScript::RejectNatively => Err(InvokeError::Native(Error("op rejected".into()))), - } + (op, reply): (&'static str, Reply), + ) -> Result<(), InvokeError> { + self.log.push(format!("op:{op}")); + self.answer(|| op.to_string()) + .map(|answer| reply.send(answer)) } fn chunk(&mut self, _: Python<'_>, chunk: std::convert::Infallible) -> PyResult> { @@ -719,7 +723,7 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri } fn close(&mut self, _: Python<'_>) { - self.log.push("route.close"); + self.log.push("host.close"); } fn traverse(&self, _: &PyVisit<'_>) -> Result<(), PyTraverseError> { @@ -828,7 +832,7 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri fn run_scripted( py: Python<'_>, - machine: ScriptedMachine, + machine: CallMachine, op: OpScript, script: AdapterScript, asynchronous: bool, @@ -848,12 +852,12 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri fn run_hosted( py: Python<'_>, - machine: ScriptedMachine, - route: SyntheticHost, + machine: CallMachine, + host: SyntheticHost, script: AdapterScript, asynchronous: bool, ) -> (PyResult>, Vec) { - let log = Log(route.log.0.clone()); + let log = Log(host.log.0.clone()); let adapter = SyntheticAdapter { log: Log(log.0.clone()), script, @@ -863,7 +867,7 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri let result = run_call( py, machine, - route, + host, Box::new(adapter), arguments.unbind(), asynchronous, @@ -884,21 +888,21 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri (result, log.entries()) } - fn success_machine() -> ScriptedMachine { - ScriptedMachine { - ops: vec![ - HostOp::Route("project"), - HostOp::BeforeSend { - wire: Box::new(wire()), - context: Box::new(context()), - }, - HostOp::Emit(MachineEvent::ResponseReceived { - raw: litellm_host::event::RawResponse { body: "raw".into() }, - }), - ], - outcome: Some(Ok("done".into())), - answers: Vec::new(), - } + /// Answers to projection, to the route op and to `before_send` all reach the + /// response, so a driver that misroutes a reply changes what the call returns. + fn success_machine() -> CallMachine { + CallMachine::new(|host| { + Box::pin(async move { + let projected = host.project().await?; + let signed = host.custom_op(|reply| ("sign", reply)).await?; + let wire = host.before_send(wire(), context()).await?; + host.emit(MachineEvent::ResponseReceived { + raw: RawResponse { body: "raw".into() }, + }) + .await?; + Ok(format!("{projected}|{signed}|{}", wire.url)) + }) + }) } #[test] @@ -917,32 +921,37 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri AdapterScript::Plain, asynchronous, ); - assert_eq!(result.unwrap().extract::(py).unwrap(), "done"); + assert_eq!( + result.unwrap().extract::(py).unwrap(), + "project:1|sign|rewritten" + ); assert_eq!( log, [ "started", "begin", - "route:project", + "project", + "op:sign", "before_send", "response:raw", "complete", "after_success", - "succeeded:done", + "succeeded:project:1|sign|rewritten", "adapter.close", - "route.close", + "host.close", ] ); } }); } - fn failing_machine() -> ScriptedMachine { - ScriptedMachine { - ops: vec![HostOp::Route("project")], - outcome: Some(Err(Error("provider exploded".into()))), - answers: Vec::new(), - } + fn failing_machine() -> CallMachine { + CallMachine::new(|host| { + Box::pin(async move { + host.project().await?; + Err(Error("provider exploded".into())) + }) + }) } #[test] @@ -969,11 +978,11 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri [ "started", "begin", - "route:project", + "project", "classify:provider exploded", "failed:Call:classified: provider exploded", "adapter.close", - "route.close", + "host.close", ] ); } @@ -1003,11 +1012,11 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri [ "started", "begin", - "route:project", + "project", "classify:op rejected", "failed:Call:classified: op rejected", "adapter.close", - "route.close", + "host.close", ] ); }); @@ -1035,10 +1044,10 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri [ "started", "begin", - "route:project", + "project", "failed:Call:op failed", "adapter.close", - "route.close", + "host.close", ] ); }); @@ -1073,11 +1082,11 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri [ "started", "begin", - "route:project", + "project", "classify:provider exploded", "failed:Call:classifier failed", "adapter.close", - "route.close", + "host.close", ] ); }); @@ -1106,7 +1115,7 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri "begin", "failed:Host:begin failed", "adapter.close", - "route.close" + "host.close" ] ); }); @@ -1130,7 +1139,7 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri ); assert_eq!(result.unwrap().extract::(py).unwrap(), "replaced"); assert!(log.contains(&"succeeded:replaced".to_string())); - assert!(!log.contains(&"succeeded:done".to_string())); + assert!(!log.contains(&"succeeded:project:1|rewritten".to_string())); } }); } @@ -1159,7 +1168,7 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri "after_success", "failed:Host:after_success failed", "adapter.close", - "route.close" + "host.close" ] ); assert!(!log.iter().any(|entry| entry.starts_with("succeeded"))); @@ -1175,18 +1184,24 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri crate::initialize_python(); Python::attach(|py| { struct Cancelling(Log); - impl RouteHost for Cancelling { - type Route = Synthetic; + impl ProtocolHost for Cancelling { + type Protocol = Synthetic; type Failure = Classified; - fn invoke( + fn project( &mut self, _: Python<'_>, _: &Bound<'_, PyDict>, - _: &'static str, ) -> Result> { - self.0.push("route"); + self.0.push("project"); Err(pyo3::exceptions::asyncio::CancelledError::new_err(()).into()) } + fn invoke( + &mut self, + _: Python<'_>, + _: (&'static str, Reply), + ) -> Result<(), InvokeError> { + Err(missing_state().into()) + } fn chunk( &mut self, _: Python<'_>, @@ -1210,7 +1225,7 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri } } let log = Log::default(); - let route = Cancelling(Log(log.0.clone())); + let host = Cancelling(Log(log.0.clone())); let adapter = SyntheticAdapter { log: Log(log.0.clone()), script: AdapterScript::Plain, @@ -1218,7 +1233,7 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri let error = run_call( py, success_machine(), - route, + host, Box::new(adapter), PyDict::new(py).unbind(), false, @@ -1227,7 +1242,7 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri assert!(!error.is_instance_of::(py)); assert_eq!( log.entries(), - ["started", "begin", "route", "adapter.close"] + ["started", "begin", "project", "adapter.close"] ); }); } diff --git a/litellm-rust/crates/host-python/src/file_reader.rs b/litellm-rust/crates/host-python/src/file_reader.rs new file mode 100644 index 00000000000..bbc7a233b28 --- /dev/null +++ b/litellm-rust/crates/host-python/src/file_reader.rs @@ -0,0 +1,241 @@ +//! A caller's file-like object: anything with a callable `read`, kept as a handle and read +//! once, on the host's thread, into bytes Rust owns. + +use bytes::Bytes; +use pyo3::{ + exceptions::PyTypeError, + gc::{PyTraverseError, PyVisit}, + prelude::*, + pybacked::PyBackedBytes, + types::{PyBytes, PyString}, +}; + +#[derive(Clone, Debug, PartialEq, Eq)] +pub struct FileContent { + pub bytes: Bytes, + pub file_name: Option, +} + +#[derive(Debug)] +pub struct PythonFileReader { + reader: Py, + name: Option, +} + +impl PythonFileReader { + /// `None` when `file` has no callable `read`. The object's `name` is read now, its + /// contents only on [`read`](Self::read). + pub fn from_file_like(file: &Bound<'_, PyAny>) -> PyResult> { + let reader = file + .getattr_opt("read")? + .filter(|value| value.is_callable()); + let Some(reader) = reader else { + return Ok(None); + }; + let name = file + .getattr_opt("name")? + .filter(|value| !value.is_none()) + .map(|value| value.extract::()) + .transpose()?; + Ok(Some(Self { + reader: reader.unbind(), + name, + })) + } + + pub fn read(&self, py: Python<'_>) -> PyResult { + let value = self.reader.bind(py).call0()?; + let bytes = if value.is_instance_of::() { + Bytes::from(value.extract::()?) + } else if value.is_instance_of::() { + py_bytes(&value)? + } else { + return Err(PyTypeError::new_err(format!( + "file read must return bytes or str, got {}", + value.get_type(), + ))); + }; + Ok(FileContent { + bytes, + file_name: self.name.clone(), + }) + } + + pub fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError> { + visit.call(&self.reader) + } +} + +/// An exact `bytes` object is shared without copying and keeps the Python object alive; +/// a `bytes` subclass is copied. +pub fn py_bytes(value: &Bound<'_, PyAny>) -> PyResult { + if value.is_exact_instance_of::() { + return Ok(Bytes::from_owner(value.extract::()?)); + } + Ok(Bytes::copy_from_slice( + value.extract::()?.as_ref(), + )) +} + +#[cfg(test)] +mod tests { + use pyo3::{exceptions::PyTypeError, types::PyDict}; + + use super::*; + + fn eval<'py>(py: Python<'py>, source: &std::ffi::CStr) -> Bound<'py, PyDict> { + let locals = PyDict::new(py); + py.run(source, Some(&locals), Some(&locals)).unwrap(); + locals + } + + fn reader<'py>(locals: &Bound<'py, PyDict>, name: &str) -> PythonFileReader { + PythonFileReader::from_file_like(&locals.get_item(name).unwrap().unwrap()) + .unwrap() + .unwrap() + } + + #[test] + fn objects_without_a_callable_read_are_not_readers() { + Python::initialize(); + Python::attach(|py| { + let locals = eval( + py, + c" +class Attribute: + read = 'not callable' +plain = object() +attribute = Attribute() +", + ); + for name in ["plain", "attribute"] { + let file = locals.get_item(name).unwrap().unwrap(); + assert!(PythonFileReader::from_file_like(&file).unwrap().is_none()); + } + }); + } + + #[test] + fn the_name_is_taken_up_front_and_the_contents_only_on_read() { + Python::initialize(); + Python::attach(|py| { + let locals = eval( + py, + c" +class Reader: + name = 'scan.png' + def __init__(self): + self.reads = 0 + def read(self): + self.reads += 1 + return b'abc' +file = Reader() +", + ); + let reads = || { + locals + .get_item("file") + .unwrap() + .unwrap() + .getattr("reads") + .unwrap() + .extract::() + .unwrap() + }; + let file = reader(&locals, "file"); + assert_eq!(reads(), 0); + let content = file.read(py).unwrap(); + assert_eq!(reads(), 1); + assert_eq!( + content, + FileContent { + bytes: b"abc".as_slice().into(), + file_name: Some("scan.png".into()), + } + ); + }); + } + + #[test] + fn read_results_are_normalized_and_exceptions_keep_their_identity() { + Python::initialize(); + Python::attach(|py| { + let locals = eval( + py, + c" +failure = KeyError('reader failed') +class Raising: + def read(self): + raise failure +class Text: + def read(self): + return 'héllo' +class Wrong: + def read(self): + return 7 +raising = Raising() +text = Text() +wrong = Wrong() +", + ); + let error = reader(&locals, "raising").read(py).unwrap_err(); + assert!( + error + .value(py) + .is(locals.get_item("failure").unwrap().unwrap()) + ); + assert_eq!( + reader(&locals, "text").read(py).unwrap().bytes.as_ref(), + "héllo".as_bytes() + ); + let error = reader(&locals, "wrong").read(py).unwrap_err(); + assert!(error.is_instance_of::(py)); + assert!(error.to_string().contains("bytes or str")); + }); + } + + #[rstest::rstest] + #[case::read("read")] + #[case::name("name")] + fn attribute_failures_keep_their_identity(#[case] attribute: &str) { + Python::initialize(); + Python::attach(|py| { + let locals = eval( + py, + c" +failure = LookupError('file property failed') +class File: + def __getattribute__(self, name): + if name == attribute: + raise failure + return super().__getattribute__(name) + name = 'scan.pdf' + def read(self): + return b'abc' +file = File() +", + ); + locals.set_item("attribute", attribute).unwrap(); + let error = + PythonFileReader::from_file_like(&locals.get_item("file").unwrap().unwrap()) + .unwrap_err(); + assert!( + error + .value(py) + .is(locals.get_item("failure").unwrap().unwrap()) + ); + }); + } + + #[test] + fn exact_python_bytes_transfer_without_copying_and_outlive_the_input() { + Python::initialize(); + let (bytes, pointer) = Python::attach(|py| { + let value = PyBytes::new(py, b"document bytes"); + let pointer = value.as_bytes().as_ptr() as usize; + (py_bytes(value.as_any()).unwrap(), pointer) + }); + assert_eq!(bytes.as_ptr() as usize, pointer); + assert_eq!(bytes.as_ref(), b"document bytes"); + } +} diff --git a/litellm-rust/crates/host-python/src/lib.rs b/litellm-rust/crates/host-python/src/lib.rs index 55b27e34b46..7e17c4da51e 100644 --- a/litellm-rust/crates/host-python/src/lib.rs +++ b/litellm-rust/crates/host-python/src/lib.rs @@ -1,6 +1,6 @@ //! The CPython runtime adapter: value marshalling, interpreter detachment, the tokio and //! asyncio glue, and the driver that runs a native [`Machine`](litellm_host::machine::Machine) -//! against a Python route host and a Python lifecycle. Everything here is Python-specific by +//! against a Python protocol host and a Python lifecycle. Everything here is Python-specific by //! construction; another host language gets its own crate of the same shape. mod adapter; @@ -8,13 +8,14 @@ mod argument; mod callable; mod driver; mod execution; +mod file_reader; mod fork_gate; mod gil; mod handle; mod marshal; pub use adapter::{ - InvokeError, LifecycleEvent, LifecycleStep, PythonLifecycle, RouteHost, missing_state, + InvokeError, LifecycleEvent, LifecycleStep, ProtocolHost, PythonLifecycle, missing_state, }; pub use argument::lookup; pub use callable::wrap_failure; @@ -24,6 +25,7 @@ pub use execution::{ reserve_process_for_forking, run_async, run_async_value, run_sync, run_sync_value, runtime_started, }; +pub use file_reader::{FileContent, PythonFileReader, py_bytes}; pub use fork_gate::RuntimeAlreadyStarted; pub use gil::{PythonContext, attach_blocking, release_count, release_gil}; pub use handle::{Execution, ExecutionBody, ExecutionStep}; diff --git a/litellm-rust/crates/host/Cargo.toml b/litellm-rust/crates/host/Cargo.toml index 0c7c46192b5..bbbed68f345 100644 --- a/litellm-rust/crates/host/Cargo.toml +++ b/litellm-rust/crates/host/Cargo.toml @@ -7,6 +7,7 @@ repository.workspace = true [dependencies] litellm-auth.workspace = true +litellm-coroutine.workspace = true serde_json.workspace = true tokio = { workspace = true, features = ["sync"] } diff --git a/litellm-rust/crates/host/src/host.rs b/litellm-rust/crates/host/src/host.rs index aba35185a18..9714b9470a3 100644 --- a/litellm-rust/crates/host/src/host.rs +++ b/litellm-rust/crates/host/src/host.rs @@ -1,28 +1,27 @@ use std::future::Future; -use crate::event::{CallEvent, MachineEvent, RequestContext, WireRequest}; -use crate::route::Route; +pub use litellm_coroutine::{Abandoned, Answer, Reply, reply}; -/// One suspension point of a native call, performed by the host. -pub enum HostOp { - Route(R::Op), +use crate::event::{CallEvent, MachineEvent, RequestContext, WireRequest}; +use crate::protocol::Protocol; + +/// One suspension point of a native call, performed by the host and answered through the +/// [`Reply`] it carries. +pub enum HostOp { + /// The first op of every call: the caller's request as the host projects it. + Project(Reply), + Custom(R::Op), BeforeSend { wire: Box, context: Box, + reply: Reply, }, - Emit(MachineEvent), + Emit(MachineEvent, Reply<()>), /// The response streams: the host hands the caller a stream and answers once the /// caller asks for the first chunk or goes away. - Open(R::StreamHead), + Open(R::StreamHead, Reply), /// The next chunk of an open stream, answered once the caller asks for the one after. - Deliver(R::Chunk), -} - -pub enum HostResult { - Route(R::OpResult), - BeforeSend(Box), - Emitted, - Demand(Demand), + Deliver(R::Chunk, Reply), } /// Whether the caller of a streamed call still reads it. @@ -39,10 +38,13 @@ pub enum HostStep { Suspend(S), } -/// An in-process host: answers route operations and observes the call without leaving +/// An in-process host: answers custom operations and observes the call without leaving /// the Rust runtime. Language hosts implement their own driver instead. -pub trait Host: Send + Sync { - fn route(&self, op: R::Op) -> impl Future> + Send; +pub trait Host: Send + Sync { + fn project(&self) -> impl Future> + Send; + + /// Answers `op` through its reply, or fails the call. + fn custom_op(&self, op: R::Op) -> impl Future> + Send; fn before_send( &self, diff --git a/litellm-rust/crates/host/src/lib.rs b/litellm-rust/crates/host/src/lib.rs index 65479c2380f..c6b9e59b65a 100644 --- a/litellm-rust/crates/host/src/lib.rs +++ b/litellm-rust/crates/host/src/lib.rs @@ -1,12 +1,13 @@ //! The contract between a native call and the host runtime that drives it. //! //! A host is whatever sits on the far side of the language boundary: CPython today, -//! another runtime later. Core runs each route on a [`machine::RouteMachine`] and never learns +//! another runtime later. Core runs each route on a [`machine::CallMachine`] and never learns //! which host is on the other end. The machine yields [`host::HostOp`]s; a driver answers -//! them, observes [`event::CallEvent`]s and may rewrite the wire request before it is sent. +//! each through the typed [`host::Reply`] it carries, observes [`event::CallEvent`]s and +//! may rewrite the wire request before it is sent. pub mod event; pub mod host; pub mod machine; -pub mod route; +pub mod protocol; pub mod run; diff --git a/litellm-rust/crates/host/src/machine/auth.rs b/litellm-rust/crates/host/src/machine/auth.rs index ba7e242e766..76e3504ca28 100644 --- a/litellm-rust/crates/host/src/machine/auth.rs +++ b/litellm-rust/crates/host/src/machine/auth.rs @@ -1,22 +1,21 @@ use std::sync::Arc; use super::{HostChannel, MachineFault}; -use crate::route::Route; +use crate::{host::Reply, protocol::Protocol}; use litellm_auth::{Error, ResolvedCredential, TokenFuture, TokenProvider, TokenProviderHandle}; -/// A route whose host can mint credentials on the call's behalf. -pub trait TokenRoute: Route { - fn acquire_token_op() -> Self::Op; - fn token_credential(result: Self::OpResult) -> Option; +/// A protocol whose host can mint credentials on the call's behalf. +pub trait TokenProtocol: Protocol { + fn acquire_token_op(reply: Reply) -> Self::Op; } /// A [`TokenProvider`] that asks the host for each credential through the call's own /// operation channel, so the host answers it on the caller's thread and context. -pub struct HostTokenProvider { +pub struct HostTokenProvider { channel: HostChannel, } -impl std::fmt::Debug for HostTokenProvider { +impl std::fmt::Debug for HostTokenProvider { fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { formatter.write_str("HostTokenProvider") } @@ -24,7 +23,7 @@ impl std::fmt::Debug for HostTokenProvider { impl HostTokenProvider where - R: TokenRoute, + R: TokenProtocol, R::Error: From + std::fmt::Display, { pub fn handle(channel: HostChannel) -> TokenProviderHandle { @@ -34,19 +33,15 @@ where impl TokenProvider for HostTokenProvider where - R: TokenRoute, + R: TokenProtocol, R::Error: From + std::fmt::Display, { fn acquire(&self) -> TokenFuture<'_> { Box::pin(async move { - let result = self - .channel - .route(R::acquire_token_op()) + self.channel + .custom_op(R::acquire_token_op) .await - .map_err(|error| Error::AzureTokenAcquisition(error.to_string()))?; - R::token_credential(result).ok_or_else(|| { - Error::AzureTokenAcquisition("invalid token provider host result".into()) - }) + .map_err(|error| Error::AzureTokenAcquisition(error.to_string())) }) } } diff --git a/litellm-rust/crates/host/src/machine/call_machine.rs b/litellm-rust/crates/host/src/machine/call_machine.rs new file mode 100644 index 00000000000..af0bc50fbe6 --- /dev/null +++ b/litellm-rust/crates/host/src/machine/call_machine.rs @@ -0,0 +1,137 @@ +//! The one machine every route runs on: the route's provider future as a +//! [`Coroutine`] that yields [`HostOp`]s, each answered through its own typed reply. No +//! task is spawned; dropping the machine drops the in-flight call. + +use std::{future::Future, pin::Pin}; + +use litellm_coroutine::{Co, Coroutine, CoroutineState, ResumeError}; + +use super::{HostFailure, Interrupted, Machine, MachineStep, Step}; +use crate::{ + event::{MachineEvent, RequestContext, WireRequest}, + host::{Demand, HostOp, Reply}, + protocol::Protocol, +}; + +/// The machine's own failures, distinct from anything the provider call reports. +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub enum MachineFault { + /// The host dropped an op's reply unanswered, or went away while the call waited. + Abandoned, + /// The host resumed the call out of turn. + Protocol(ResumeError), +} + +pub type ExecuteFuture = + Pin::Response, ::Error>> + Send>>; + +/// The provider side of the machine: how the in-flight call reaches its host. +pub struct HostChannel { + co: Co>, +} + +impl Clone for HostChannel { + fn clone(&self) -> Self { + Self { + co: self.co.clone(), + } + } +} + +impl HostChannel +where + R::Error: From, +{ + async fn yield_( + &self, + ask: impl FnOnce(Reply) -> HostOp + Send, + ) -> Result { + self.co + .yield_(ask) + .await + .map_err(|_| MachineFault::Abandoned.into()) + } + + pub async fn project(&self) -> Result { + self.yield_(HostOp::Project).await + } + + /// Asks the host to perform the custom operation `ask` builds around its reply, as in + /// `host.custom_op(OcrOp::AcquireAzureAdToken)`. + pub async fn custom_op( + &self, + ask: impl FnOnce(Reply) -> R::Op + Send, + ) -> Result { + self.yield_(|reply| HostOp::Custom(ask(reply))).await + } + + pub async fn before_send( + &self, + wire: WireRequest, + context: RequestContext, + ) -> Result { + self.yield_(|reply| HostOp::BeforeSend { + wire: Box::new(wire), + context: Box::new(context), + reply, + }) + .await + } + + pub async fn emit(&self, event: MachineEvent) -> Result<(), R::Error> { + self.yield_(|reply| HostOp::Emit(event, reply)).await + } + + pub async fn open(&self, head: R::StreamHead) -> Result { + self.yield_(|reply| HostOp::Open(head, reply)).await + } + + pub async fn deliver(&self, chunk: R::Chunk) -> Result { + self.yield_(|reply| HostOp::Deliver(chunk, reply)).await + } +} + +type CallCoroutine = + Coroutine, Result<::Response, ::Error>>; + +pub struct CallMachine { + coroutine: CallCoroutine, +} + +impl CallMachine +where + R::Error: From, +{ + pub fn new(execute: impl FnOnce(HostChannel) -> ExecuteFuture + Send + 'static) -> Self { + Self { + coroutine: Coroutine::new(|co| execute(HostChannel { co })), + } + } +} + +impl Machine for CallMachine +where + R::Error: From, +{ + type Protocol = R; + type Complete = R::Response; + + fn resume(&mut self) -> Step<'_, Self> { + Box::pin(async move { + match self + .coroutine + .resume() + .await + .map_err(MachineFault::Protocol)? + { + CoroutineState::Yielded(op) => Ok(MachineStep::Host(op)), + CoroutineState::Complete(outcome) => outcome.map(MachineStep::Complete), + } + }) + } + + fn interrupt(&mut self, failure: HostFailure) -> Interrupted<'_, Self> { + self.coroutine.cancel(); + Box::pin(async move { Err(failure.into_error()) }) + } +} diff --git a/litellm-rust/crates/host/src/machine/mod.rs b/litellm-rust/crates/host/src/machine/mod.rs index 2c26db61582..0c7501633fa 100644 --- a/litellm-rust/crates/host/src/machine/mod.rs +++ b/litellm-rust/crates/host/src/machine/mod.rs @@ -1,16 +1,16 @@ mod auth; -mod route_machine; +mod call_machine; use std::future::Future; use std::pin::Pin; -pub use auth::{HostTokenProvider, TokenRoute}; -pub use route_machine::{ExecuteFuture, HostChannel, MachineFault, RouteMachine}; +pub use auth::{HostTokenProvider, TokenProtocol}; +pub use call_machine::{CallMachine, ExecuteFuture, HostChannel, MachineFault}; -use crate::host::{HostOp, HostResult}; -use crate::route::Route; +use crate::host::HostOp; +use crate::protocol::Protocol; -pub enum MachineStep { +pub enum MachineStep { Host(HostOp), Complete(C), } @@ -19,8 +19,8 @@ pub type Step<'a, M> = Pin< Box< dyn Future< Output = Result< - MachineStep<::Route, ::Complete>, - <::Route as Route>::Error, + MachineStep<::Protocol, ::Complete>, + <::Protocol as Protocol>::Error, >, > + Send + 'a, @@ -30,7 +30,10 @@ pub type Step<'a, M> = Pin< pub type Interrupted<'a, M> = Pin< Box< dyn Future< - Output = Result<::Complete, <::Route as Route>::Error>, + Output = Result< + ::Complete, + <::Protocol as Protocol>::Error, + >, > + Send + 'a, >, @@ -51,19 +54,18 @@ impl HostFailure { } /// A resumable call. Core implements it per route; a host drives it. Every suspension -/// point is an op the host performs and answers with a result. +/// point is an op the host performs and answers through the op's own reply before it +/// resumes the call again. pub trait Machine: Send { - type Route: Route; + type Protocol: Protocol; type Complete: Send + 'static; - /// `None` on the first call and whenever the previous step completed without - /// yielding an op; otherwise the result of the op last yielded. - fn resume(&mut self, result: Option>) -> Step<'_, Self>; + fn resume(&mut self) -> Step<'_, Self>; /// The host failed to perform the pending op, or the caller cancelled. The call /// yields no further ops. fn interrupt( &mut self, - failure: HostFailure<::Error>, + failure: HostFailure<::Error>, ) -> Interrupted<'_, Self>; } diff --git a/litellm-rust/crates/host/src/machine/route_machine.rs b/litellm-rust/crates/host/src/machine/route_machine.rs deleted file mode 100644 index 38a0b8bc16a..00000000000 --- a/litellm-rust/crates/host/src/machine/route_machine.rs +++ /dev/null @@ -1,199 +0,0 @@ -//! The one machine every route runs on: it owns the route's provider future, polls it in -//! place, and turns the host operations that future requests into [`Machine`] steps. No -//! task is spawned; dropping the machine drops the in-flight call. - -use std::{future::Future, pin::Pin}; - -use tokio::sync::{mpsc, oneshot}; - -use super::{HostFailure, Interrupted, Machine, MachineStep, Step}; -use crate::{ - event::{MachineEvent, RequestContext, WireRequest}, - host::{Demand, HostOp, HostResult}, - route::Route, -}; - -/// The machine's own failures, distinct from anything the provider call reports. -#[derive(Clone, Copy, Debug, PartialEq, Eq)] -pub enum MachineFault { - /// The host driver went away while the call was waiting on it. - Abandoned, - /// The host answered out of turn: a result with nothing pending, or nothing when a - /// result was pending. - Protocol(&'static str), - /// The host answered a route operation with the wrong result variant. - Mismatch, -} - -pub type ExecuteFuture = - Pin::Response, ::Error>> + Send>>; - -struct PendingOp { - op: HostOp, - reply: oneshot::Sender>, -} - -/// The provider side of the machine: how the in-flight call reaches its host. -pub struct HostChannel { - ops: mpsc::UnboundedSender>, -} - -impl Clone for HostChannel { - fn clone(&self) -> Self { - Self { - ops: self.ops.clone(), - } - } -} - -impl HostChannel -where - R::Error: From, -{ - async fn invoke(&self, op: HostOp) -> Result, R::Error> { - let (reply, answer) = oneshot::channel(); - self.ops - .send(PendingOp { op, reply }) - .map_err(|_| MachineFault::Abandoned)?; - answer.await.map_err(|_| MachineFault::Abandoned.into()) - } - - pub async fn route(&self, op: R::Op) -> Result { - match self.invoke(HostOp::Route(op)).await? { - HostResult::Route(result) => Ok(result), - _ => Err(MachineFault::Mismatch.into()), - } - } - - pub async fn before_send( - &self, - wire: WireRequest, - context: RequestContext, - ) -> Result { - let op = HostOp::BeforeSend { - wire: Box::new(wire), - context: Box::new(context), - }; - match self.invoke(op).await? { - HostResult::BeforeSend(wire) => Ok(*wire), - _ => Err(MachineFault::Mismatch.into()), - } - } - - pub async fn emit(&self, event: MachineEvent) -> Result<(), R::Error> { - match self.invoke(HostOp::Emit(event)).await? { - HostResult::Emitted => Ok(()), - _ => Err(MachineFault::Mismatch.into()), - } - } - - pub async fn open(&self, head: R::StreamHead) -> Result { - self.demand(HostOp::Open(head)).await - } - - pub async fn deliver(&self, chunk: R::Chunk) -> Result { - self.demand(HostOp::Deliver(chunk)).await - } - - async fn demand(&self, op: HostOp) -> Result { - match self.invoke(op).await? { - HostResult::Demand(demand) => Ok(demand), - _ => Err(MachineFault::Mismatch.into()), - } - } -} - -enum Execution { - Unstarted(Box) -> ExecuteFuture + Send>), - Running(ExecuteFuture), - Done, -} - -pub struct RouteMachine { - execution: Execution, - ops: mpsc::UnboundedReceiver>, - channel: HostChannel, - reply: Option>>, -} - -impl RouteMachine -where - R::Error: From, -{ - pub fn new(execute: impl FnOnce(HostChannel) -> ExecuteFuture + Send + 'static) -> Self { - let (ops_tx, ops) = mpsc::unbounded_channel(); - Self { - execution: Execution::Unstarted(Box::new(execute)), - ops, - channel: HostChannel { ops: ops_tx }, - reply: None, - } - } - - async fn step( - &mut self, - result: Option>, - ) -> Result, R::Error> { - match (self.reply.take(), result) { - (Some(reply), Some(result)) => { - reply - .send(result) - .map_err(|_| MachineFault::Protocol("the call stopped waiting on the host"))?; - } - (None, None) if matches!(self.execution, Execution::Unstarted(_)) => {} - (Some(reply), None) => { - self.reply = Some(reply); - return Err(MachineFault::Protocol("host operation result is required").into()); - } - (None, Some(_)) => { - return Err(MachineFault::Protocol("unexpected host operation result").into()); - } - (None, None) => { - return Err( - MachineFault::Protocol("call cannot be resumed after completion").into(), - ); - } - } - if let Execution::Unstarted(_) = self.execution { - let Execution::Unstarted(start) = - std::mem::replace(&mut self.execution, Execution::Done) - else { - unreachable!() - }; - self.execution = Execution::Running(start(self.channel.clone())); - } - let Execution::Running(future) = &mut self.execution else { - return Err(MachineFault::Protocol("call cannot be resumed after completion").into()); - }; - tokio::select! { - biased; - pending = self.ops.recv() => { - let pending = pending.ok_or(MachineFault::Abandoned)?; - self.reply = Some(pending.reply); - Ok(MachineStep::Host(pending.op)) - } - outcome = future => { - self.execution = Execution::Done; - outcome.map(MachineStep::Complete) - } - } - } -} - -impl Machine for RouteMachine -where - R::Error: From, -{ - type Route = R; - type Complete = R::Response; - - fn resume(&mut self, result: Option>) -> Step<'_, Self> { - Box::pin(self.step(result)) - } - - fn interrupt(&mut self, failure: HostFailure) -> Interrupted<'_, Self> { - self.reply = None; - self.execution = Execution::Done; - Box::pin(async move { Err(failure.into_error()) }) - } -} diff --git a/litellm-rust/crates/host/src/protocol.rs b/litellm-rust/crates/host/src/protocol.rs new file mode 100644 index 00000000000..a7c0f3470b2 --- /dev/null +++ b/litellm-rust/crates/host/src/protocol.rs @@ -0,0 +1,17 @@ +/// One public call surface: what a completed call produces, how it fails, what the host +/// projects the caller's request into, and the protocol-specific operations only its host +/// can perform mid-call (token acquisition, for one). +pub trait Protocol: Send + Sync + 'static { + type Response: Send + 'static; + type Error: Clone + Send + Sync + 'static; + /// The caller's request as the host projects it, answered once before anything else. + type Projection: Send + 'static; + /// Each operation carries the [`Reply`](crate::host::Reply) its answer goes through. + /// A protocol with no operations of its own uses `Infallible`. + type Op: Send + 'static; + /// One piece of a streamed response, handed to the caller as it arrives. A protocol + /// that never streams uses `Infallible`. + type Chunk: Send + 'static; + /// What the call knows once a streamed response starts, before its first chunk. + type StreamHead: Send + 'static; +} diff --git a/litellm-rust/crates/host/src/route.rs b/litellm-rust/crates/host/src/route.rs deleted file mode 100644 index 8ab2b125760..00000000000 --- a/litellm-rust/crates/host/src/route.rs +++ /dev/null @@ -1,14 +0,0 @@ -/// One public call surface: what a completed call produces, how it fails, and the -/// route-specific operations only its host can perform (request projection, file reads, -/// token acquisition). -pub trait Route: Send + Sync + 'static { - type Response: Send + 'static; - type Error: Clone + Send + Sync + 'static; - type Op: Send + 'static; - type OpResult: Send + 'static; - /// One piece of a streamed response, handed to the caller as it arrives. A route - /// that never streams uses `Infallible`. - type Chunk: Send + 'static; - /// What the route knows once a streamed response starts, before its first chunk. - type StreamHead: Send + 'static; -} diff --git a/litellm-rust/crates/host/src/run.rs b/litellm-rust/crates/host/src/run.rs index 6a0c08fba68..baa3b58e058 100644 --- a/litellm-rust/crates/host/src/run.rs +++ b/litellm-rust/crates/host/src/run.rs @@ -1,40 +1,28 @@ use crate::event::{CallEvent, FailureOrigin, Timing, epoch_seconds}; -use crate::host::{Host, HostOp, HostResult}; +use crate::host::{Host, HostOp}; use crate::machine::{HostFailure, Machine, MachineStep}; -use crate::route::Route; +use crate::protocol::Protocol; /// Drives a machine to completion against an in-process host and emits exactly one /// terminal event. -pub async fn run(mut machine: M, host: &H) -> Result::Error> +pub async fn run( + mut machine: M, + host: &H, +) -> Result::Error> where M: Machine, - H: Host, + H: Host, { let start_time = epoch_seconds(); let _ = host.emit(&CallEvent::Started { start_time }).await; - let mut result = None; let outcome = loop { - let step = match machine.resume(result.take()).await { + let op = match machine.resume().await { Ok(MachineStep::Complete(complete)) => break Ok(complete), Ok(MachineStep::Host(op)) => op, Err(error) => break Err(error), }; - let answer = match step { - HostOp::Route(op) => host.route(op).await.map(HostResult::Route), - HostOp::BeforeSend { wire, context } => host - .before_send(*wire, &context) - .await - .map(|wire| HostResult::BeforeSend(Box::new(wire))), - HostOp::Emit(event) => host - .emit(&CallEvent::Machine(event)) - .await - .map(|()| HostResult::Emitted), - HostOp::Open(head) => host.open(head).await.map(HostResult::Demand), - HostOp::Deliver(chunk) => host.deliver(chunk).await.map(HostResult::Demand), - }; - match answer { - Ok(answer) => result = Some(answer), - Err(error) => break machine.interrupt(HostFailure::Error(error)).await, + if let Err(error) = perform(host, op).await { + break machine.interrupt(HostFailure::Error(error)).await; } }; let timing = Timing { @@ -52,44 +40,52 @@ where outcome } +async fn perform>(host: &H, op: HostOp) -> Result<(), R::Error> { + match op { + HostOp::Project(reply) => host + .project() + .await + .map(|projection| reply.send(projection)), + HostOp::Custom(op) => host.custom_op(op).await, + HostOp::BeforeSend { + wire, + context, + reply, + } => host + .before_send(*wire, &context) + .await + .map(|wire| reply.send(wire)), + HostOp::Emit(event, reply) => host + .emit(&CallEvent::Machine(event)) + .await + .map(|()| reply.send(())), + HostOp::Open(head, reply) => host.open(head).await.map(|demand| reply.send(demand)), + HostOp::Deliver(chunk, reply) => host.deliver(chunk).await.map(|demand| reply.send(demand)), + } +} + #[cfg(test)] mod tests { use std::sync::Mutex; use super::*; - use crate::machine::{Interrupted, Step}; + use crate::host::Reply; + use crate::machine::{CallMachine, MachineFault}; struct Unit; - impl Route for Unit { + impl Protocol for Unit { type Response = (); type Error = &'static str; - type Op = &'static str; - type OpResult = (); + type Projection = (); + type Op = (&'static str, Reply<()>); type Chunk = std::convert::Infallible; type StreamHead = std::convert::Infallible; } - struct Scripted { - ops: Vec<&'static str>, - outcome: Result<(), &'static str>, - } - - impl Machine for Scripted { - type Route = Unit; - type Complete = (); - - fn resume(&mut self, _: Option>) -> Step<'_, Self> { - Box::pin(async move { - if !self.ops.is_empty() { - return Ok(MachineStep::Host(HostOp::Route(self.ops.remove(0)))); - } - self.outcome.map(MachineStep::Complete) - }) - } - - fn interrupt(&mut self, failure: HostFailure<&'static str>) -> Interrupted<'_, Self> { - Box::pin(async move { Err(failure.into_error()) }) + impl From for &'static str { + fn from(_: MachineFault) -> Self { + "machine fault" } } @@ -100,12 +96,21 @@ mod tests { } impl Host for Recording { - async fn route(&self, op: &'static str) -> Result<(), &'static str> { - self.seen.lock().unwrap().push(format!("route:{op}")); - match self.fail { - Some(failing) if failing == op => Err("host failed"), - _ => Ok(()), + async fn project(&self) -> Result<(), &'static str> { + self.seen.lock().unwrap().push("project".into()); + Ok(()) + } + + async fn custom_op( + &self, + (op, reply): (&'static str, Reply<()>), + ) -> Result<(), &'static str> { + self.seen.lock().unwrap().push(format!("op:{op}")); + if self.fail == Some(op) { + return Err("host failed"); } + reply.send(()); + Ok(()) } async fn emit(&self, event: &CallEvent) -> Result<(), &'static str> { @@ -119,21 +124,29 @@ mod tests { } } - fn scripted(ops: &[&'static str], outcome: Result<(), &'static str>) -> Scripted { - Scripted { - ops: ops.to_vec(), - outcome, - } + fn scripted( + ops: &'static [&'static str], + outcome: Result<(), &'static str>, + ) -> CallMachine { + CallMachine::new(move |host| { + Box::pin(async move { + host.project().await?; + for op in ops { + host.custom_op(|reply| (*op, reply)).await?; + } + outcome + }) + }) } #[tokio::test] async fn forwards_every_op_then_emits_one_succeeded() { let host = Recording::default(); - let outcome = run(scripted(&["project", "send"], Ok(())), &host).await; + let outcome = run(scripted(&["sign", "send"], Ok(())), &host).await; assert_eq!(outcome, Ok(())); assert_eq!( *host.seen.lock().unwrap(), - ["started", "route:project", "route:send", "succeeded"] + ["started", "project", "op:sign", "op:send", "succeeded"] ); } @@ -142,24 +155,32 @@ mod tests { let host = Recording::default(); let outcome = run(scripted(&[], Err("boom")), &host).await; assert_eq!(outcome, Err("boom")); - assert_eq!(*host.seen.lock().unwrap(), ["started", "failed"]); + assert_eq!(*host.seen.lock().unwrap(), ["started", "project", "failed"]); let host = Recording { fail: Some("send"), ..Recording::default() }; - let outcome = run(scripted(&["project", "send", "never"], Ok(())), &host).await; + let outcome = run(scripted(&["sign", "send", "never"], Ok(())), &host).await; assert_eq!(outcome, Err("host failed")); assert_eq!( *host.seen.lock().unwrap(), - ["started", "route:project", "route:send", "failed"] + ["started", "project", "op:sign", "op:send", "failed"] ); } struct StartTimes(Mutex>); impl Host for StartTimes { - async fn route(&self, _: &'static str) -> Result<(), &'static str> { + async fn project(&self) -> Result<(), &'static str> { + Ok(()) + } + + async fn custom_op( + &self, + (_, reply): (&'static str, Reply<()>), + ) -> Result<(), &'static str> { + reply.send(()); Ok(()) } @@ -178,7 +199,7 @@ mod tests { #[tokio::test] async fn started_opens_the_call_at_the_terminal_start_time_and_cannot_fail_it() { let host = StartTimes(Mutex::default()); - assert_eq!(run(scripted(&["project"], Ok(())), &host).await, Ok(())); + assert_eq!(run(scripted(&["send"], Ok(())), &host).await, Ok(())); let times = host.0.lock().unwrap(); assert_eq!(times.len(), 2); assert_eq!(times[0], times[1]); diff --git a/litellm-rust/crates/llms/src/base_llm/ocr/error.rs b/litellm-rust/crates/llms/src/base_llm/ocr/error.rs index b3df8fc18c8..5d7be8b10df 100644 --- a/litellm-rust/crates/llms/src/base_llm/ocr/error.rs +++ b/litellm-rust/crates/llms/src/base_llm/ocr/error.rs @@ -118,7 +118,6 @@ impl From for Error { Self::InvalidRequest(match fault { MachineFault::Abandoned => "OCR host driver was abandoned".into(), MachineFault::Protocol(message) => format!("OCR {message}"), - MachineFault::Mismatch => "invalid OCR host operation result".into(), }) } } diff --git a/litellm-rust/crates/python-bridge/src/logger/machine.rs b/litellm-rust/crates/python-bridge/src/logger/machine.rs index 7234308e67e..54f8f3d4b3f 100644 --- a/litellm-rust/crates/python-bridge/src/logger/machine.rs +++ b/litellm-rust/crates/python-bridge/src/logger/machine.rs @@ -1,9 +1,8 @@ use std::sync::OnceLock; use litellm_host::{ - host::HostResult, machine::{HostFailure, Interrupted, Machine, Step}, - route::Route, + protocol::Protocol, }; use litellm_tracing::Logger; use pyo3::Python; @@ -23,17 +22,17 @@ impl LoggedMachine { } impl Machine for LoggedMachine { - type Route = M::Route; + type Protocol = M::Protocol; type Complete = M::Complete; - fn resume(&mut self, result: Option>) -> Step<'_, Self> { + fn resume(&mut self) -> Step<'_, Self> { let logger = self.logger.get_or_init(|| Python::attach(super::capture)); - Box::pin(logger.instrument(logger.scope(|| self.machine.resume(result)))) + Box::pin(logger.instrument(logger.scope(|| self.machine.resume()))) } fn interrupt( &mut self, - failure: HostFailure<::Error>, + failure: HostFailure<::Error>, ) -> Interrupted<'_, Self> { let logger = self.logger.get_or_init(|| Python::attach(super::capture)); Box::pin(logger.instrument(logger.scope(|| self.machine.interrupt(failure)))) diff --git a/litellm-rust/crates/python-bridge/src/logger/tests.rs b/litellm-rust/crates/python-bridge/src/logger/tests.rs index 9312d4c187c..b65e7d37023 100644 --- a/litellm-rust/crates/python-bridge/src/logger/tests.rs +++ b/litellm-rust/crates/python-bridge/src/logger/tests.rs @@ -1,29 +1,28 @@ use std::{process::Command, task::Poll}; use litellm_host::{ - host::HostResult, machine::{HostFailure, Interrupted, Machine, MachineStep, Step}, - route::Route, + protocol::Protocol, }; use pyo3::{prelude::*, types::PyDict}; struct DiagnosticMachine; -impl Route for DiagnosticMachine { +impl Protocol for DiagnosticMachine { type Response = (); type Error = String; + type Projection = (); type Op = (); - type OpResult = (); type Chunk = (); type StreamHead = (); } impl Machine for DiagnosticMachine { - type Route = Self; + type Protocol = Self; type Complete = (); - fn resume(&mut self, _: Option>) -> Step<'_, Self> { + fn resume(&mut self) -> Step<'_, Self> { litellm_tracing::warn!("machine started"); Box::pin(async { tokio::task::yield_now().await; @@ -45,7 +44,7 @@ fn machine_warning(py: Python<'_>) -> PyResult> { let mut machine = super::LoggedMachine::new(DiagnosticMachine); let mut future = Box::pin(async move { machine - .resume(None) + .resume() .await .map_err(pyo3::exceptions::PyValueError::new_err)?; machine diff --git a/litellm-rust/crates/python-bridge/src/routes/messages/host.rs b/litellm-rust/crates/python-bridge/src/routes/messages/host.rs index 9d97094aeda..6de4e1320e1 100644 --- a/litellm-rust/crates/python-bridge/src/routes/messages/host.rs +++ b/litellm-rust/crates/python-bridge/src/routes/messages/host.rs @@ -1,10 +1,12 @@ +use std::convert::Infallible; + use bytes::Bytes; use litellm_core::messages::{ Error, - route::{Messages, MessagesCall, MessagesOp, MessagesOpResult, MessagesOutput}, + route::{Messages, MessagesCall, MessagesOutput}, types::MessagesShaping, }; -use litellm_host_python::{InvokeError, RouteHost, from_py, lookup, to_py}; +use litellm_host_python::{InvokeError, ProtocolHost, from_py, lookup, to_py}; use litellm_http::transport::Error as TransportError; use litellm_types::utils::ProviderSpecificHeaders; use pyo3::{ @@ -80,16 +82,16 @@ fn native_error(py: Python<'_>, error: Error) -> PyResult { /// The Python side of the Messages route: projects the prepared arguments and builds the /// public response, chunks and exceptions. -pub(super) struct MessagesRouteHost { +pub(super) struct MessagesPythonHost { request: Py, } -impl MessagesRouteHost { +impl MessagesPythonHost { pub(super) fn new(request: Py) -> Self { Self { request } } - fn project(&self, py: Python<'_>, arguments: &Bound<'_, PyDict>) -> PyResult { + fn projection(&self, py: Python<'_>, arguments: &Bound<'_, PyDict>) -> PyResult { let request = self.request.bind(py); let argument = |name: &str| -> PyResult>> { Ok(lookup(arguments, request, name)?.filter(|value| !value.is_none())) @@ -208,22 +210,21 @@ impl MessagesRouteHost { } } -impl RouteHost for MessagesRouteHost { - type Route = Messages; +impl ProtocolHost for MessagesPythonHost { + type Protocol = Messages; type Failure = PyErr; - fn invoke( + fn project( &mut self, py: Python<'_>, arguments: &Bound<'_, PyDict>, - op: MessagesOp, - ) -> Result> { - match op { - MessagesOp::ProjectRequest => self - .project(py, arguments) - .map(|call| MessagesOpResult::Request(Box::new(call))) - .map_err(|error| InvokeError::Python(self.map_failure(py, error))), - } + ) -> Result> { + self.projection(py, arguments) + .map_err(|error| InvokeError::Python(self.map_failure(py, error))) + } + + fn invoke(&mut self, _: Python<'_>, op: Infallible) -> Result<(), InvokeError> { + match op {} } fn complete(&mut self, py: Python<'_>, response: MessagesOutput) -> PyResult> { diff --git a/litellm-rust/crates/python-bridge/src/routes/messages/mod.rs b/litellm-rust/crates/python-bridge/src/routes/messages/mod.rs index fd474e6b2d4..65040f31684 100644 --- a/litellm-rust/crates/python-bridge/src/routes/messages/mod.rs +++ b/litellm-rust/crates/python-bridge/src/routes/messages/mod.rs @@ -1,6 +1,6 @@ mod host; -use host::MessagesRouteHost; +use host::MessagesPythonHost; use litellm_callbacks_legacy_python::{ LegacySurface, PassThroughStream, PublicCall, run_legacy_call, }; @@ -45,7 +45,7 @@ fn run_messages( SURFACE, PublicCall::capture(&request, &args, &kwargs)?, crate::logger::LoggedMachine::new(messages_machine(secrets)), - MessagesRouteHost::new(request.unbind()), + MessagesPythonHost::new(request.unbind()), asynchronous, ) } diff --git a/litellm-rust/crates/python-bridge/src/routes/ocr/document.rs b/litellm-rust/crates/python-bridge/src/routes/ocr/document.rs index a928e62d5b7..e6821241c89 100644 --- a/litellm-rust/crates/python-bridge/src/routes/ocr/document.rs +++ b/litellm-rust/crates/python-bridge/src/routes/ocr/document.rs @@ -1,58 +1,38 @@ use std::path::PathBuf; -use bytes::Bytes; -use litellm_core::ocr::types::{OcrDocumentInput, OcrFileContent}; +use litellm_core::ocr::types::OcrDocumentInput; +use litellm_host_python::{PythonFileReader, py_bytes}; use pyo3::{ - exceptions::{PyTypeError, PyValueError}, - gc::{PyTraverseError, PyVisit}, + exceptions::PyValueError, prelude::*, - pybacked::PyBackedBytes, sync::PyOnceLock, types::{PyBytes, PyString, PyType}, }; -#[derive(Debug)] -pub(super) struct PythonFileReader { - reader: Py, - name: Option, +/// A `type='file'` document as projected: paths and bytes are typed inputs already; a +/// file-like object is a reader the projection consumes once every other field is read. +pub(super) enum FileDocumentInput { + Ready(OcrDocumentInput), + Deferred { + reader: PythonFileReader, + mime_type: Option, + }, } -impl PythonFileReader { - pub(super) fn read(&self, py: Python<'_>) -> PyResult { - let value = self.reader.bind(py).call0()?; - let bytes = if value.is_instance_of::() { - Bytes::from(value.extract::()?) - } else if value.is_instance_of::() { - extract_bytes(&value)? - } else { - return Err(PyTypeError::new_err(format!( - "OCR file read must return bytes or str, got {}", - value.get_type(), - ))); - }; - Ok(OcrFileContent { - bytes, - file_name: self.name.clone(), - }) +impl FileDocumentInput { + pub(super) fn resolve(self, py: Python<'_>) -> PyResult { + match self { + Self::Ready(input) => Ok(input), + Self::Deferred { reader, mime_type } => { + let content = reader.read(py)?; + Ok(OcrDocumentInput::Bytes { + bytes: content.bytes, + file_name: content.file_name, + mime_type, + }) + } + } } - - pub(super) fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError> { - visit.call(&self.reader) - } -} - -fn extract_bytes(value: &Bound<'_, PyAny>) -> PyResult { - if value.is_exact_instance_of::() { - return Ok(Bytes::from_owner(value.extract::()?)); - } - Ok(Bytes::copy_from_slice( - value.extract::()?.as_ref(), - )) -} - -pub(super) struct FileDocumentInput { - pub input: OcrDocumentInput, - pub reader: Option, } impl FromPyObject<'_, '_> for FileDocumentInput { @@ -87,51 +67,31 @@ impl FromPyObject<'_, '_> for FileDocumentInput { } static PATH_LIKE: PyOnceLock> = PyOnceLock::new(); if file.is_instance(PATH_LIKE.import(py, "os", "PathLike")?)? { - return Ok(Self { - input: OcrDocumentInput::Path { - path: file.extract::()?, - mime_type, - }, - reader: None, - }); + return Ok(Self::Ready(OcrDocumentInput::Path { + path: file.extract::()?, + mime_type, + })); } if file.is_instance_of::() { - return Ok(Self { - input: OcrDocumentInput::Bytes { - bytes: extract_bytes(&file)?, - file_name: None, - mime_type, - }, - reader: None, - }); + return Ok(Self::Ready(OcrDocumentInput::Bytes { + bytes: py_bytes(&file)?, + file_name: None, + mime_type, + })); } - let reader = file - .getattr_opt("read")? - .filter(|value| value.is_callable()); - let Some(reader) = reader else { - return Err(PyValueError::new_err(format!( + match PythonFileReader::from_file_like(&file)? { + Some(reader) => Ok(Self::Deferred { reader, mime_type }), + None => Err(PyValueError::new_err(format!( "Unsupported file input type: {}. Expected pathlib.Path, bytes, or a file-like object.", file.get_type(), - ))); - }; - let name = file - .getattr_opt("name")? - .filter(|value| !value.is_none()) - .map(|value| value.extract::()) - .transpose()?; - Ok(Self { - input: OcrDocumentInput::HostReader { mime_type }, - reader: Some(PythonFileReader { - reader: reader.unbind(), - name, - }), - }) + ))), + } } } #[cfg(test)] mod tests { - use pyo3::types::PyDict; + use pyo3::{exceptions::PyTypeError, types::PyDict}; use super::*; @@ -141,6 +101,13 @@ mod tests { locals } + fn ready(input: FileDocumentInput) -> OcrDocumentInput { + match input { + FileDocumentInput::Ready(input) => input, + FileDocumentInput::Deferred { .. } => panic!("expected a ready document"), + } + } + #[test] fn extraction_validates_required_file_and_optional_mime_type() { Python::initialize(); @@ -167,13 +134,19 @@ mod tests { .unwrap(); assert!(error.is_instance_of::(py)); assert!(error.to_string().contains("bare str")); + let error = py + .eval(c"{'file': object()}", None, None) + .unwrap() + .extract::() + .err() + .unwrap(); + assert!(error.is_instance_of::(py)); + assert!(error.to_string().contains("Unsupported file input type")); let document = py .eval(c"{'file': b'abc', 'mime_type': 'image/png'}", None, None) .unwrap(); - let input: FileDocumentInput = document.extract().unwrap(); - assert!(input.reader.is_none()); assert_eq!( - input.input, + ready(document.extract().unwrap()), OcrDocumentInput::Bytes { bytes: b"abc".as_slice().into(), file_name: None, @@ -199,7 +172,7 @@ class Reader: return b'abc' reader = Reader() document = {'file': reader, 'mime_type': 7} -reader_document = {'file': reader} +reader_document = {'file': reader, 'mime_type': 'application/pdf'} path_document = {'file': Path('/nonexistent/ocr-projection-test.pdf'), 'mime_type': 'image/png'}", ); let document = locals.get_item("document").unwrap().unwrap(); @@ -208,10 +181,6 @@ path_document = {'file': Path('/nonexistent/ocr-projection-test.pdf'), 'mime_typ let document = locals.get_item("reader_document").unwrap().unwrap(); let input: FileDocumentInput = document.extract().unwrap(); - assert_eq!( - input.input, - OcrDocumentInput::HostReader { mime_type: None } - ); let reads = || { locals .get_item("reader") @@ -223,21 +192,20 @@ path_document = {'file': Path('/nonexistent/ocr-projection-test.pdf'), 'mime_typ .unwrap() }; assert_eq!(reads(), 0); - let content = input.reader.unwrap().read(py).unwrap(); + let resolved = input.resolve(py).unwrap(); assert_eq!(reads(), 1); assert_eq!( - content, - OcrFileContent { + resolved, + OcrDocumentInput::Bytes { bytes: b"abc".as_slice().into(), file_name: Some("scan.png".into()), + mime_type: Some("application/pdf".into()), } ); let document = locals.get_item("path_document").unwrap().unwrap(); - let input: FileDocumentInput = document.extract().unwrap(); - assert!(input.reader.is_none()); assert_eq!( - input.input, + ready(document.extract().unwrap()), OcrDocumentInput::Path { path: PathBuf::from("/nonexistent/ocr-projection-test.pdf"), mime_type: Some("image/png".into()), @@ -245,97 +213,4 @@ path_document = {'file': Path('/nonexistent/ocr-projection-test.pdf'), 'mime_typ ); }); } - - #[test] - fn reader_results_are_normalized_and_exceptions_keep_their_identity() { - Python::initialize(); - Python::attach(|py| { - let locals = eval( - py, - c"failure = KeyError('reader failed') -class Raising: - def read(self): - raise failure -class Text: - def read(self): - return 'héllo' -class Wrong: - def read(self): - return 7 -raising = {'file': Raising()} -text = {'file': Text()} -wrong = {'file': Wrong()}", - ); - let reader = |name: &str| { - locals - .get_item(name) - .unwrap() - .unwrap() - .extract::() - .unwrap() - .reader - .unwrap() - }; - let error = reader("raising").read(py).unwrap_err(); - assert!( - error - .value(py) - .is(locals.get_item("failure").unwrap().unwrap()) - ); - assert_eq!( - reader("text").read(py).unwrap().bytes.as_ref(), - "héllo".as_bytes() - ); - let error = reader("wrong").read(py).unwrap_err(); - assert!(error.is_instance_of::(py)); - assert!(error.to_string().contains("bytes or str")); - }); - } - - #[rstest::rstest] - #[case::read("read")] - #[case::name("name")] - fn reader_attribute_failures_keep_their_identity(#[case] attribute: &str) { - Python::initialize(); - Python::attach(|py| { - let locals = eval( - py, - c"failure = LookupError('file property failed') -class File: - def __getattribute__(self, name): - if name == attribute: - raise failure - return super().__getattribute__(name) - name = 'scan.pdf' - def read(self): - return b'abc' -document = {'file': File()}", - ); - locals.set_item("attribute", attribute).unwrap(); - let error = locals - .get_item("document") - .unwrap() - .unwrap() - .extract::() - .err() - .unwrap(); - assert!( - error - .value(py) - .is(locals.get_item("failure").unwrap().unwrap()) - ); - }); - } - - #[test] - fn exact_python_bytes_transfer_without_copying_and_outlive_the_input() { - Python::initialize(); - let (bytes, pointer) = Python::attach(|py| { - let value = PyBytes::new(py, b"document bytes"); - let pointer = value.as_bytes().as_ptr() as usize; - (extract_bytes(value.as_any()).unwrap(), pointer) - }); - assert_eq!(bytes.as_ptr() as usize, pointer); - assert_eq!(bytes.as_ref(), b"document bytes"); - } } diff --git a/litellm-rust/crates/python-bridge/src/routes/ocr/host.rs b/litellm-rust/crates/python-bridge/src/routes/ocr/host.rs index 8bf99cd355f..5a3806e61e3 100644 --- a/litellm-rust/crates/python-bridge/src/routes/ocr/host.rs +++ b/litellm-rust/crates/python-bridge/src/routes/ocr/host.rs @@ -1,6 +1,6 @@ use litellm_auth::ResolvedCredential; -use litellm_core::ocr::route::{Ocr, OcrOp, OcrOpResult}; -use litellm_host_python::{InvokeError, RouteHost, missing_state, to_py}; +use litellm_core::ocr::route::{Ocr, OcrOp, OcrProjection}; +use litellm_host_python::{InvokeError, ProtocolHost, missing_state, to_py}; use litellm_llms::base_llm::ocr::{error::Error, transformation::LiteLLMOcrResponse}; use pyo3::{ exceptions::{PyBaseException, PyException}, @@ -20,14 +20,15 @@ enum OcrHostData { Released, } -/// The Python side of the OCR route: projects the prepared arguments, reads file-like -/// documents, acquires Azure AD tokens, and builds the public response and exception. -pub(super) struct OcrRouteHost { +/// The Python side of the OCR route: projects the prepared arguments (reading a file-like +/// document as it goes), acquires Azure AD tokens, and builds the public response and +/// exception. +pub(super) struct OcrPythonHost { request: Py, data: OcrHostData, } -impl OcrRouteHost { +impl OcrPythonHost { pub(super) fn new(request: Py) -> Self { Self { request, @@ -42,14 +43,6 @@ impl OcrRouteHost { } } - fn read_document(&self, py: Python<'_>) -> PyResult { - self.handles()? - .reader - .as_ref() - .ok_or_else(missing_state)? - .read(py) - } - fn acquire_azure_ad_token(&self, py: Python<'_>) -> PyResult { self.handles()? .azure_ad_token_provider @@ -58,30 +51,21 @@ impl OcrRouteHost { .acquire(py) } - fn answer( + fn projection( &mut self, py: Python<'_>, arguments: &Bound<'_, PyDict>, - op: OcrOp, - ) -> PyResult { - match op { - OcrOp::ProjectRequest => { - let OcrHostData::Unprojected = self.data else { - return Err(missing_state()); - }; - let (request, handles) = project_request(self.request.bind(py), arguments)?; - let caller_token = handles.azure_ad_token_provider.is_some(); - self.data = OcrHostData::Projected(Box::new(handles)); - Ok(OcrOpResult::Request { - request: Box::new(request), - caller_token, - }) - } - OcrOp::ReadDocument => self.read_document(py).map(OcrOpResult::Document), - OcrOp::AcquireAzureAdToken => self - .acquire_azure_ad_token(py) - .map(OcrOpResult::AzureAdToken), - } + ) -> PyResult { + let OcrHostData::Unprojected = self.data else { + return Err(missing_state()); + }; + let (request, handles) = project_request(self.request.bind(py), arguments)?; + let caller_token = handles.azure_ad_token_provider.is_some(); + self.data = OcrHostData::Projected(Box::new(handles)); + Ok(OcrProjection { + request, + caller_token, + }) } fn map_failure(&self, py: Python<'_>, error: PyErr) -> PyErr { @@ -104,20 +88,28 @@ impl OcrRouteHost { } } -impl RouteHost for OcrRouteHost { - type Route = Ocr; +impl ProtocolHost for OcrPythonHost { + type Protocol = Ocr; type Failure = PyErr; - fn invoke( + fn project( &mut self, py: Python<'_>, arguments: &Bound<'_, PyDict>, - op: OcrOp, - ) -> Result> { - self.answer(py, arguments, op) + ) -> Result> { + self.projection(py, arguments) .map_err(|error| InvokeError::Python(self.map_failure(py, error))) } + fn invoke(&mut self, py: Python<'_>, op: OcrOp) -> Result<(), InvokeError> { + match op { + OcrOp::AcquireAzureAdToken(reply) => self + .acquire_azure_ad_token(py) + .map(|token| reply.send(token)) + .map_err(|error| InvokeError::Python(self.map_failure(py, error))), + } + } + fn complete(&mut self, py: Python<'_>, response: LiteLLMOcrResponse) -> PyResult> { py.import("litellm.rust_bridge.ocr.route_host")? .getattr("response")? @@ -148,13 +140,10 @@ impl RouteHost for OcrRouteHost { fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError> { visit.call(&self.request)?; - if let OcrHostData::Projected(handles) = &self.data { - if let Some(reader) = &handles.reader { - reader.traverse(visit)?; - } - if let Some(provider) = &handles.azure_ad_token_provider { - provider.traverse(visit)?; - } + if let OcrHostData::Projected(handles) = &self.data + && let Some(provider) = &handles.azure_ad_token_provider + { + provider.traverse(visit)?; } Ok(()) } @@ -205,20 +194,13 @@ del provider .unwrap() .cast_into::() .unwrap(); - let mut host = OcrRouteHost::new(py.None()); - let projected = host.invoke(py, &kwargs, OcrOp::ProjectRequest).unwrap(); - assert!(matches!( - projected, - OcrOpResult::Request { - caller_token: true, - .. - } - )); + let mut host = OcrPythonHost::new(py.None()); + assert!(host.project(py, &kwargs).unwrap().caller_token); locals.del_item("kwargs").unwrap(); drop(kwargs); + let (reply, _) = litellm_host::host::reply(); assert_eq!( - host.invoke(py, &PyDict::new(py), OcrOp::AcquireAzureAdToken) - .is_ok(), + host.invoke(py, OcrOp::AcquireAzureAdToken(reply)).is_ok(), succeeds ); let alive = || { diff --git a/litellm-rust/crates/python-bridge/src/routes/ocr/mod.rs b/litellm-rust/crates/python-bridge/src/routes/ocr/mod.rs index b54316b258b..a4f2bf851d7 100644 --- a/litellm-rust/crates/python-bridge/src/routes/ocr/mod.rs +++ b/litellm-rust/crates/python-bridge/src/routes/ocr/mod.rs @@ -5,7 +5,7 @@ mod project; use std::sync::LazyLock; -use host::OcrRouteHost; +use host::OcrPythonHost; use litellm_auth_gcp::VertexAuth; use litellm_callbacks_legacy_python::{LegacySurface, PublicCall, run_legacy_call}; use litellm_core::ocr::{provider_config, route::ocr_machine}; @@ -69,7 +69,7 @@ fn run_ocr( if asynchronous { ASYNC_SURFACE } else { SURFACE }, PublicCall::capture(&request, &args, &kwargs)?, crate::logger::LoggedMachine::new(ocr_machine(client)), - OcrRouteHost::new(request.unbind()), + OcrPythonHost::new(request.unbind()), asynchronous, ) } diff --git a/litellm-rust/crates/python-bridge/src/routes/ocr/project.rs b/litellm-rust/crates/python-bridge/src/routes/ocr/project.rs index 697b935a1d4..be43a1b7711 100644 --- a/litellm-rust/crates/python-bridge/src/routes/ocr/project.rs +++ b/litellm-rust/crates/python-bridge/src/routes/ocr/project.rs @@ -8,19 +8,15 @@ use litellm_llms::base_llm::ocr::error::Error; use pyo3::{exceptions::PyValueError, prelude::*, types::PyDict}; use serde_json::{Map, Value}; -use super::{ - document::{FileDocumentInput, PythonFileReader}, - errors::to_pyerr as ocr_error_to_pyerr, -}; +use super::{document::FileDocumentInput, errors::to_pyerr as ocr_error_to_pyerr}; use crate::{ credentials::{self, CallerTokenProvider}, marshal::{project_optional_fields, python_timeout_seconds, request_input_sources}, }; -/// What the host keeps after projection: the caller's callables that answer the document -/// read and token operations, and the provider name the failure mapping reports. +/// What the host keeps after projection: the caller's token callable that answers the +/// token operation, and the provider name the failure mapping reports. pub(super) struct OcrHostHandles { - pub reader: Option, pub azure_ad_token_provider: Option, pub provider: &'static str, } @@ -104,13 +100,11 @@ impl ProjectedDocument { Ok(Self::File(document.extract()?)) } - fn into_parts(self) -> PyResult<(OcrDocumentInput, Option)> { + /// Reads a file-like document now, so it runs after every other argument was read. + fn resolve(self, py: Python<'_>) -> PyResult { match self { - Self::File(FileDocumentInput { input, reader }) => Ok((input, reader)), - Self::Other(wire) => Ok(( - decode_document(wire).map_err(ocr_error_to_pyerr)?.into(), - None, - )), + Self::File(file) => file.resolve(py), + Self::Other(wire) => Ok(decode_document(wire).map_err(ocr_error_to_pyerr)?.into()), } } } @@ -136,24 +130,25 @@ pub(super) fn project_request( .chain(["api_key", "api_base", "extra_headers"]), )?; let azure_ad_token_provider = credentials::azure_ad_token_provider(kwargs)?; - let (document, reader) = document.into_parts()?; + let api_base = arguments.api_base()?; + let extra_headers = arguments.extra_headers()?; + let timeout_seconds = arguments.timeout_seconds()?; let wire = OcrWireRequest { model, - document, + document: document.resolve(request.py())?, api_key, - api_base: arguments.api_base()?, + api_base, custom_llm_provider, - extra_headers: arguments.extra_headers()?, + extra_headers, optional_params, input_sources, - timeout_seconds: arguments.timeout_seconds()?, + timeout_seconds, }; let request = decode_request_input(wire).map_err(ocr_error_to_pyerr)?; let provider = request.provider_name(); Ok(( request, OcrHostHandles { - reader, azure_ad_token_provider, provider, }, @@ -180,10 +175,8 @@ mod tests { OcrArguments { request, kwargs } } - fn project_document( - document: &Bound<'_, PyAny>, - ) -> PyResult<(OcrDocumentInput, Option)> { - ProjectedDocument::project(document)?.into_parts() + fn project_document(document: &Bound<'_, PyAny>) -> PyResult { + ProjectedDocument::project(document)?.resolve(document.py()) } fn url_document(url: &str) -> OcrDocumentInput { @@ -342,8 +335,11 @@ kwargs = {} }); } + /// A reader that rewrites the request while it runs shows which arguments projection + /// read before it and which after: every other argument is read first, and the read + /// happens exactly once. #[test] - fn document_readers_are_not_consumed_during_projection() { + fn document_readers_are_read_once_after_every_other_argument() { Python::initialize(); Python::attach(|py| { stub_timeout_conversion(py); @@ -351,17 +347,24 @@ kwargs = {} py, c" class Request: - api_base = 'original' + model = 'mistral/mistral-ocr-latest' + custom_llm_provider = None + api_key = None + api_base = 'https://original.example.com' + extra_headers = {'x-source': 'original'} timeout = 1 @property def document(self): return document class Reader: + reads = 0 def read(self): - Request.api_base = 'mutated' + Reader.reads += 1 + Request.api_base = 'https://mutated.example.com' + Request.extra_headers = {'x-source': 'mutated'} Request.timeout = 9 return b'abc' -document = {'type': 'file', 'file': Reader()} +document = {'type': 'file', 'file': Reader(), 'mime_type': 'application/pdf'} request = Request() kwargs = {} ", @@ -373,15 +376,38 @@ kwargs = {} .unwrap() .cast_into::() .unwrap(); - let arguments = arguments(&request, &kwargs); - let document = arguments.document().unwrap(); - let (input, reader) = project_document(&document).unwrap(); - assert_eq!(input, OcrDocumentInput::HostReader { mime_type: None }); - assert_eq!(arguments.api_base().unwrap().as_deref(), Some("original")); - assert_eq!(arguments.timeout_seconds().unwrap(), Some(1.0)); - reader.unwrap().read(py).unwrap(); - assert_eq!(arguments.api_base().unwrap().as_deref(), Some("mutated")); - assert_eq!(arguments.timeout_seconds().unwrap(), Some(9.0)); + let (projected, _) = project_request(&request, &kwargs).unwrap(); + assert_eq!( + py.eval(c"Reader.reads", Some(&locals), Some(&locals)) + .unwrap() + .extract::() + .unwrap(), + 1 + ); + assert_eq!( + projected.document, + OcrDocumentInput::Bytes { + bytes: b"abc".as_slice().into(), + file_name: None, + mime_type: Some("application/pdf".into()), + } + ); + assert_eq!( + projected + .credentials + .api_base + .as_ref() + .map(|base| base.value().as_str()), + Some("https://original.example.com") + ); + assert_eq!( + projected.transport.extra_headers, + [("x-source".to_string(), "original".to_string())] + ); + assert_eq!( + projected.transport.timeout, + Some(std::time::Duration::from_secs(1)) + ); }); } @@ -396,16 +422,14 @@ kwargs = {} None, ) .unwrap(); - let (input, reader) = project_document(&file).unwrap(); assert_eq!( - input, + project_document(&file).unwrap(), OcrDocumentInput::Bytes { bytes: b"%PDF-1.4".as_slice().into(), file_name: None, mime_type: Some("application/pdf".into()), } ); - assert!(reader.is_none()); let original = py .eval( @@ -414,8 +438,10 @@ kwargs = {} None, ) .unwrap(); - let (input, _) = project_document(&original).unwrap(); - assert_eq!(input, url_document("https://example.com/a.pdf")); + assert_eq!( + project_document(&original).unwrap(), + url_document("https://example.com/a.pdf") + ); }); } @@ -617,7 +643,7 @@ document = Document() ", ); let document = locals.get_item("document").unwrap().unwrap(); - let (input, _) = project_document(&document).unwrap(); + let input = project_document(&document).unwrap(); assert!(matches!(input, OcrDocumentInput::Bytes { .. })); let reads: Vec = document.getattr("reads").unwrap().extract().unwrap(); assert_eq!(reads, ["type", "mime_type", "file"]); From efbb3ac87e9a6fe12b0356e1670523543cf5a12b Mon Sep 17 00:00:00 2001 From: "berriai-litellm-provider-info-sync[bot]" <328147090+berriai-litellm-provider-info-sync[bot]@users.noreply.github.com> Date: Thu, 24 Sep 2026 21:25:23 -0700 Subject: [PATCH 09/10] chore(cost-map): add together-ai deprecation dates for gpt-oss-20b and gemma-4-31B-it (#43127) Price-Sync: litellm-providers Co-authored-by: berriai-litellm-provider-info-sync[bot] <328147090+berriai-litellm-provider-info-sync[bot]@users.noreply.github.com> --- litellm/model_prices_and_context_window_backup.json | 2 ++ model_prices_and_context_window.json | 2 ++ 2 files changed, 4 insertions(+) diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index efc0e0e2877..f875936b0bf 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -69078,6 +69078,7 @@ "source": "https://api.together.ai/v1/models" }, "together_ai/google/gemma-4-31B-it": { + "deprecation_date": "2026-09-14", "input_cost_per_token": 3.9e-07, "litellm_provider": "together_ai", "max_input_tokens": 262144, @@ -69207,6 +69208,7 @@ "source": "https://api.together.ai/v1/models" }, "together_ai/openai/gpt-oss-20b": { + "deprecation_date": "2026-09-14", "input_cost_per_token": 5e-08, "litellm_provider": "together_ai", "max_input_tokens": 131072, diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index efc0e0e2877..f875936b0bf 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -69078,6 +69078,7 @@ "source": "https://api.together.ai/v1/models" }, "together_ai/google/gemma-4-31B-it": { + "deprecation_date": "2026-09-14", "input_cost_per_token": 3.9e-07, "litellm_provider": "together_ai", "max_input_tokens": 262144, @@ -69207,6 +69208,7 @@ "source": "https://api.together.ai/v1/models" }, "together_ai/openai/gpt-oss-20b": { + "deprecation_date": "2026-09-14", "input_cost_per_token": 5e-08, "litellm_provider": "together_ai", "max_input_tokens": 131072, From c976c16a82244807e4e0355e92c453a007ae37a9 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Thu, 24 Sep 2026 21:34:57 -0700 Subject: [PATCH 10/10] feat(mcp): allow ["*"] wildcard in mcp_tool_permissions to grant all current and future tools (#43108) * feat(mcp): allow ["*"] wildcard in mcp_tool_permissions to grant all current and future tools Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * style(ui): run prettier on MCPToolPermissions files Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(mcp): keep ["*"] wildcard through toolset union and move constant to litellm.constants Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * style(mcp): format user_api_key_auth_mcp with ruff format Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(mcp): restore wildcard ceiling and deny-all regression tests Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(mcp): treat an empty team tool list as deny-all regardless of key grants Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * revert(mcp): keep legacy [] merge semantics, the truthiness check predates this PR Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(mcp): drop banner comment that repeats the wildcard test docstring Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: joshua Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/constants.py | 3 + .../mcp_server/auth/user_api_key_auth_mcp.py | 13 +- .../mcp_server/mcp_server_manager.py | 25 +-- .../auth/test_user_api_key_auth_mcp.py | 144 ++++++++++++++++++ .../mcp_server/test_mcp_server_manager.py | 29 ++++ .../MCPToolPermissions.test.tsx | 80 +++++++++- .../MCPToolPermissions.tsx | 18 ++- .../effectiveMcpServers.test.ts | 11 ++ .../effectiveMcpServers.ts | 7 + .../src/components/mcp_tools/constants.ts | 3 + 10 files changed, 312 insertions(+), 21 deletions(-) diff --git a/litellm/constants.py b/litellm/constants.py index a86be55d654..67021ae2abc 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -167,6 +167,9 @@ MCP_OAUTH2_TOKEN_CACHE_MAX_SIZE: Final = int(os.getenv("MCP_OAUTH2_TOKEN_CACHE_M MCP_OAUTH2_TOKEN_CACHE_DEFAULT_TTL: Final = int(os.getenv("MCP_OAUTH2_TOKEN_CACHE_DEFAULT_TTL", "3600")) MCP_SSO_ASSERTION_CACHE_TTL_SECONDS: Final = int(os.getenv("MCP_SSO_ASSERTION_CACHE_TTL_SECONDS", "60")) +# mcp_tool_permissions entry that grants every current and future tool on a server +MCP_ALL_TOOLS_WILDCARD: Final = "*" + # Default npm cache directory for STDIO MCP servers. # npm/npx needs a writable cache dir; in containers the default (~/.npm) # may not exist or be read-only. /tmp is always writable. diff --git a/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py b/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py index cdf52e6dc8d..a93ffaeac9f 100644 --- a/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py +++ b/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py @@ -13,6 +13,7 @@ from typing_extensions import assert_never import litellm from litellm._logging import verbose_logger +from litellm.constants import MCP_ALL_TOOLS_WILDCARD from litellm.proxy._experimental.mcp_server.oauth_utils import ( get_passthrough_resource_metadata_url, get_passthrough_www_authenticate, @@ -2136,7 +2137,11 @@ class MCPRequestHandler: via_toolsets: Sequence[str] | None, ) -> Sequence[str] | None: """Union of one level's direct tool grants and its toolset-granted tools on one server, - ``None`` when neither source restricts (allow-all from this level).""" + ``None`` when neither source restricts (allow-all from this level). A direct grant + containing ``MCP_ALL_TOOLS_WILDCARD`` makes the level unrestricted, so it returns + ``None`` whatever the toolsets name.""" + if direct is not None and MCP_ALL_TOOLS_WILDCARD in direct: + return None if direct is None and via_toolsets is None: return None return tuple({*(direct or ()), *(via_toolsets or ())}) @@ -2251,11 +2256,7 @@ class MCPRequestHandler: else None ) - key_tools: Final = ( - list(set(key_direct_tools or []) | set(key_toolset_tools or [])) - if key_direct_tools is not None or key_toolset_tools is not None - else None - ) + key_tools: Final = _as_list(MCPRequestHandler._union_tool_grants(key_direct_tools, key_toolset_tools)) team_direct_tools: Final = ( global_mcp_server_manager.expand_tool_permissions(team_obj_perm.mcp_tool_permissions).get(server_id) if team_obj_perm diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 312dcb27d89..be4df55ff58 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -27,7 +27,8 @@ from collections.abc import ( from contextlib import asynccontextmanager from dataclasses import dataclass, replace from functools import lru_cache -from itertools import chain +from itertools import chain, groupby +from operator import itemgetter from types import MappingProxyType from typing import TYPE_CHECKING, Any, Final, Generic, Literal, TypeAlias, TypedDict, TypeVar, cast from urllib.parse import ParseResult, urlparse @@ -6812,9 +6813,11 @@ class MCPServerManager: """ Rewrite an ``mcp_tool_permissions`` dict keyed by id/name/alias so every key is a concrete server_id where possible. Tool lists from - keys that point at the same server are unioned, matching the - "duplicate names grant access to all matches" semantics of - ``expand_permission_list``. + keys that point at the same server are unioned and deduplicated + first-seen, matching the "duplicate names grant access to all + matches" semantics of ``expand_permission_list``; the + ``MCP_ALL_TOOLS_WILDCARD`` entry is preserved as an ordinary list + entry for the caller to interpret. Required so name-based keys don't silently drop their tool restrictions when the lookup uses the resolved server_id. Unresolved @@ -6823,11 +6826,15 @@ class MCPServerManager: """ if not tool_permissions: return {} - result: Final[dict[str, list[str]]] = {} - for key, tools in tool_permissions.items(): - for server_id in self.expand_permission_list([key]): - result.setdefault(server_id, []).extend(tools or []) - return result + expanded: Final = tuple( + (server_id, tuple(tools or ())) + for key, tools in tool_permissions.items() + for server_id in self.expand_permission_list([key]) + ) + return { + server_id: list(dict.fromkeys(tool for _, tools in group for tool in tools)) + for server_id, group in groupby(sorted(expanded, key=itemgetter(0)), key=itemgetter(0)) + } def get_mcp_server_by_name(self, server_name: str, client_ip: str | None = None) -> MCPServer | None: """ diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py b/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py index 24f09e79dbb..fc4d7b45785 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py @@ -342,6 +342,15 @@ class TestMCPRequestHandler: mock_manager.resolve_toolset_tool_permissions = AsyncMock(return_value=toolset_perms) return mock_manager + def _real_manager_with_toolsets(self, toolset_perms): + """A real MCPServerManager so the real expand_tool_permissions runs; + only the DB-backed toolset lookup is stubbed""" + from litellm.proxy._experimental.mcp_server.mcp_server_manager import MCPServerManager + + manager = MCPServerManager() + manager.resolve_toolset_tool_permissions = AsyncMock(return_value=toolset_perms) + return manager + async def test_get_allowed_mcp_servers_for_key_includes_toolset_servers(self): """A key granted only mcp_toolsets must reach the toolset's servers on every path (list, call, REST); regression for the list-ok/call-403 bug""" @@ -508,6 +517,141 @@ class TestMCPRequestHandler: assert result is None + @pytest.mark.parametrize( + "direct,via_toolsets,expected", + [ + (["*"], None, None), + (["*"], ["read_file"], None), + (None, None, None), + ([], None, ()), + (None, ["read_file"], ("read_file",)), + ], + ) + def test_union_tool_grants_wildcard_and_union_cases(self, direct, via_toolsets, expected): + """A direct ["*"] makes the level unrestricted even beside a toolset + list (regression: mapping ["*"] to None in expand_tool_permissions let + a same-level toolset list deny every other tool)""" + result = MCPRequestHandler._union_tool_grants(direct, via_toolsets) + + if expected is None: + assert result is None + else: + assert result is not None + assert set(result) == set(expected) + + def test_union_tool_grants_unions_two_concrete_lists(self): + result = MCPRequestHandler._union_tool_grants(["read_file"], ["write_file"]) + + assert result is not None + assert set(result) == {"read_file", "write_file"} + + async def test_key_wildcard_allows_a_tool_never_enumerated(self): + """End to end at the key level: object_permission sits on the auth + object already, no team named, so no patching is needed; the real + global manager expands ["*"] and the level reads unrestricted""" + user_api_key_auth = UserAPIKeyAuth( + api_key="test-key", + object_permission=LiteLLM_ObjectPermissionTable( + object_permission_id="perm-1", + mcp_tool_permissions={"server-a": ["*"]}, + ), + ) + + allowed = await MCPRequestHandler.get_allowed_tools_for_server( + server_id="server-a", user_api_key_auth=user_api_key_auth + ) + brand_new_tool_allowed = await MCPRequestHandler.is_tool_allowed_for_server( + tool_name="brand_new_tool", server_id="server-a", user_api_key_auth=user_api_key_auth + ) + + assert allowed is None + assert brand_new_tool_allowed is True + + async def test_key_wildcard_stays_capped_by_team_allowlist(self): + """A wildcard on the key must never widen a team's enumerated ceiling: + the intersection keeps only the team's named tools""" + user_api_key_auth = UserAPIKeyAuth( + api_key="test-key", + team_id="team-1", + object_permission=LiteLLM_ObjectPermissionTable( + object_permission_id="perm-1", + mcp_tool_permissions={"server-a": ["*"]}, + ), + ) + team_object_permission = self._toolset_only_object_permission([]) + team_object_permission.mcp_tool_permissions = {"server-a": ["read_file"]} + manager = self._real_manager_with_toolsets({}) + + with ( + patch.object( # test-quality-ok: stub the DB team loader to drive the real team-server resolution path + MCPRequestHandler, "_get_team_object_permission", AsyncMock(return_value=team_object_permission) + ), + patch( # test-quality-ok: isolate the MCP registry, same seam as the sibling tests + "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager", + manager, + ), + ): + allowed = await MCPRequestHandler.get_allowed_tools_for_server( + server_id="server-a", user_api_key_auth=user_api_key_auth + ) + brand_new_tool_allowed = await MCPRequestHandler.is_tool_allowed_for_server( + tool_name="brand_new_tool", server_id="server-a", user_api_key_auth=user_api_key_auth + ) + + assert allowed == ["read_file"] + assert brand_new_tool_allowed is False + + async def test_team_wildcard_stays_capped_by_key_allowlist(self): + """A wildcard on the team leaves the key's enumerated list as the + effective ceiling""" + user_api_key_auth = UserAPIKeyAuth( + api_key="test-key", + team_id="team-1", + object_permission=LiteLLM_ObjectPermissionTable( + object_permission_id="perm-1", + mcp_tool_permissions={"server-a": ["read_file"]}, + ), + ) + team_object_permission = self._toolset_only_object_permission([]) + team_object_permission.mcp_tool_permissions = {"server-a": ["*"]} + manager = self._real_manager_with_toolsets({}) + + with ( + patch.object( # test-quality-ok: stub the DB team loader to drive the real team-server resolution path + MCPRequestHandler, "_get_team_object_permission", AsyncMock(return_value=team_object_permission) + ), + patch( # test-quality-ok: isolate the MCP registry, same seam as the sibling tests + "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager", + manager, + ), + ): + allowed = await MCPRequestHandler.get_allowed_tools_for_server( + server_id="server-a", user_api_key_auth=user_api_key_auth + ) + + assert allowed == ["read_file"] + + async def test_key_empty_tool_list_stays_deny_all(self): + """[] on the key is deny-all, distinct from the wildcard: it must not + be widened into allow-all""" + user_api_key_auth = UserAPIKeyAuth( + api_key="test-key", + object_permission=LiteLLM_ObjectPermissionTable( + object_permission_id="perm-1", + mcp_tool_permissions={"server-a": []}, + ), + ) + + allowed = await MCPRequestHandler.get_allowed_tools_for_server( + server_id="server-a", user_api_key_auth=user_api_key_auth + ) + read_file_allowed = await MCPRequestHandler.is_tool_allowed_for_server( + tool_name="read_file", server_id="server-a", user_api_key_auth=user_api_key_auth + ) + + assert allowed == [] + assert read_file_allowed is False + # ------------------------------------------------------------------ # LIT-5749: toolsets attached to a TEAM, ORG, or internal USER must be # enforced exactly like inline tool allowlists, on both axes diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py index cd5dae1269a..f3d37a858ca 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -8709,6 +8709,35 @@ class TestMCPServerManagerExpandToolPermissions: result = manager.expand_tool_permissions({"uuid-a": ["read_file"], "alias-a": ["write_file"]}) assert sorted(result["uuid-a"]) == ["read_file", "write_file"] + def test_wildcard_survives_expansion_as_list_entry(self): + """["*"] stays in the expanded list so the caller's wildcard check + (``_union_tool_grants``) can read it; this function only normalizes + keys and never maps grants to None.""" + manager = MCPServerManager() + manager.config_mcp_servers["uuid-a"] = self._make_server("uuid-a", server_name="alpha") + + result = manager.expand_tool_permissions({"uuid-a": ["*"]}) + assert result == {"uuid-a": ["*"]} + + def test_wildcard_unions_with_concrete_names_across_keys_for_same_server(self): + """An alias key carrying ["*"] unioned with an id key naming one tool + keeps both entries; interpretation of the wildcard belongs to the + caller, not the expansion.""" + manager = MCPServerManager() + manager.config_mcp_servers["uuid-a"] = self._make_server("uuid-a", server_name="alias-a", alias="alias-a") + + result = manager.expand_tool_permissions({"uuid-a": ["read_file"], "alias-a": ["*"]}) + assert sorted(result["uuid-a"]) == ["*", "read_file"] + + def test_empty_list_stays_deny_all(self): + """[] is deny-all, a distinct meaning from no entry (unrestricted); + the key must survive expansion rather than disappear.""" + manager = MCPServerManager() + manager.config_mcp_servers["uuid-a"] = self._make_server("uuid-a", server_name="alpha") + + result = manager.expand_tool_permissions({"uuid-a": []}) + assert result == {"uuid-a": []} + class TestOAuthDiscoverySSRFGuard: """SSRF guard for the OAuth metadata discovery follow-up fetches. diff --git a/ui/litellm-dashboard/src/components/mcp_server_management/MCPToolPermissions.test.tsx b/ui/litellm-dashboard/src/components/mcp_server_management/MCPToolPermissions.test.tsx index 149d231fff6..d3e3f204813 100644 --- a/ui/litellm-dashboard/src/components/mcp_server_management/MCPToolPermissions.test.tsx +++ b/ui/litellm-dashboard/src/components/mcp_server_management/MCPToolPermissions.test.tsx @@ -133,9 +133,10 @@ describe("MCPToolPermissions", () => { const selectAllButton = screen.getByRole("button", { name: "Select All" }); await userEvent.click(selectAllButton); - // Verify onChange was called with all tools selected + // Selecting every displayed tool writes the wildcard, which also covers tools the + // server adds later. expect(mockOnChange).toHaveBeenCalledWith({ - [mockServerId]: ["read_wiki_structure", "read_wiki_contents", "ask_question"], + [mockServerId]: ["*"], }); }); @@ -190,6 +191,77 @@ describe("MCPToolPermissions", () => { }); }); + describe("wildcard all-tools grant", () => { + const wildcardServerId = "server-1"; + const wildcardServer = { server_id: wildcardServerId, server_name: "Wildcard Server", alias: "Wildcard Server" }; + const wildcardTools = [ + { name: "read_wiki_structure", description: "Get documentation topics" }, + { name: "read_wiki_contents", description: "View documentation" }, + { name: "ask_question", description: "Ask questions" }, + ]; + + beforeEach(() => { + vi.mocked(networking.fetchMCPServers).mockResolvedValue([wildcardServer]); + vi.mocked(networking.listMCPTools).mockResolvedValue({ tools: wildcardTools, error: false }); + }); + + it("renders every tool checked with the future-tools note when the entry is the wildcard", async () => { + renderWithProviders( + , + ); + + expect(await screen.findByText("Wildcard Server")).toBeInTheDocument(); + expect(screen.getByText("All tools allowed, including tools added to this server later")).toBeInTheDocument(); + + await userEvent.click(screen.getByText("Flat List")); + for (const checkbox of screen.getAllByRole("checkbox")) { + expect(checkbox).toBeChecked(); + } + }); + + it("writes the wildcard when Select All covers every displayed tool", async () => { + const mockOnChange = vi.fn(); + renderWithProviders( + , + ); + + expect(await screen.findByText("read_wiki_structure")).toBeInTheDocument(); + await userEvent.click(screen.getByRole("button", { name: "Select All" })); + + expect(mockOnChange).toHaveBeenCalledWith({ [wildcardServerId]: ["*"] }); + }); + + it("converts back to an enumerated list when one tool is unchecked from a wildcard grant", async () => { + const mockOnChange = vi.fn(); + renderWithProviders( + , + ); + + expect(await screen.findByText("read_wiki_structure")).toBeInTheDocument(); + await userEvent.click(screen.getByText("Flat List")); + await userEvent.click(screen.getByRole("checkbox", { name: "ask_question" })); + + expect(mockOnChange).toHaveBeenCalledWith({ + [wildcardServerId]: ["read_wiki_structure", "read_wiki_contents"], + }); + }); + }); + describe("servers reached indirectly", () => { const groupServer = { server_id: "srv-group-1", @@ -428,6 +500,8 @@ describe("MCPToolPermissions", () => { expect(await screen.findByText("list_issues")).toBeInTheDocument(); await userEvent.click(screen.getByText("Select All")); + // A toolset-sourced server never writes the wildcard: that would create a standing direct + // grant outliving the toolset. The write keeps only the tools this level grants itself. expect(mockOnChange).toHaveBeenCalledWith({ [toolsetServer.server_id]: ["delete_issue"] }); }); @@ -849,7 +923,7 @@ describe("MCPToolPermissions", () => { const written = mockOnChange.mock.calls.at(-1)?.[0] as Record; expect(written["github_mcp"]).toEqual(["list_issues"]); - expect(written[twin.server_id]).toEqual(["list_issues", "create_issue", "delete_issue"]); + expect(written[twin.server_id]).toEqual(["*"]); }); it("says nothing about shared names when every key names one server", async () => { diff --git a/ui/litellm-dashboard/src/components/mcp_server_management/MCPToolPermissions.tsx b/ui/litellm-dashboard/src/components/mcp_server_management/MCPToolPermissions.tsx index e26f1a6f511..7edeaedff2e 100644 --- a/ui/litellm-dashboard/src/components/mcp_server_management/MCPToolPermissions.tsx +++ b/ui/litellm-dashboard/src/components/mcp_server_management/MCPToolPermissions.tsx @@ -8,13 +8,14 @@ import { useMCPAccessGroups } from "../../app/(dashboard)/hooks/mcpServers/useMC import { useMCPToolsets } from "../../app/(dashboard)/hooks/mcpServers/useMCPToolsets"; import McpCrudPermissionPanel from "../mcp_tools/McpCrudPermissionPanel"; import { classifyToolOp } from "../../utils/mcpToolCrudClassification"; -import { NO_MCP_SERVERS_SENTINEL } from "../mcp_tools/constants"; +import { MCP_ALL_TOOLS_WILDCARD, NO_MCP_SERVERS_SENTINEL } from "../mcp_tools/constants"; import { EffectiveMcpServer, McpGrantSource, applyToolPermissionWrite, emptyMcpAccessGroups, mcpAllowedToolsFor, + mcpGrantsAllTools, resolveEffectiveMcpServers, } from "./effectiveMcpServers"; @@ -150,7 +151,12 @@ const MCPToolPermissions: React.FC = ({ // Every write goes through here so an edit is authoritative for the SERVER, not for one of the // equivalent keys that may name it. const writeAllowedTools = (entry: EffectiveMcpServer, allowed: string[]) => { - onChange(applyToolPermissionWrite({ toolPermissions, entry, allowed })); + const names = (serverTools[entry.server.server_id] ?? []).map((t) => t.name); + const next = + entry.source.kind !== "toolset" && names.length > 0 && names.every((n) => allowed.includes(n)) + ? [MCP_ALL_TOOLS_WILDCARD] + : allowed; + onChange(applyToolPermissionWrite({ toolPermissions, entry, allowed: next })); }; const handleSelectAll = (entry: EffectiveMcpServer) => { @@ -222,7 +228,8 @@ const MCPToolPermissions: React.FC = ({ const serverId = server.server_id; const serverName = server.server_name || server.alias || serverId; const tools = serverTools[serverId] || []; - const selectedTools = entry.allowedTools ?? tools.map((t) => t.name); + const grantsAll = mcpGrantsAllTools(entry.keyedTools); + const selectedTools = grantsAll ? tools.map((t) => t.name) : entry.allowedTools ?? tools.map((t) => t.name); const isLoading = loadingTools[serverId]; const error = toolErrors[serverId]; const viewMode = viewModes[serverId] ?? "crud"; @@ -247,6 +254,11 @@ const MCPToolPermissions: React.FC = ({ )} {server.description &&

{server.description}

} + {grantsAll && ( +

+ All tools allowed, including tools added to this server later +

+ )} {entry.ambiguousKeys.length > 0 && (

{`Also granted by ${entry.ambiguousKeys.map((key) => `"${key}"`).join(", ")}, which names another server too. Those tools stay allowed here until the servers no longer share that name`} diff --git a/ui/litellm-dashboard/src/components/mcp_server_management/effectiveMcpServers.test.ts b/ui/litellm-dashboard/src/components/mcp_server_management/effectiveMcpServers.test.ts index 487c6f9f55e..07f2b0e2508 100644 --- a/ui/litellm-dashboard/src/components/mcp_server_management/effectiveMcpServers.test.ts +++ b/ui/litellm-dashboard/src/components/mcp_server_management/effectiveMcpServers.test.ts @@ -4,6 +4,7 @@ import { applyToolPermissionWrite, emptyMcpAccessGroups, mcpAllowedToolsFor, + mcpGrantsAllTools, mcpServersForIdentifier, mcpToolPermissionKeyFor, resolveEffectiveMcpServers, @@ -66,6 +67,16 @@ describe("mcpServersForIdentifier", () => { }); }); +describe("mcpGrantsAllTools", () => { + it("is true only when the union carries the wildcard, never for an absent grant", () => { + expect(mcpGrantsAllTools(["*"])).toBe(true); + expect(mcpGrantsAllTools(["read_file", "*"])).toBe(true); + expect(mcpGrantsAllTools(["read_file"])).toBe(false); + expect(mcpGrantsAllTools([])).toBe(false); + expect(mcpGrantsAllTools(undefined)).toBe(false); + }); +}); + describe("mcpToolPermissionKeyFor", () => { const target = server({ server_id: "uuid-1", server_name: "github_mcp", alias: "GitHub" }); diff --git a/ui/litellm-dashboard/src/components/mcp_server_management/effectiveMcpServers.ts b/ui/litellm-dashboard/src/components/mcp_server_management/effectiveMcpServers.ts index b3ba24f3c59..c9e85fef31b 100644 --- a/ui/litellm-dashboard/src/components/mcp_server_management/effectiveMcpServers.ts +++ b/ui/litellm-dashboard/src/components/mcp_server_management/effectiveMcpServers.ts @@ -1,5 +1,6 @@ import { z } from "zod/v4"; import { MCPServer, MCPToolset } from "../mcp_tools/types"; +import { MCP_ALL_TOOLS_WILDCARD } from "../mcp_tools/constants"; // Mirrors the backend resolver's union (direct + access_group + tool_perm + toolset), so the // editor shows exactly the servers this permission level entitles. @@ -121,6 +122,12 @@ export const mcpAllowedToolsFor = ( return [...new Set(keys.flatMap((key) => toolPermissions[key] ?? []))]; }; +// An allowed-tools union carrying the wildcard grants every current and future tool on the +// server; `undefined` (no entry at all) is unrestricted for a different reason and is not a +// wildcard grant the editor should expand. +export const mcpGrantsAllTools = (allowed: readonly string[] | undefined): boolean => + allowed !== undefined && allowed.includes(MCP_ALL_TOOLS_WILDCARD); + // Tool names the given toolsets grant on this server, `undefined` when they grant none. const mcpToolsetToolsFor = ( server: MCPServer, diff --git a/ui/litellm-dashboard/src/components/mcp_tools/constants.ts b/ui/litellm-dashboard/src/components/mcp_tools/constants.ts index 66ab1a352f4..eef98383d2a 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/constants.ts +++ b/ui/litellm-dashboard/src/components/mcp_tools/constants.ts @@ -3,5 +3,8 @@ export const NO_MCP_SERVERS_SENTINEL = "no-mcp-servers"; export const ALL_PROXY_MCP_SERVERS_SENTINEL = "all-proxy-mcpservers"; +// Must match the backend MCP_ALL_TOOLS_WILDCARD constant in litellm/constants.py. +export const MCP_ALL_TOOLS_WILDCARD = "*"; + export const MCP_TOOLS_PREVIEW_FORBIDDEN_MESSAGE = "Tool preview is not available for submissions. Tools will be verified by an admin during review.";