From 1cd481d4f2e9d236b77ea61cf5c0bbe0e9ee4c46 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Tue, 4 Aug 2026 14:36:04 -0700 Subject: [PATCH] fix(proxy): enforce per-model budgets against resolved cursor model variants --- .../proxy/response_api_endpoints/endpoints.py | 30 +++- .../response_api_endpoints/test_endpoints.py | 159 +++++++++++++++++- 2 files changed, 180 insertions(+), 9 deletions(-) diff --git a/litellm/proxy/response_api_endpoints/endpoints.py b/litellm/proxy/response_api_endpoints/endpoints.py index 5f4a386c8c1..f752986d7f4 100644 --- a/litellm/proxy/response_api_endpoints/endpoints.py +++ b/litellm/proxy/response_api_endpoints/endpoints.py @@ -19,6 +19,10 @@ from litellm.proxy.auth.user_api_key_auth import ( user_api_key_auth_websocket, ) from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing +from litellm.proxy.common_utils.http_parsing_utils import ( + _read_request_body, + _safe_set_request_parsed_body, +) from litellm.types.llms.openai import REASONING_EFFORT, ResponseAPIUsage, ResponsesAPIResponse from litellm.types.responses.main import DeleteResponseResult @@ -148,6 +152,18 @@ def _resolve_cursor_model_variant( return {**resolved, "reasoning": {"effort": variant.reasoning_effort}} # mutable-ok: plain body dict +async def _resolve_cursor_model_variant_before_auth(request: Request) -> None: + from litellm.proxy.proxy_server import llm_router + + try: + raw_body: Final = await _read_request_body(request=request) + except (json.JSONDecodeError, ProxyException): + return + resolved: Final = _resolve_cursor_model_variant(raw_body, llm_router) + if resolved is not raw_body: + _safe_set_request_parsed_body(request=request, parsed_body=resolved) + + @router.post( "/v1/responses", dependencies=[Depends(user_api_key_auth)], @@ -440,7 +456,10 @@ async def cursor_model_list( @router.post( "/cursor/chat/completions", - dependencies=[Depends(user_api_key_auth)], + dependencies=[ + Depends(_resolve_cursor_model_variant_before_auth), + Depends(user_api_key_auth), + ], tags=["responses"], ) async def cursor_chat_completions( @@ -479,9 +498,7 @@ async def cursor_chat_completions( responses_api_bridge, ) from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper - from litellm.proxy.common_utils.http_parsing_utils import _safe_set_request_parsed_body from litellm.proxy.proxy_server import ( - _read_request_body, async_data_generator, chat_completion, general_settings, @@ -499,14 +516,13 @@ async def cursor_chat_completions( from litellm.types.utils import ModelResponse raw_body: Final = await _read_request_body(request=request) - data = _resolve_cursor_model_variant(raw_body, llm_router) - if _is_chat_completions_body(data): + if _is_chat_completions_body(raw_body): # Genuine chat completions body (Cursor sends these for models whose BYOK it # already fixed); delegate so behavior matches /chat/completions exactly. # Keyed on messages CONTENT, not key presence: Cursor can send a null or # empty messages stub alongside a real agent-mode input array - normalized: Final = _normalize_tool_dialect(data, to_chat=True) + normalized: Final = _normalize_tool_dialect(raw_body, to_chat=True) if normalized is not raw_body: _safe_set_request_parsed_body(request=request, parsed_body=normalized) return await chat_completion( @@ -521,7 +537,7 @@ async def cursor_chat_completions( # Rebuild rather than pop: _read_request_body can return the request-scope # cached parsed-body dict itself, and removing keys from it corrupts the # cache's key snapshot so later readers get an empty body - data = {key: value for key, value in data.items() if key != "stream_options"} # mutable-ok: plain body dict + data = {key: value for key, value in raw_body.items() if key != "stream_options"} # mutable-ok: plain body dict data = _normalize_tool_dialect(data, to_chat=False) diff --git a/tests/test_litellm/proxy/response_api_endpoints/test_endpoints.py b/tests/test_litellm/proxy/response_api_endpoints/test_endpoints.py index 60168e7f912..a064c8de985 100644 --- a/tests/test_litellm/proxy/response_api_endpoints/test_endpoints.py +++ b/tests/test_litellm/proxy/response_api_endpoints/test_endpoints.py @@ -851,8 +851,9 @@ def test_cursor_chat_completions_input_body_uses_responses_pipeline_and_strips_s app.dependency_overrides[user_api_key_auth] = _auth_override try: - with patch.object(ps, "llm_router", mock_router), patch.object( - ps, "_read_request_body", side_effect=capturing_read_request_body + with patch.object(ps, "llm_router", mock_router), patch( + "litellm.proxy.response_api_endpoints.endpoints._read_request_body", + side_effect=capturing_read_request_body, ): client = TestClient(app) response = client.post( @@ -1568,3 +1569,157 @@ class TestCursorModelSuffixResolutionEndToEnd: assert mock_router.aresponses.call_args is not None assert mock_router.aresponses.call_args.kwargs["model"] == "claude-opus-5" assert mock_router.aresponses.call_args.kwargs["reasoning"] == {"effort": "high"} + + +def _cursor_budget_auth_env(base_model: str, spend: float): + from litellm import Router + from litellm.caching.dual_cache import DualCache + from litellm.proxy._types import UserAPIKeyAuth + from litellm.proxy.hooks.model_max_budget_limiter import ( + VIRTUAL_KEY_SPEND_CACHE_KEY_PREFIX, + _PROXY_VirtualKeyModelMaxBudgetLimiter, + ) + + valid_token = UserAPIKeyAuth( + api_key="sk-cursor-budget-test", + token="hashed-cursor-budget-token", + model_max_budget={base_model: {"budget_limit": 0.00001, "time_period": "1d"}}, + ) + limiter = _PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=DualCache()) + limiter.dual_cache.in_memory_cache.set_cache( + key=f"{VIRTUAL_KEY_SPEND_CACHE_KEY_PREFIX}:{valid_token.token}:{base_model}:1d", + value=spend, + ) + router = Router( + model_list=[{"model_name": "anthropic/*", "litellm_params": {"model": "anthropic/*", "api_key": "fake"}}] + ) + + mock_proxy_logging_obj = MagicMock() + mock_proxy_logging_obj.post_call_failure_hook = AsyncMock(return_value=None) + + proxy_server_attrs = { + "prisma_client": MagicMock(), + "user_api_key_cache": DualCache(), + "proxy_logging_obj": mock_proxy_logging_obj, + "master_key": "sk-master-key", + "general_settings": {}, + "llm_model_list": [], + "llm_router": router, + "open_telemetry_logger": None, + "model_max_budget_limiter": limiter, + "user_custom_auth": None, + "jwt_handler": None, + "litellm_proxy_admin_name": "admin", + } + return valid_token, proxy_server_attrs + + +def _post_cursor_with_real_auth(valid_token, proxy_server_attrs, request_model: str): + with ( + patch.multiple("litellm.proxy.proxy_server", **proxy_server_attrs), + patch( + "litellm.proxy.auth.resolvers.store.IdentityStore._resolve_key", + new_callable=AsyncMock, + return_value=valid_token, + ), + ): + client = TestClient(app) + return client.post( + "/cursor/chat/completions", + json={"model": request_model, "input": [{"role": "user", "content": "hi"}]}, + headers={"Authorization": "Bearer sk-cursor-budget-test"}, + ) + + +class TestCursorVariantPerModelBudgetEnforcement: + """Regression tests for the per-model budget bypass on /cursor/chat/completions. + + user_api_key_auth enforced key model_max_budget against the raw request model, + but _resolve_cursor_model_variant only rewrote minted aliases like + -thinking- to inside the handler, after auth had already + run. A key whose budget for was exhausted could keep calling + through any unconfigured alias. The variant must now be resolved in a + route-level dependency that runs before user_api_key_auth, so these tests + exercise the real dependency chain (real auth, real budget limiter) through + TestClient and fail if that ordering ever breaks.""" + + def test_minted_alias_rejected_when_base_model_budget_exhausted(self): + valid_token, attrs = _cursor_budget_auth_env(base_model="claude-opus-5", spend=1.0) + + response = _post_cursor_with_real_auth(valid_token, attrs, request_model="claude-opus-5-thinking-high") + + assert response.status_code == 429, response.text + error = response.json()["error"] + assert error["type"] == "budget_exceeded" + assert "exceeded budget for model=claude-opus-5" in error["message"] + + def test_alias_rejection_matches_base_model_rejection(self): + valid_token, attrs = _cursor_budget_auth_env(base_model="claude-opus-5", spend=1.0) + + base_response = _post_cursor_with_real_auth(valid_token, attrs, request_model="claude-opus-5") + alias_response = _post_cursor_with_real_auth(valid_token, attrs, request_model="claude-opus-5-fast") + + assert base_response.status_code == 429, base_response.text + assert alias_response.status_code == 429, alias_response.text + assert alias_response.json() == base_response.json() + + +class TestCursorVariantResolvedBeforeAuth: + """The route-level resolver dependency must rewrite the parsed body before + user_api_key_auth reads it, so every auth check (model access, key and + end-user model budgets, rate limits) sees the base model, and names the + router already serves must reach auth untouched.""" + + def _run_with_recording_auth(self, mock_router, request_model: str): + from litellm.proxy._types import UserAPIKeyAuth + from litellm.proxy.auth.user_api_key_auth import user_api_key_auth + from litellm.proxy.common_utils.http_parsing_utils import _read_request_body + + from fastapi import Request + + bodies_seen_by_auth = [] + + async def recording_auth(request: Request) -> UserAPIKeyAuth: + bodies_seen_by_auth.append(await _read_request_body(request=request)) + return UserAPIKeyAuth(api_key="sk-test-cursor") + + async def fake_chat_completion(request, fastapi_response, model, user_api_key_dict): + return {"id": "chatcmpl-fake", "object": "chat.completion", "choices": []} + + app.dependency_overrides[user_api_key_auth] = recording_auth + try: + with ( + patch("litellm.proxy.proxy_server.llm_router", new=mock_router), + patch("litellm.proxy.proxy_server.chat_completion", new=fake_chat_completion), + ): + client = TestClient(app) + response = client.post( + "/cursor/chat/completions", + json={"model": request_model, "messages": [{"role": "user", "content": "hi"}]}, + headers={"Authorization": "Bearer sk-test-cursor"}, + ) + finally: + app.dependency_overrides.pop(user_api_key_auth, None) + + assert response.status_code == 200, response.text + assert len(bodies_seen_by_auth) == 1 + return bodies_seen_by_auth[0] + + def test_auth_sees_base_model_for_minted_alias(self): + auth_body = self._run_with_recording_auth( + mock_router=_router_serving_only("claude-opus-5"), + request_model="claude-opus-5-thinking-xhigh-fast", + ) + assert auth_body["model"] == "claude-opus-5" + assert auth_body["reasoning_effort"] == "xhigh" + + def test_auth_sees_servable_model_name_untouched(self): + mock_router = _router_serving_only("claude-opus-5") + mock_router.model_names = {"claude-opus-5-thinking-high"} + + auth_body = self._run_with_recording_auth( + mock_router=mock_router, + request_model="claude-opus-5-thinking-high", + ) + assert auth_body["model"] == "claude-opus-5-thinking-high" + assert "reasoning_effort" not in auth_body