mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-02 02:11:58 +00:00
fix(proxy): preserve decision request bodies under token limits (#43920)
This commit is contained in:
parent
c2ae483782
commit
d96477abce
6 changed files with 123 additions and 12 deletions
|
|
@ -75,6 +75,7 @@ from litellm.router_utils.add_retry_fallback_headers import (
|
|||
from litellm.router_utils.common_utils import resolve_model_group_alias
|
||||
from litellm.types.caching import RedisPipelineIncrementOperation
|
||||
from litellm.types.llms.openai import BaseLiteLLMOpenAIResponseObject, ResponseAPIUsage
|
||||
from litellm.types.passthrough_endpoints.pass_through_endpoints import EndpointType
|
||||
from litellm.types.utils import (
|
||||
CallTypes,
|
||||
EmbeddingResponse,
|
||||
|
|
@ -978,6 +979,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
|
|||
min_configured_limit: int | None,
|
||||
call_type: str | None,
|
||||
configured_output_tokens: int | None = None,
|
||||
endpoint_type: EndpointType = EndpointType.GENERIC,
|
||||
) -> None:
|
||||
"""Hard-cap generation length when the request has no explicit cap.
|
||||
|
||||
|
|
@ -1006,6 +1008,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
|
|||
capped_floor >= baseline_floor
|
||||
or _PROXY_MaxParallelRequestsHandler_v3._has_explicit_output_cap(data, call_type)
|
||||
or is_embedding
|
||||
or endpoint_type == EndpointType.DECISIONS
|
||||
):
|
||||
return
|
||||
effective_cap: Final = max(capped_floor, configured_output_tokens or 0)
|
||||
|
|
@ -3757,6 +3760,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
|
|||
tpm_reservation_scopes: Sequence[tuple[str, str]],
|
||||
tpm_reservation_amount: int,
|
||||
call_type: str | None = None,
|
||||
endpoint_type: EndpointType = EndpointType.GENERIC,
|
||||
) -> None:
|
||||
"""
|
||||
Reserve project-scoped ITPM/OTPM tokens (Bedrock Mantle-style
|
||||
|
|
@ -3810,6 +3814,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
|
|||
data=data,
|
||||
min_configured_limit=min_configured_otpm_limit,
|
||||
call_type=call_type,
|
||||
endpoint_type=endpoint_type,
|
||||
)
|
||||
|
||||
io_response, itpm_reserved, otpm_reserved = await self.reserve_io_tokens(
|
||||
|
|
@ -3984,6 +3989,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
|
|||
cache: DualCache,
|
||||
data: dict,
|
||||
call_type: str,
|
||||
endpoint_type: EndpointType = EndpointType.GENERIC,
|
||||
):
|
||||
"""
|
||||
Pre-call hook to check rate limits before making the API call.
|
||||
|
|
@ -4115,6 +4121,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
|
|||
min_configured_limit=min_configured_tpm_limit,
|
||||
call_type=call_type,
|
||||
configured_output_tokens=configured_output_tokens,
|
||||
endpoint_type=endpoint_type,
|
||||
)
|
||||
|
||||
# Floor at 1 token so contentless requests (/responses,
|
||||
|
|
@ -4201,6 +4208,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
|
|||
tpm_reservation_scopes=tpm_reservation_scopes,
|
||||
tpm_reservation_amount=tpm_reservation_amount,
|
||||
call_type=call_type,
|
||||
endpoint_type=endpoint_type,
|
||||
)
|
||||
|
||||
def _create_pipeline_operations(
|
||||
|
|
|
|||
|
|
@ -377,8 +377,10 @@ class HttpPassThroughEndpointHelpers(BasePassthroughUtils):
|
|||
return return_headers
|
||||
|
||||
@staticmethod
|
||||
def get_endpoint_type(url: str) -> EndpointType:
|
||||
def get_endpoint_type(url: str, custom_llm_provider: str | None = None) -> EndpointType:
|
||||
parsed_url: Final = urlparse(url)
|
||||
if custom_llm_provider == "typesafe" and parsed_url.path.removesuffix("/").endswith("/v1/systemone"):
|
||||
return EndpointType.DECISIONS
|
||||
if (
|
||||
("generateContent") in url
|
||||
or ("streamGenerateContent") in url
|
||||
|
|
@ -1093,7 +1095,9 @@ async def pass_through_request(
|
|||
|
||||
requested_query_params: dict | None = query_params or dict(request.query_params) or None
|
||||
|
||||
endpoint_type: Final[EndpointType] = HttpPassThroughEndpointHelpers.get_endpoint_type(str(url))
|
||||
endpoint_type: Final[EndpointType] = HttpPassThroughEndpointHelpers.get_endpoint_type(
|
||||
str(url), custom_llm_provider
|
||||
)
|
||||
|
||||
# SigV4-signed callers (e.g. Bedrock) attach the exact bytes that were
|
||||
# signed via request.state; we must send those instead of re-encoding the
|
||||
|
|
@ -1180,6 +1184,7 @@ async def pass_through_request(
|
|||
user_api_key_dict=user_api_key_dict,
|
||||
data=_parsed_body,
|
||||
call_type="pass_through_endpoint",
|
||||
endpoint_type=endpoint_type,
|
||||
)
|
||||
resolved_timeout: Final = resolve_pass_through_request_timeout(timeout)
|
||||
async_client_obj: Final = get_async_httpx_client(
|
||||
|
|
|
|||
|
|
@ -245,6 +245,7 @@ from litellm.types.mcp import (
|
|||
MCPPreCallRequestObject,
|
||||
MCPPreCallResponseObject,
|
||||
)
|
||||
from litellm.types.passthrough_endpoints.pass_through_endpoints import EndpointType
|
||||
from litellm.types.proxy.policy_engine.pipeline_types import PipelineExecutionResult
|
||||
from litellm.types.utils import LLMResponseTypes, LoggedLiteLLMParams
|
||||
from litellm.utils import (
|
||||
|
|
@ -2353,6 +2354,7 @@ class ProxyLogging:
|
|||
call_type: CallTypesLiteral,
|
||||
guardrails_only: bool = False,
|
||||
skip_guardrails: bool = False,
|
||||
endpoint_type: EndpointType = EndpointType.GENERIC,
|
||||
) -> None:
|
||||
pass
|
||||
|
||||
|
|
@ -2364,6 +2366,7 @@ class ProxyLogging:
|
|||
call_type: CallTypesLiteral,
|
||||
guardrails_only: bool = False,
|
||||
skip_guardrails: bool = False,
|
||||
endpoint_type: EndpointType = EndpointType.GENERIC,
|
||||
) -> dict:
|
||||
pass
|
||||
|
||||
|
|
@ -2374,6 +2377,7 @@ class ProxyLogging:
|
|||
call_type: CallTypesLiteral,
|
||||
guardrails_only: bool = False,
|
||||
skip_guardrails: bool = False,
|
||||
endpoint_type: EndpointType = EndpointType.GENERIC,
|
||||
) -> dict | None:
|
||||
"""
|
||||
Allows users to modify/reject the incoming request to the proxy, without having to deal with parsing Request body.
|
||||
|
|
@ -2512,11 +2516,21 @@ class ProxyLogging:
|
|||
if call_type in MCP_GUARDRAIL_CALL_TYPES and user_api_key_dict is None:
|
||||
continue
|
||||
|
||||
response: Exception | str | Mapping[str, object] | None = await _callback.async_pre_call_hook(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
cache=self.call_details["user_api_key_cache"],
|
||||
data=data,
|
||||
call_type=call_type,
|
||||
response: Exception | str | Mapping[str, object] | None = (
|
||||
await _callback.async_pre_call_hook(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
cache=self.call_details["user_api_key_cache"],
|
||||
data=data,
|
||||
call_type=call_type,
|
||||
endpoint_type=endpoint_type,
|
||||
)
|
||||
if isinstance(_callback, _PROXY_MaxParallelRequestsHandler_v3)
|
||||
else await _callback.async_pre_call_hook(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
cache=self.call_details["user_api_key_cache"],
|
||||
data=data,
|
||||
call_type=call_type,
|
||||
)
|
||||
)
|
||||
if response is not None:
|
||||
data = await self.process_pre_call_hook_response(
|
||||
|
|
|
|||
|
|
@ -30,6 +30,7 @@ class EndpointType(str, Enum):
|
|||
OPENAI = "openai"
|
||||
TINYFISH = "tinyfish"
|
||||
GENERIC = "generic"
|
||||
DECISIONS = "decisions"
|
||||
|
||||
|
||||
class PassthroughStandardLoggingPayload(TypedDict, total=False):
|
||||
|
|
|
|||
|
|
@ -7,7 +7,7 @@ import os
|
|||
import traceback
|
||||
from collections.abc import AsyncIterator, Awaitable, Callable, Iterator, Mapping
|
||||
from types import MappingProxyType, SimpleNamespace
|
||||
from typing import Final
|
||||
from typing import Final, Literal
|
||||
from unittest import mock
|
||||
from unittest.mock import AsyncMock, MagicMock, Mock, patch
|
||||
from urllib.parse import parse_qs
|
||||
|
|
@ -7292,6 +7292,82 @@ class TestTypeSafePassthroughRoute:
|
|||
assert sent.headers["authorization"] == "Bearer typesafe-test-key"
|
||||
assert json.loads(sent.content or b"{}") == (body or {})
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"provider, endpoint, is_decision_request",
|
||||
(
|
||||
("typesafe", "systemone", True),
|
||||
("typesafe", "systemone/", True),
|
||||
("typesafe", "systemone?trace=1", True),
|
||||
("typesafe", "systemone/?trace=1", True),
|
||||
("typesafe", "systemone/other", False),
|
||||
("typesafe", "systemone/other/", False),
|
||||
("typesafe", "systemone-other", False),
|
||||
("typesafe", "chat/completions?next=/typesafe/v1/systemone", False),
|
||||
("openrouter", "systemone", False),
|
||||
("openrouter", "systemone/", False),
|
||||
("openrouter", "chat/completions", False),
|
||||
),
|
||||
)
|
||||
@pytest.mark.parametrize("quota_scope", ("key", "project_output"))
|
||||
@pytest.mark.parametrize("token_limit", (0, 1000))
|
||||
def test_token_limits_preserve_decisions_cap_generation_and_enforce_quota(
|
||||
self,
|
||||
client: TestClient,
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
provider: Literal["typesafe", "openrouter"],
|
||||
endpoint: str,
|
||||
is_decision_request: bool,
|
||||
quota_scope: Literal["key", "project_output"],
|
||||
token_limit: int,
|
||||
) -> None:
|
||||
from litellm.caching.caching import DualCache
|
||||
from litellm.proxy import proxy_server
|
||||
from litellm.proxy.hooks.cache_control_check import _PROXY_CacheControlCheck
|
||||
from litellm.proxy.hooks.parallel_request_limiter_v3 import (
|
||||
_PROXY_MaxParallelRequestsHandler_v3,
|
||||
get_request_stash,
|
||||
)
|
||||
from litellm.proxy.utils import InternalUsageCache, ProxyLogging
|
||||
|
||||
cache: Final = DualCache()
|
||||
limiter: Final = _PROXY_MaxParallelRequestsHandler_v3(internal_usage_cache=InternalUsageCache(cache))
|
||||
monkeypatch.setattr(litellm, "callbacks", list((limiter, _PROXY_CacheControlCheck())))
|
||||
monkeypatch.setattr(proxy_server, "proxy_logging_obj", ProxyLogging(user_api_key_cache=cache))
|
||||
monkeypatch.setenv("OPENROUTER_API_KEY", "openrouter-test-key")
|
||||
monkeypatch.setenv("OPENROUTER_API_BASE", "https://typesafe.example/base")
|
||||
model: Final = "jev-latest" if provider == "typesafe" else "test-generative-model"
|
||||
auth: Final = UserAPIKeyAuth(
|
||||
api_key="sk-limited",
|
||||
tpm_limit=token_limit if quota_scope == "key" else None,
|
||||
project_id="test-project" if quota_scope == "project_output" else None,
|
||||
project_metadata={"model_otpm_limit": {model: token_limit}} if quota_scope == "project_output" else {},
|
||||
)
|
||||
monkeypatch.setitem(proxy_server.app.dependency_overrides, user_api_key_auth, lambda: auth)
|
||||
body: Final = (
|
||||
{
|
||||
"model": model,
|
||||
"state": "A request for help",
|
||||
"questions": {"urgent": {"type": "noul", "instructions": "Is this urgent?"}},
|
||||
}
|
||||
if is_decision_request
|
||||
else {"model": model, "messages": [{"role": "user", "content": "Hello"}]}
|
||||
)
|
||||
|
||||
def upstream_response(request: httpx.Request) -> httpx.Response:
|
||||
expected_body: Final = body if is_decision_request else {**body, "max_tokens": token_limit // 4}
|
||||
assert json.loads(request.content) == expected_body
|
||||
stash: Final = get_request_stash()
|
||||
assert stash is not None
|
||||
assert (stash.reserved_tokens if quota_scope == "key" else stash.otpm_reserved_tokens) > 0
|
||||
return httpx.Response(200, json={"model": model})
|
||||
|
||||
with respx.mock(assert_all_called=False) as upstream:
|
||||
route: Final = upstream.post(f"https://typesafe.example/base/v1/{endpoint}").mock(side_effect=upstream_response)
|
||||
response: Final = client.post(f"/{provider}/v1/{endpoint}", json=body)
|
||||
|
||||
assert response.status_code == (429 if token_limit == 0 else 200), response.text
|
||||
assert route.call_count == (0 if token_limit == 0 else 1)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_forwards_target_auth_headers_provider_and_query(self, monkeypatch):
|
||||
monkeypatch.setenv("TYPESAFE_API_KEY", "typesafe-test-key")
|
||||
|
|
|
|||
|
|
@ -49,6 +49,7 @@ from litellm.proxy.pass_through_endpoints.success_handler import (
|
|||
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 (
|
||||
EndpointType,
|
||||
LITELLM_PASS_THROUGH_DEPLOYMENT_MODEL_INFO_STATE_KEY,
|
||||
LITELLM_PASS_THROUGH_RAW_BODY_STATE_KEY,
|
||||
)
|
||||
|
|
@ -1689,7 +1690,7 @@ async def test_pass_through_request_streamed_response_is_owned_by_the_caller():
|
|||
cache_dict[cache_key] = SimpleNamespace(client=httpx.AsyncClient(transport=httpx.MockTransport(transport_handler)))
|
||||
|
||||
mock_proxy_logging = MagicMock()
|
||||
mock_proxy_logging.pre_call_hook = AsyncMock(side_effect=lambda user_api_key_dict, data, call_type: data)
|
||||
mock_proxy_logging.pre_call_hook = AsyncMock(side_effect=lambda user_api_key_dict, data, call_type, endpoint_type: data)
|
||||
mock_proxy_logging.post_call_failure_hook = AsyncMock()
|
||||
mock_proxy_logging.post_call_response_headers_hook = AsyncMock(return_value={})
|
||||
mock_proxy_logging.get_proxy_hook = MagicMock(return_value=MagicMock())
|
||||
|
|
@ -2591,7 +2592,9 @@ async def _run_pass_through_and_capture_wire_url(
|
|||
mock_request.body = AsyncMock(return_value=b"")
|
||||
|
||||
mock_proxy_logging = MagicMock()
|
||||
mock_proxy_logging.pre_call_hook = AsyncMock(side_effect=lambda user_api_key_dict, data, call_type: data)
|
||||
mock_proxy_logging.pre_call_hook = AsyncMock(
|
||||
side_effect=lambda user_api_key_dict, data, call_type, endpoint_type=None: data
|
||||
)
|
||||
mock_proxy_logging.post_call_failure_hook = AsyncMock()
|
||||
mock_proxy_logging.post_call_response_headers_hook = AsyncMock(return_value={})
|
||||
mock_proxy_logging.get_proxy_hook = MagicMock(return_value=managed_files_hook)
|
||||
|
|
@ -4889,7 +4892,9 @@ async def test_pass_through_request_mid_stream_upstream_drop_fires_failure_hook(
|
|||
cache_dict[cache_key] = SimpleNamespace(client=httpx.AsyncClient(transport=httpx.MockTransport(transport_handler)))
|
||||
|
||||
mock_proxy_logging = MagicMock()
|
||||
mock_proxy_logging.pre_call_hook = AsyncMock(side_effect=lambda user_api_key_dict, data, call_type: data)
|
||||
mock_proxy_logging.pre_call_hook = AsyncMock(
|
||||
side_effect=lambda user_api_key_dict, data, call_type, endpoint_type=None: data
|
||||
)
|
||||
mock_proxy_logging.post_call_failure_hook = AsyncMock()
|
||||
mock_proxy_logging.post_call_response_headers_hook = AsyncMock(return_value=None)
|
||||
mock_proxy_logging.get_proxy_hook = MagicMock(return_value=None)
|
||||
|
|
@ -7093,7 +7098,9 @@ async def _drive_passthrough_request_and_capture_logging(
|
|||
|
||||
captured_data: dict = {} # mutable-ok: the pre-call hook records the request data into it
|
||||
|
||||
async def capture_pre_call_hook(user_api_key_dict, data, call_type):
|
||||
async def capture_pre_call_hook(
|
||||
user_api_key_dict, data, call_type, endpoint_type: EndpointType = EndpointType.GENERIC
|
||||
):
|
||||
captured_data.update(data)
|
||||
if on_pre_call is not None:
|
||||
on_pre_call(data.get("litellm_logging_obj"))
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue