mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
fix(proxy): enforce per-model budgets against resolved cursor model variants
This commit is contained in:
parent
27885076e7
commit
1cd481d4f2
2 changed files with 180 additions and 9 deletions
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue