mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
fix(proxy): stop /utils/transform_request from calling the provider and blocking the event loop (#33954)
Co-authored-by: yassin <yassin@berri.ai> Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
4584958574
commit
5afb80742d
5 changed files with 91 additions and 3 deletions
|
|
@ -385,6 +385,10 @@ def _get_cached_prometheus_logger():
|
|||
return _PrometheusLogger
|
||||
|
||||
|
||||
class RawRequestCaptured(Exception):
|
||||
pass
|
||||
|
||||
|
||||
_DEPLOYMENT_PRICING_KEYS: Final = (
|
||||
"input_cost_per_token",
|
||||
"output_cost_per_token",
|
||||
|
|
@ -591,6 +595,7 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
kwargs: dict | None = None,
|
||||
log_raw_request_response: bool = False,
|
||||
supports_correlation_logging: bool = True,
|
||||
raw_request_only: bool = False,
|
||||
):
|
||||
_input: Final[str | None] = messages # save original value of messages
|
||||
if messages is not None:
|
||||
|
|
@ -650,6 +655,7 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
self.streaming_chunks: list[Any] = [] # for generating complete stream response
|
||||
self.sync_streaming_chunks: list[Any] = [] # for generating complete stream response
|
||||
self.log_raw_request_response = log_raw_request_response
|
||||
self.raw_request_only = raw_request_only
|
||||
|
||||
# Initialize dynamic callbacks
|
||||
self.dynamic_input_callbacks: list[str | Callable | CustomLogger] | None = dynamic_input_callbacks
|
||||
|
|
@ -1476,6 +1482,9 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
if capture_exception: # log this error to sentry for debugging
|
||||
capture_exception(e)
|
||||
|
||||
if self.raw_request_only:
|
||||
raise RawRequestCaptured()
|
||||
|
||||
def _print_llm_call_debugging_log(
|
||||
self,
|
||||
api_base: str,
|
||||
|
|
|
|||
|
|
@ -13766,7 +13766,7 @@ async def transform_request(request: TransformRequestBody):
|
|||
except ValueError as e:
|
||||
raise HTTPException(status_code=400, detail={"error": str(e)})
|
||||
|
||||
return return_raw_request(endpoint=request.call_type, kwargs=request.request_body)
|
||||
return await asyncio.to_thread(return_raw_request, request.call_type, request.request_body)
|
||||
|
||||
|
||||
async def _check_if_model_is_user_added(
|
||||
|
|
|
|||
|
|
@ -10273,7 +10273,7 @@ def return_raw_request(endpoint: CallTypes, kwargs: dict) -> RawRequestTypedDict
|
|||
"""
|
||||
from datetime import datetime
|
||||
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging, RawRequestCaptured
|
||||
|
||||
litellm_logging_obj: Final = Logging(
|
||||
model="gpt-3.5-turbo",
|
||||
|
|
@ -10284,6 +10284,7 @@ def return_raw_request(endpoint: CallTypes, kwargs: dict) -> RawRequestTypedDict
|
|||
start_time=datetime.now(),
|
||||
function_id="1234",
|
||||
log_raw_request_response=True,
|
||||
raw_request_only=True,
|
||||
)
|
||||
|
||||
llm_api_endpoint: Final = getattr(litellm, endpoint.value)
|
||||
|
|
@ -10294,7 +10295,11 @@ def return_raw_request(endpoint: CallTypes, kwargs: dict) -> RawRequestTypedDict
|
|||
llm_api_endpoint(
|
||||
**kwargs,
|
||||
litellm_logging_obj=litellm_logging_obj,
|
||||
api_key="my-fake-api-key", # 👈 ensure the request fails
|
||||
api_key="my-fake-api-key",
|
||||
)
|
||||
except RawRequestCaptured:
|
||||
received_exception = (
|
||||
"raw request was not captured before the provider call; check the proxy logs for the pre_call error"
|
||||
)
|
||||
except Exception as e:
|
||||
received_exception = str(e)
|
||||
|
|
|
|||
|
|
@ -11168,6 +11168,49 @@ class TestTransformRequestBannedParams:
|
|||
)
|
||||
|
||||
|
||||
class TestTransformRequestOffEventLoop:
|
||||
@pytest.fixture
|
||||
def client(self):
|
||||
mock_auth = UserAPIKeyAuth(user_id="test-internal", user_role=LitellmUserRoles.INTERNAL_USER)
|
||||
original = app.dependency_overrides.copy()
|
||||
app.dependency_overrides[user_api_key_auth] = lambda: mock_auth
|
||||
try:
|
||||
yield TestClient(app)
|
||||
finally:
|
||||
app.dependency_overrides = original
|
||||
|
||||
def test_transform_request_runs_return_raw_request_off_the_event_loop(self, client, monkeypatch):
|
||||
import litellm.utils
|
||||
from litellm.types.utils import RawRequestTypedDict
|
||||
|
||||
seen: dict[str, bool] = {}
|
||||
|
||||
def fake_return_raw_request(endpoint, kwargs):
|
||||
try:
|
||||
asyncio.get_running_loop()
|
||||
seen["on_event_loop"] = True
|
||||
except RuntimeError:
|
||||
seen["on_event_loop"] = False
|
||||
return RawRequestTypedDict(
|
||||
raw_request_api_base="https://api.openai.com/v1/",
|
||||
raw_request_body=kwargs,
|
||||
raw_request_headers={},
|
||||
error=None,
|
||||
)
|
||||
|
||||
monkeypatch.setattr(litellm.utils, "return_raw_request", fake_return_raw_request)
|
||||
response = client.post(
|
||||
"/utils/transform_request",
|
||||
json={
|
||||
"call_type": "completion",
|
||||
"request_body": {"model": "gpt-5.6-sol", "messages": [{"role": "user", "content": "hi"}]},
|
||||
},
|
||||
)
|
||||
assert response.status_code == 200, response.text
|
||||
assert response.json()["raw_request_body"]["model"] == "gpt-5.6-sol"
|
||||
assert seen == {"on_event_loop": False}, "return_raw_request ran on the event loop thread"
|
||||
|
||||
|
||||
class TestSortModelsByDisplayName:
|
||||
"""Regression: team BYOK rows persist an internal `model_name` like
|
||||
`model_name_{team_id}_{uuid}` and expose the user-facing name via
|
||||
|
|
|
|||
|
|
@ -713,6 +713,37 @@ def _mocked_openai_chat_response(model: str) -> httpx.Response:
|
|||
)
|
||||
|
||||
|
||||
def test_return_raw_request_does_not_call_provider(respx_mock: respx.MockRouter):
|
||||
"""Regression for #33952: return_raw_request must transform without contacting the provider.
|
||||
|
||||
Previously return_raw_request invoked the real endpoint with a fake key and relied on the
|
||||
provider rejecting it, which sent an unintended inference request and (in the async proxy
|
||||
route) blocked the event loop on provider I/O.
|
||||
"""
|
||||
from litellm.types.utils import CallTypes
|
||||
from litellm.utils import return_raw_request
|
||||
|
||||
model = "gpt-4o"
|
||||
route = respx_mock.post("https://api.openai.com/v1/chat/completions").mock(
|
||||
return_value=_mocked_openai_chat_response(model)
|
||||
)
|
||||
|
||||
request = return_raw_request(
|
||||
endpoint=CallTypes.completion,
|
||||
kwargs={
|
||||
"model": model,
|
||||
"messages": [{"role": "user", "content": "hi"}],
|
||||
},
|
||||
)
|
||||
|
||||
assert route.call_count == 0
|
||||
assert request.get("error") is None
|
||||
assert request["raw_request_body"]["model"] == model
|
||||
assert request["raw_request_body"]["messages"] == [
|
||||
{"role": "user", "content": "hi"}
|
||||
]
|
||||
|
||||
|
||||
def test_completion_forwards_verbosity_in_raw_request(respx_mock: respx.MockRouter):
|
||||
"""Regression test: completion() must forward the verbosity param to the provider request body."""
|
||||
from litellm.types.utils import CallTypes
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue