fix(proxy): enforce per-model budgets against resolved cursor model variants

This commit is contained in:
mateo-berri 2026-08-04 14:36:04 -07:00
parent 27885076e7
commit 1cd481d4f2
2 changed files with 180 additions and 9 deletions

View file

@ -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)

View file

@ -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
<base>-thinking-<level> to <base> inside the handler, after auth had already
run. A key whose budget for <base> was exhausted could keep calling <base>
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