mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
fix(live): send Live HTTP through the shared async handler
Greptile flagged that `LiveTransport.request` reached past LiteLLM's HTTP handler onto `AsyncHTTPHandler.client` to call `request` directly, which is the custom request path the repository rules ask us to avoid. The GET and POST verbs Live uses now go through `AsyncHTTPHandler.get` and `AsyncHTTPHandler.post`, with redirect-following turned off on the shared handler itself rather than per call, so the connection pool stays shared. The handler raises `MaskedHTTPStatusError` for a non-2xx POST, so that error is turned back into its `httpx.Response`: the proxy keeps answering with the upstream status, body and `x-request-id`, which the parametrised regressions now assert for 204, 404 and 503 on every Live operation. `http_client` became `http_handler` on the constructor, the only caller being the proxy, which never passed it.
This commit is contained in:
parent
c88ced1cae
commit
8c93bf3ee4
2 changed files with 71 additions and 29 deletions
|
|
@ -17,6 +17,7 @@ from litellm.llms.chatgpt.realtime import (
|
|||
realtime_headers,
|
||||
)
|
||||
from litellm.llms.custom_httpx.http_handler import (
|
||||
AsyncHTTPHandler,
|
||||
get_async_httpx_client, # pyright: ignore[reportUnknownVariableType] # shared client has legacy untyped optional params
|
||||
get_shared_realtime_ssl_context,
|
||||
)
|
||||
|
|
@ -84,10 +85,10 @@ class LiveTransport:
|
|||
deployment: LiveDeployment,
|
||||
inbound_headers: Mapping[str, str],
|
||||
*,
|
||||
http_client: httpx.AsyncClient | None = None,
|
||||
http_handler: AsyncHTTPHandler | None = None,
|
||||
) -> None:
|
||||
self.deployment = deployment
|
||||
self._http_client = http_client
|
||||
self._http_handler = http_handler
|
||||
params: Final = GenericLiteLLMParams.model_validate(
|
||||
MappingProxyType(
|
||||
{
|
||||
|
|
@ -156,20 +157,22 @@ class LiveTransport:
|
|||
):
|
||||
raise ValueError("Invalid Live HTTP operation")
|
||||
url: Final = self._url(path, query, websocket=False)
|
||||
client: Final = (
|
||||
self._http_client
|
||||
or get_async_httpx_client(
|
||||
llm_provider=LlmProviders.CHATGPT if self.deployment.provider == "chatgpt" else LlmProviders.OPENAI
|
||||
).client
|
||||
)
|
||||
return await client.request(
|
||||
method,
|
||||
url,
|
||||
headers=MappingProxyType({**self._headers, "content-type": "application/json"}),
|
||||
json=dict(body) if body is not None else None, # mutable-ok: JSON encoder requires a concrete dict
|
||||
timeout=60,
|
||||
follow_redirects=False,
|
||||
handler: Final = self._http_handler or get_async_httpx_client(
|
||||
llm_provider=LlmProviders.CHATGPT if self.deployment.provider == "chatgpt" else LlmProviders.OPENAI,
|
||||
params={"follow_redirects": False},
|
||||
)
|
||||
headers: Final = {**self._headers, "content-type": "application/json"}
|
||||
if method == "GET":
|
||||
return await handler.get(url, headers=headers, timeout=60, follow_redirects=False)
|
||||
try:
|
||||
return await handler.post(
|
||||
url,
|
||||
headers=headers,
|
||||
json=dict(body) if body is not None else None, # mutable-ok: JSON encoder requires a concrete dict
|
||||
timeout=60,
|
||||
)
|
||||
except httpx.HTTPStatusError as error:
|
||||
return error.response
|
||||
|
||||
async def connect(self, path: str, query: LiveQuery | None = None) -> "ClientConnection":
|
||||
import websockets
|
||||
|
|
|
|||
|
|
@ -1,9 +1,21 @@
|
|||
import json
|
||||
from contextlib import asynccontextmanager
|
||||
from unittest.mock import AsyncMock
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
|
||||
from litellm.llms.chatgpt.live import LiveDeployment, LiveTransport, live_session_path
|
||||
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
|
||||
|
||||
|
||||
@asynccontextmanager
|
||||
async def live_handler(respond):
|
||||
handler = AsyncHTTPHandler(transport=httpx.MockTransport(respond), follow_redirects=False)
|
||||
try:
|
||||
yield handler
|
||||
finally:
|
||||
await handler.close()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -32,7 +44,7 @@ async def test_live_request_preserves_payload_status_and_selected_credentials(pr
|
|||
assert not {"model", "call_id", "session_id", "api_key"}.intersection(request.url.params)
|
||||
return httpx.Response(status, json={"result": "upstream"}, headers={"x-request-id": "provider-id"})
|
||||
|
||||
async with httpx.AsyncClient(transport=httpx.MockTransport(respond)) as client:
|
||||
async with live_handler(respond) as handler:
|
||||
transport = LiveTransport(
|
||||
LiveDeployment(
|
||||
model="deployment-model",
|
||||
|
|
@ -43,7 +55,7 @@ async def test_live_request_preserves_payload_status_and_selected_credentials(pr
|
|||
extra_query={"gateway": "trusted", "tag": ("a +/&", "b"), "model": "bad", "session_id": "bad"},
|
||||
),
|
||||
{"Authorization": "Bearer proxy-key", "Cookie": "private", "OpenAI-Beta": "feature=v1"},
|
||||
http_client=client,
|
||||
http_handler=handler,
|
||||
)
|
||||
response = await transport.request(
|
||||
"POST",
|
||||
|
|
@ -54,26 +66,53 @@ async def test_live_request_preserves_payload_status_and_selected_credentials(pr
|
|||
assert response.status_code == status
|
||||
assert response.json() == {"result": "upstream"}
|
||||
assert response.headers["x-request-id"] == "provider-id"
|
||||
assert not client.is_closed
|
||||
assert not handler.client.is_closed
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("operation", ["fork", "accept", "reject", "refer", "hangup", "content"])
|
||||
async def test_live_all_http_operations(operation):
|
||||
@pytest.mark.parametrize("status", [204, 404, 503])
|
||||
async def test_live_all_http_operations(operation, status):
|
||||
def respond(request):
|
||||
assert request.url.path == f"/v1/live/sessions/sess_new-ID/{operation}"
|
||||
assert request.method == ("GET" if operation == "content" else "POST")
|
||||
assert request.url.params["output_format"] == "json"
|
||||
return httpx.Response(204)
|
||||
return httpx.Response(status)
|
||||
|
||||
async with httpx.AsyncClient(transport=httpx.MockTransport(respond)) as client:
|
||||
transport = LiveTransport(LiveDeployment("model", provider="openai", api_key="key"), {}, http_client=client)
|
||||
async with live_handler(respond) as handler:
|
||||
transport = LiveTransport(LiveDeployment("model", provider="openai", api_key="key"), {}, http_handler=handler)
|
||||
response = await transport.request(
|
||||
"GET" if operation == "content" else "POST",
|
||||
live_session_path("sess_new-ID", operation),
|
||||
query={"output_format": "json"},
|
||||
)
|
||||
assert response.status_code == 204
|
||||
assert response.status_code == status
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_live_request_uses_handler_methods():
|
||||
handler = AsyncMock(spec=AsyncHTTPHandler)
|
||||
handler.get.return_value = httpx.Response(404)
|
||||
handler.post.return_value = httpx.Response(503)
|
||||
transport = LiveTransport(LiveDeployment("model", provider="openai", api_key="key"), {}, http_handler=handler)
|
||||
|
||||
get_response = await transport.request("GET", live_session_path("sess_1", "content"))
|
||||
post_response = await transport.request("POST", "live/sessions", {"transport": {"type": "webrtc"}})
|
||||
|
||||
assert get_response.status_code == 404
|
||||
assert post_response.status_code == 503
|
||||
handler.get.assert_awaited_once_with(
|
||||
"https://api.openai.com/v1/live/sessions/sess_1/content",
|
||||
headers={"authorization": "Bearer key", "content-type": "application/json"},
|
||||
timeout=60,
|
||||
follow_redirects=False,
|
||||
)
|
||||
handler.post.assert_awaited_once_with(
|
||||
"https://api.openai.com/v1/live/sessions",
|
||||
headers={"authorization": "Bearer key", "content-type": "application/json"},
|
||||
json={"transport": {"type": "webrtc"}},
|
||||
timeout=60,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -149,8 +188,8 @@ async def test_live_preserves_opaque_session_ids(session_id):
|
|||
assert request.url.params == httpx.QueryParams()
|
||||
return httpx.Response(200, json={"session_id": session_id})
|
||||
|
||||
async with httpx.AsyncClient(transport=httpx.MockTransport(respond)) as client:
|
||||
transport = LiveTransport(LiveDeployment("model", provider="openai", api_key="key"), {}, http_client=client)
|
||||
async with live_handler(respond) as handler:
|
||||
transport = LiveTransport(LiveDeployment("model", provider="openai", api_key="key"), {}, http_handler=handler)
|
||||
response = await transport.request("GET", live_session_path(session_id, "content"))
|
||||
assert response.json()["session_id"] == session_id
|
||||
|
||||
|
|
@ -161,8 +200,8 @@ async def test_live_does_not_redirect_credentials():
|
|||
assert request.url.host == "api.openai.com"
|
||||
return httpx.Response(307, headers={"location": "https://elsewhere.example/collect"})
|
||||
|
||||
async with httpx.AsyncClient(transport=httpx.MockTransport(respond), follow_redirects=True) as client:
|
||||
transport = LiveTransport(LiveDeployment("model", provider="openai", api_key="key"), {}, http_client=client)
|
||||
async with live_handler(respond) as handler:
|
||||
transport = LiveTransport(LiveDeployment("model", provider="openai", api_key="key"), {}, http_handler=handler)
|
||||
response = await transport.request("POST", "live/sessions", {})
|
||||
assert response.status_code == 307
|
||||
|
||||
|
|
@ -211,11 +250,11 @@ async def test_live_rejects_invalid_api_base_before_network(api_base):
|
|||
requests.append(request)
|
||||
return httpx.Response(200, json={})
|
||||
|
||||
async with httpx.AsyncClient(transport=httpx.MockTransport(respond)) as client:
|
||||
async with live_handler(respond) as handler:
|
||||
transport = LiveTransport(
|
||||
LiveDeployment("deployment-model", provider="openai", api_key="deployment-key", api_base=api_base),
|
||||
{},
|
||||
http_client=client,
|
||||
http_handler=handler,
|
||||
)
|
||||
with pytest.raises(ValueError, match="Invalid Live API base"):
|
||||
await transport.request("POST", "live/sessions", {})
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue