diff --git a/litellm/proxy/shunt_endpoints/endpoints.py b/litellm/proxy/shunt_endpoints/endpoints.py index a89dfe2b2b5..7033ebbfef0 100644 --- a/litellm/proxy/shunt_endpoints/endpoints.py +++ b/litellm/proxy/shunt_endpoints/endpoints.py @@ -16,7 +16,7 @@ from typing import TYPE_CHECKING, Annotated, Final from fastapi import APIRouter, Depends, File, Form, Header, HTTPException, Query, Request, UploadFile from litellm._logging import verbose_proxy_logger -from litellm.litellm_core_utils.env_utils import get_env_int +from litellm.litellm_core_utils.env_utils import get_env_int_in_range from litellm.proxy._types import LitellmUserRoles, ProxyErrorTypes, ProxyException, UserAPIKeyAuth from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing from litellm.proxy.guardrails.auto_router_shunt import ShuntConfig, shunt_config_for_model @@ -156,13 +156,16 @@ async def _worker_text( route_type="acompletion", llm_router=llm_router, ) - response: Final = await route_request( # pyright: ignore[reportUnknownVariableType] # route_request's own return type is intentionally an untyped union (see its ANN202 suppression) + llm_call: Final = await route_request( # pyright: ignore[reportUnknownVariableType] # route_request's own return type is intentionally an untyped union (see its ANN202 suppression) data=data, route_type="acompletion", llm_router=llm_router, user_model=None, user_api_key_dict=user_api_key_dict, ) + # Two awaits: route_request resolves the deployment and hands back the provider + # coroutine unawaited (see its own ANN202 note), so this second await is the call. + response: Final = await llm_call # pyright: ignore[reportUnknownVariableType] # same untyped union as above processed: Final = await proxy_logging_obj.post_call_success_hook( data=data, user_api_key_dict=user_api_key_dict, @@ -186,9 +189,17 @@ async def _worker_text( # multipart uploads specifically to hand their contents to a worker model, an authenticated # caller with a valid capability token could otherwise upload enough data to exhaust a proxy # worker's memory before the size limit ever runs. -_MAX_UPLOAD_BYTES_PER_FILE: Final = get_env_int("LITELLM_SHUNT_MAX_UPLOAD_BYTES_PER_FILE", 1024 * 1024) -_MAX_UPLOAD_BYTES_TOTAL: Final = get_env_int("LITELLM_SHUNT_MAX_UPLOAD_BYTES_TOTAL", 8 * 1024 * 1024) -_MAX_UPLOAD_FILE_COUNT: Final = get_env_int("LITELLM_SHUNT_MAX_UPLOAD_FILE_COUNT", 20) +# +# Range-constrained rather than plain get_env_int: a zero or negative override would make the +# per-read limit non-positive, and UploadFile.read() treats a negative size as "read the whole +# file", so a typo'd env var would silently turn the bound it configures into no bound at all. +_MAX_UPLOAD_BYTES_PER_FILE: Final = get_env_int_in_range( + "LITELLM_SHUNT_MAX_UPLOAD_BYTES_PER_FILE", 1024 * 1024, minimum=1, maximum=128 * 1024 * 1024 +) +_MAX_UPLOAD_BYTES_TOTAL: Final = get_env_int_in_range( + "LITELLM_SHUNT_MAX_UPLOAD_BYTES_TOTAL", 8 * 1024 * 1024, minimum=1, maximum=512 * 1024 * 1024 +) +_MAX_UPLOAD_FILE_COUNT: Final = get_env_int_in_range("LITELLM_SHUNT_MAX_UPLOAD_FILE_COUNT", 20, minimum=1, maximum=1000) async def _read_upload_text(upload: UploadFile, *, remaining_total_bytes: int) -> str: @@ -196,22 +207,20 @@ async def _read_upload_text(upload: UploadFile, *, remaining_total_bytes: int) - remaining-total byte budgets -- never the whole file, so a caller can't force this endpoint to buffer more than that regardless of how large the real upload is. - `remaining_total_bytes <= 0` is checked explicitly rather than left to `min()` + - `.read(limit + 1)`: a non-positive `remaining_total_bytes` would make `limit` zero or - negative, and `UploadFile.read` treats a negative size as "read the whole file", which - would silently defeat this budget for any caller of this function that ever passes one. - `_read_upload_texts` below never actually produces a negative value (each read is already - bounded by what was left when it started), so this is the function's own contract holding - regardless of caller, not a path reachable through that call site today. + The computed `limit` is checked before it reaches `.read()`, not just the incoming + `remaining_total_bytes`: `UploadFile.read` treats a negative size as "read the whole file", + so a non-positive limit from *either* input -- an exhausted total budget, or a zero/negative + `_MAX_UPLOAD_BYTES_PER_FILE` from a bad env override -- would silently turn this bound into + no bound at all. Checking the value actually passed to `.read()` covers both at once. """ - if remaining_total_bytes <= 0: + limit: Final = min(_MAX_UPLOAD_BYTES_PER_FILE, remaining_total_bytes) + if limit <= 0: raise ProxyException( message=f"uploads exceed the {_MAX_UPLOAD_BYTES_TOTAL}-byte total limit for this call", type=ProxyErrorTypes.bad_request_error, param="paths", code=400, ) - limit: Final = min(_MAX_UPLOAD_BYTES_PER_FILE, remaining_total_bytes) content: Final = await upload.read(limit + 1) if len(content) > limit: raise ProxyException( diff --git a/tests/test_litellm/proxy/client/cli/autoroute/test_settings.py b/tests/test_litellm/proxy/client/cli/autoroute/test_settings.py index 1a456f781d5..87a33c79a79 100644 --- a/tests/test_litellm/proxy/client/cli/autoroute/test_settings.py +++ b/tests/test_litellm/proxy/client/cli/autoroute/test_settings.py @@ -9,36 +9,49 @@ def test_preserves_unrelated_top_level_keys(): assert merged["theme"] == "dark" -def test_sets_base_url_and_auth_token(): - merged = merge_claude_settings_static_token({}, "http://127.0.0.1:4000", "token-abc") - assert merged["env"]["ANTHROPIC_BASE_URL"] == "http://127.0.0.1:4000" - assert merged["env"]["ANTHROPIC_AUTH_TOKEN"] == "token-abc" - - -def test_strips_trailing_slash_from_base_url(): - merged = merge_claude_settings_static_token({}, "http://127.0.0.1:4000/", "token-abc") - assert merged["env"]["ANTHROPIC_BASE_URL"] == "http://127.0.0.1:4000" - - -def test_clears_existing_api_key_env_var(): - settings = {"env": {"ANTHROPIC_API_KEY": "sk-old"}} - merged = merge_claude_settings_static_token(settings, "http://127.0.0.1:4000", "token-abc") - assert "ANTHROPIC_API_KEY" not in merged["env"] - - -def test_clears_existing_api_key_helper(): - settings = {"apiKeyHelper": "some-script.sh"} - merged = merge_claude_settings_static_token(settings, "http://127.0.0.1:4000", "token-abc") - assert "apiKeyHelper" not in merged - - -def test_preserves_other_env_vars(): +def test_preserves_unrelated_env_keys(): settings = {"env": {"SOME_OTHER_VAR": "value"}} merged = merge_claude_settings_static_token(settings, "http://127.0.0.1:4000", "token-abc") assert merged["env"]["SOME_OTHER_VAR"] == "value" -def test_sets_every_default_model_env_key_to_autorouter(): +def test_sets_base_url_and_auth_token(): + merged = merge_claude_settings_static_token({}, "http://127.0.0.1:4000/", "token-abc") + assert merged["env"]["ANTHROPIC_BASE_URL"] == "http://127.0.0.1:4000" + assert merged["env"]["ANTHROPIC_AUTH_TOKEN"] == "token-abc" + assert merged["env"]["ENABLE_TOOL_SEARCH"] == "true" + + +def test_preserves_existing_tool_search(): + settings = {"env": {"ENABLE_TOOL_SEARCH": "false"}} + merged = merge_claude_settings_static_token(settings, "http://127.0.0.1:4000", "token-abc") + assert merged["env"]["ENABLE_TOOL_SEARCH"] == "false" + + +def test_drops_stray_api_key(): + settings = {"env": {"ANTHROPIC_API_KEY": "leaked-key"}} + merged = merge_claude_settings_static_token(settings, "http://127.0.0.1:4000", "token-abc") + assert "ANTHROPIC_API_KEY" not in merged["env"] + + +def test_removes_existing_api_key_helper(): + settings = {"apiKeyHelper": "/usr/local/bin/lite auth print-token"} + merged = merge_claude_settings_static_token(settings, "http://127.0.0.1:4000", "token-abc") + assert "apiKeyHelper" not in merged + + +def test_does_not_mutate_input(): + settings = {"env": {"FOO": "bar"}, "apiKeyHelper": "old-helper"} + merge_claude_settings_static_token(settings, "http://127.0.0.1:4000", "token-abc") + assert settings == {"env": {"FOO": "bar"}, "apiKeyHelper": "old-helper"} + + +def test_forces_all_claude_code_default_model_tiers_to_the_autorouter(): + # A bare "*" model_name deployment looks like the obvious way to catch every request + # regardless of which model Claude Code thinks it's using, but Router's auto-router + # registry is keyed by the literal requested model string with no wildcard resolution + # (litellm/router.py:10711-10717) -- so the only reliable way to make every one of Claude + # Code's own tiers hit the auto-router is to override the env vars it reads per tier. merged = merge_claude_settings_static_token({}, "http://127.0.0.1:4000", "token-abc") for key in ANTHROPIC_DEFAULT_MODEL_ENV_KEYS: assert merged["env"][key] == "autorouter" diff --git a/tests/test_litellm/proxy/shunt_endpoints/test_shunt_worker_endpoints.py b/tests/test_litellm/proxy/shunt_endpoints/test_shunt_worker_endpoints.py index e529153e346..b3d927ba0b2 100644 --- a/tests/test_litellm/proxy/shunt_endpoints/test_shunt_worker_endpoints.py +++ b/tests/test_litellm/proxy/shunt_endpoints/test_shunt_worker_endpoints.py @@ -152,10 +152,19 @@ class TestWorkerTextGoesThroughTheSharedPipeline: async def _fake_pre_call_logic(self, **kwargs): return self.data, object() + # Mirrors route_request's real contract: awaiting it resolves the deployment and + # hands back the provider coroutine *unawaited*, so the caller must await twice. + # A fake that returned the ModelResponse directly would pass against a caller that + # forgets the second await and hands a raw coroutine to the rest of the pipeline. async def _fake_route_request(**kwargs): from litellm.types.utils import Choices, Message, ModelResponse - return ModelResponse(choices=[Choices(index=0, message=Message(role="assistant", content=response_text))]) + async def _provider_call(): + return ModelResponse( + choices=[Choices(index=0, message=Message(role="assistant", content=response_text))] + ) + + return _provider_call() monkeypatch.setattr( "litellm.proxy.shunt_endpoints.endpoints.ProxyBaseLLMRequestProcessing.common_processing_pre_call_logic", @@ -229,6 +238,31 @@ class TestUploadLimits: await endpoints_mod._read_upload_texts([self._upload("big.py", b"x" * 50)]) assert exc_info.value.code == "400" + @pytest.mark.asyncio + async def test_a_misconfigured_negative_per_file_limit_still_bounds_the_read(self, monkeypatch): + """A negative _MAX_UPLOAD_BYTES_PER_FILE (e.g. from a bad env var override) must never + reach UploadFile.read(): a negative size there means "read the whole file", which would + silently defeat this limit for every upload rather than enforce it. The constant is + range-validated at import time via get_env_int_in_range specifically to prevent this, + but this asserts the read-path behavior directly regardless of how the value got here.""" + monkeypatch.setattr(endpoints_mod, "_MAX_UPLOAD_BYTES_PER_FILE", -5) + upload = self._upload("big.py", b"this must never be read without a positive bound") + real_read = upload.read + read_calls: list[int] = [] + + async def _tracking_read(size: int = -1): + read_calls.append(size) + return await real_read(size) + + upload.read = _tracking_read # rebind-ok: test spy on this one instance + with pytest.raises(ProxyException) as exc_info: + await endpoints_mod._read_upload_texts([upload]) + assert exc_info.value.code == "400" + assert all(size > 0 for size in read_calls), ( + f"read() was called with a non-positive size in {read_calls}, " + "which UploadFile.read() treats as 'read the whole file'" + ) + @pytest.mark.asyncio async def test_files_under_the_limits_are_read_in_full(self, monkeypatch): monkeypatch.setattr(endpoints_mod, "_MAX_UPLOAD_BYTES_PER_FILE", 100)