mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
Merge pull request #22559 from BerriAI/litellm_responses_websocket
[Feat] Add support for Responses Websocket
This commit is contained in:
commit
8764e5da8c
11 changed files with 1162 additions and 6 deletions
|
|
@ -1246,6 +1246,7 @@ from .ocr.main import *
|
|||
from .rag.main import *
|
||||
from .search.main import *
|
||||
from .realtime_api.main import _arealtime
|
||||
from .responses.main import _aresponses_websocket
|
||||
from .fine_tuning.main import *
|
||||
from .files.main import *
|
||||
from .vector_store_files.main import (
|
||||
|
|
|
|||
|
|
@ -69,6 +69,7 @@ from litellm.responses.streaming_iterator import (
|
|||
BaseResponsesAPIStreamingIterator,
|
||||
MockResponsesAPIStreamingIterator,
|
||||
ResponsesAPIStreamingIterator,
|
||||
ResponsesWebSocketStreaming,
|
||||
SyncResponsesAPIStreamingIterator,
|
||||
)
|
||||
from litellm.types.containers.main import (
|
||||
|
|
@ -4731,6 +4732,98 @@ class BaseLLMHTTPHandler:
|
|||
f"Unexpected error while closing WebSocket: {close_error}"
|
||||
)
|
||||
|
||||
async def async_responses_websocket(
|
||||
self,
|
||||
model: str,
|
||||
websocket: Any,
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
responses_api_provider_config: BaseResponsesAPIConfig,
|
||||
api_base: Optional[str] = None,
|
||||
api_key: Optional[str] = None,
|
||||
timeout: Optional[float] = None,
|
||||
user_api_key_dict: Optional[Any] = None,
|
||||
litellm_metadata: Optional[Dict[str, Any]] = None,
|
||||
):
|
||||
"""
|
||||
Handles Responses API WebSocket mode.
|
||||
|
||||
Opens a persistent WebSocket to the provider's /v1/responses endpoint
|
||||
and proxies response.create events bidirectionally for lower-latency
|
||||
agentic workflows.
|
||||
"""
|
||||
import websockets
|
||||
from websockets.asyncio.client import ClientConnection
|
||||
|
||||
litellm_params = GenericLiteLLMParams()
|
||||
headers = responses_api_provider_config.validate_environment(
|
||||
headers={},
|
||||
model=model,
|
||||
litellm_params=litellm_params,
|
||||
)
|
||||
if api_key:
|
||||
headers["Authorization"] = f"Bearer {api_key}"
|
||||
|
||||
http_url = responses_api_provider_config.get_complete_url(
|
||||
api_base=api_base,
|
||||
litellm_params={},
|
||||
)
|
||||
# /responses -> wss:// URL
|
||||
ws_url = http_url.replace("https://", "wss://").replace("http://", "ws://")
|
||||
|
||||
try:
|
||||
ssl_context = get_shared_realtime_ssl_context()
|
||||
if ws_url.startswith("wss://") and ssl_context is False:
|
||||
ssl_context = ssl.SSLContext(ssl.PROTOCOL_TLS_CLIENT)
|
||||
ssl_context.check_hostname = False
|
||||
ssl_context.verify_mode = ssl.CERT_NONE
|
||||
|
||||
logging_obj.pre_call(
|
||||
input=None,
|
||||
api_key=api_key or "",
|
||||
additional_args={
|
||||
"api_base": ws_url,
|
||||
"headers": headers,
|
||||
"complete_input_dict": {"mode": "responses_websocket"},
|
||||
},
|
||||
)
|
||||
|
||||
async with websockets.connect( # type: ignore
|
||||
ws_url,
|
||||
additional_headers=headers,
|
||||
max_size=REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES,
|
||||
ssl=ssl_context,
|
||||
) as backend_ws:
|
||||
_request_data: Dict[str, Any] = {}
|
||||
if litellm_metadata:
|
||||
_request_data["litellm_metadata"] = litellm_metadata
|
||||
streaming = ResponsesWebSocketStreaming(
|
||||
websocket=websocket,
|
||||
backend_ws=cast(ClientConnection, backend_ws),
|
||||
logging_obj=logging_obj,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
request_data=_request_data,
|
||||
)
|
||||
await streaming.bidirectional_forward()
|
||||
|
||||
except websockets.exceptions.InvalidStatusCode as e: # type: ignore
|
||||
verbose_logger.exception(f"Error connecting to responses WS backend: {e}")
|
||||
await websocket.close(code=e.status_code, reason=str(e))
|
||||
except Exception as e:
|
||||
verbose_logger.exception(f"Error in responses WS: {e}")
|
||||
try:
|
||||
await websocket.close(
|
||||
code=1011, reason=f"Internal server error: {str(e)}"
|
||||
)
|
||||
except RuntimeError as close_error:
|
||||
if "already completed" in str(close_error) or "websocket.close" in str(
|
||||
close_error
|
||||
):
|
||||
pass
|
||||
else:
|
||||
raise Exception(
|
||||
f"Unexpected error while closing WebSocket: {close_error}"
|
||||
)
|
||||
|
||||
def image_edit_handler(
|
||||
self,
|
||||
model: str,
|
||||
|
|
|
|||
|
|
@ -499,6 +499,7 @@ class ProxyBaseLLMRequestProcessing:
|
|||
"aembedding",
|
||||
"aresponses",
|
||||
"_arealtime",
|
||||
"_aresponses_websocket",
|
||||
"aget_responses",
|
||||
"adelete_responses",
|
||||
"acancel_responses",
|
||||
|
|
|
|||
|
|
@ -1,14 +1,21 @@
|
|||
import asyncio
|
||||
import json
|
||||
import time
|
||||
from typing import Any, AsyncIterator, Optional, cast
|
||||
from typing import Any, AsyncIterator, Dict, Optional, cast
|
||||
from uuid import uuid4
|
||||
|
||||
import fastapi
|
||||
from fastapi import APIRouter, Depends, HTTPException, Request, Response
|
||||
from starlette.websockets import WebSocket
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.integrations.custom_guardrail import ModifyResponseException
|
||||
from litellm.proxy._types import *
|
||||
from litellm.proxy.auth.user_api_key_auth import UserAPIKeyAuth, user_api_key_auth
|
||||
from litellm.proxy.auth.user_api_key_auth import (
|
||||
UserAPIKeyAuth,
|
||||
user_api_key_auth,
|
||||
user_api_key_auth_websocket,
|
||||
)
|
||||
from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing
|
||||
from litellm.types.llms.openai import ResponseAPIUsage, ResponsesAPIResponse
|
||||
from litellm.types.responses.main import DeleteResponseResult
|
||||
|
|
@ -904,3 +911,121 @@ async def cancel_response(
|
|||
proxy_logging_obj=proxy_logging_obj,
|
||||
version=version,
|
||||
)
|
||||
|
||||
|
||||
@router.websocket("/v1/responses")
|
||||
@router.websocket("/responses")
|
||||
async def responses_websocket_endpoint(
|
||||
websocket: WebSocket,
|
||||
model: str = fastapi.Query(
|
||||
..., description="The model to use for the responses WebSocket session."
|
||||
),
|
||||
user_api_key_dict=Depends(user_api_key_auth_websocket),
|
||||
):
|
||||
"""
|
||||
Responses API WebSocket mode endpoint.
|
||||
|
||||
Keeps a persistent WebSocket connection for response.create events,
|
||||
enabling lower-latency agentic workflows with many tool-call round trips.
|
||||
|
||||
See: https://developers.openai.com/api/docs/guides/websocket-mode/
|
||||
"""
|
||||
from litellm.proxy.proxy_server import (
|
||||
general_settings,
|
||||
llm_router,
|
||||
proxy_config,
|
||||
proxy_logging_obj,
|
||||
user_api_base,
|
||||
user_max_tokens,
|
||||
user_model,
|
||||
user_request_timeout,
|
||||
user_temperature,
|
||||
version,
|
||||
)
|
||||
from litellm.proxy.route_llm_request import route_request
|
||||
|
||||
# Accept the WebSocket handshake
|
||||
requested_protocols = [
|
||||
p.strip()
|
||||
for p in (websocket.headers.get("sec-websocket-protocol") or "").split(",")
|
||||
if p.strip()
|
||||
]
|
||||
accept_kwargs: dict = {}
|
||||
if requested_protocols:
|
||||
accept_kwargs["subprotocol"] = requested_protocols[0]
|
||||
await websocket.accept(**accept_kwargs)
|
||||
|
||||
data: Dict[str, Any] = {
|
||||
"model": model,
|
||||
"websocket": websocket,
|
||||
}
|
||||
|
||||
# Construct a synthetic Request for pre-call processing
|
||||
headers_list = list(websocket.scope.get("headers") or [])
|
||||
scope: Dict[str, Any] = {
|
||||
"type": "http",
|
||||
"method": "POST",
|
||||
"path": "/v1/responses",
|
||||
"headers": headers_list,
|
||||
}
|
||||
request = Request(scope=scope)
|
||||
request._url = websocket.url
|
||||
|
||||
async def return_body():
|
||||
return f'{{"model": "{model}"}}'.encode()
|
||||
|
||||
request.body = return_body # type: ignore
|
||||
|
||||
# Phase 1: pre-call processing (auth, guardrails, rate limits)
|
||||
base_llm_response_processor = ProxyBaseLLMRequestProcessing(data=data)
|
||||
try:
|
||||
(
|
||||
data,
|
||||
litellm_logging_obj,
|
||||
) = await base_llm_response_processor.common_processing_pre_call_logic(
|
||||
request=request,
|
||||
general_settings=general_settings,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
version=version,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
proxy_config=proxy_config,
|
||||
user_model=user_model,
|
||||
user_temperature=user_temperature,
|
||||
user_request_timeout=user_request_timeout,
|
||||
user_max_tokens=user_max_tokens,
|
||||
user_api_base=user_api_base,
|
||||
model=model,
|
||||
route_type="_aresponses_websocket",
|
||||
)
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.exception("Responses WebSocket pre-call error")
|
||||
try:
|
||||
await websocket.send_text(
|
||||
json.dumps(
|
||||
{
|
||||
"type": "error",
|
||||
"error": {
|
||||
"type": "pre_call_error",
|
||||
"message": str(e),
|
||||
},
|
||||
}
|
||||
)
|
||||
)
|
||||
except Exception:
|
||||
pass
|
||||
await websocket.close(code=1011, reason="Pre-call error")
|
||||
return
|
||||
|
||||
# Phase 2: route to upstream provider
|
||||
try:
|
||||
data["user_api_key_dict"] = user_api_key_dict
|
||||
llm_call = await route_request(
|
||||
data=data,
|
||||
route_type="_aresponses_websocket",
|
||||
llm_router=llm_router,
|
||||
user_model=user_model,
|
||||
)
|
||||
await llm_call
|
||||
except Exception:
|
||||
verbose_proxy_logger.exception("Responses WebSocket error")
|
||||
await websocket.close(code=1011, reason="Internal server error")
|
||||
|
|
|
|||
|
|
@ -42,6 +42,7 @@ ROUTE_ENDPOINT_MAPPING = {
|
|||
"amoderation": "/moderations",
|
||||
"arerank": "/rerank",
|
||||
"aresponses": "/responses",
|
||||
"_aresponses_websocket": "/responses",
|
||||
"alist_input_items": "/responses/{response_id}/input_items",
|
||||
"aimage_edit": "/images/edits",
|
||||
"acancel_responses": "/responses/{response_id}/cancel",
|
||||
|
|
@ -163,6 +164,7 @@ async def route_request( # noqa: PLR0915 - Complex routing function, refactorin
|
|||
"acreate_response_reply",
|
||||
"alist_input_items",
|
||||
"_arealtime", # private function for realtime API
|
||||
"_aresponses_websocket", # private function for responses WebSocket mode
|
||||
"aimage_edit",
|
||||
"agenerate_content",
|
||||
"agenerate_content_stream",
|
||||
|
|
|
|||
|
|
@ -51,6 +51,8 @@ if TYPE_CHECKING:
|
|||
from litellm.types.llms.openai import ResponseText # type: ignore
|
||||
else:
|
||||
ResponseText = str # Fallback for ResponseText import
|
||||
from litellm.litellm_core_utils.get_litellm_params import get_litellm_params
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
from litellm.types.responses.main import *
|
||||
from litellm.types.router import GenericLiteLLMParams
|
||||
from litellm.utils import ProviderConfigManager, client
|
||||
|
|
@ -182,8 +184,6 @@ async def aresponses_api_with_mcp(
|
|||
mcp_server_auth_headers: Optional[Dict[str, Dict[str, str]]] = None
|
||||
secret_fields = kwargs.get("secret_fields")
|
||||
if secret_fields and isinstance(secret_fields, dict):
|
||||
from litellm.responses.utils import ResponsesAPIRequestUtils
|
||||
|
||||
mcp_auth_header, mcp_server_auth_headers, _, _ = (
|
||||
ResponsesAPIRequestUtils.extract_mcp_headers_from_request(
|
||||
secret_fields=secret_fields, tools=tools
|
||||
|
|
@ -1662,3 +1662,100 @@ def compact_responses(
|
|||
completion_kwargs=local_vars,
|
||||
extra_kwargs=kwargs,
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Responses API WebSocket mode
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _build_litellm_metadata_for_ws(kwargs: dict) -> dict:
|
||||
metadata: dict = {**(kwargs.get("litellm_metadata") or {})}
|
||||
guardrails = (
|
||||
(kwargs.get("metadata") or {}).get("guardrails")
|
||||
or kwargs.get("guardrails")
|
||||
or []
|
||||
)
|
||||
if guardrails:
|
||||
metadata["guardrails"] = guardrails
|
||||
return metadata
|
||||
|
||||
|
||||
@client
|
||||
async def _aresponses_websocket(
|
||||
model: str,
|
||||
websocket: Any,
|
||||
api_base: Optional[str] = None,
|
||||
api_key: Optional[str] = None,
|
||||
timeout: Optional[float] = None,
|
||||
**kwargs,
|
||||
):
|
||||
"""
|
||||
Private function to handle the Responses API WebSocket mode.
|
||||
|
||||
For PROXY use only.
|
||||
|
||||
Resolves the LLM provider from ``model``, looks up the matching
|
||||
``BaseResponsesAPIConfig``, and hands off to
|
||||
``BaseLLMHTTPHandler.async_responses_websocket``.
|
||||
"""
|
||||
litellm_logging_obj: LiteLLMLoggingObj = kwargs.get("litellm_logging_obj") # type: ignore
|
||||
user = kwargs.get("user", None)
|
||||
litellm_params = GenericLiteLLMParams(**kwargs)
|
||||
litellm_params_dict = get_litellm_params(**kwargs)
|
||||
|
||||
model, _custom_llm_provider, dynamic_api_key, dynamic_api_base = (
|
||||
litellm.get_llm_provider(
|
||||
model=model,
|
||||
api_base=api_base,
|
||||
api_key=api_key,
|
||||
)
|
||||
)
|
||||
|
||||
litellm_logging_obj.update_environment_variables(
|
||||
model=model,
|
||||
user=user,
|
||||
optional_params={},
|
||||
litellm_params=litellm_params_dict,
|
||||
custom_llm_provider=_custom_llm_provider,
|
||||
)
|
||||
|
||||
responses_api_provider_config: Optional[BaseResponsesAPIConfig] = None
|
||||
if _custom_llm_provider is not None:
|
||||
responses_api_provider_config = (
|
||||
ProviderConfigManager.get_provider_responses_api_config(
|
||||
model=model,
|
||||
provider=litellm.LlmProviders(_custom_llm_provider),
|
||||
)
|
||||
)
|
||||
|
||||
if responses_api_provider_config is None:
|
||||
raise ValueError(
|
||||
f"Responses API WebSocket mode is not supported for provider: {_custom_llm_provider}"
|
||||
)
|
||||
|
||||
resolved_api_base = (
|
||||
dynamic_api_base
|
||||
or litellm_params.api_base
|
||||
or litellm.api_base
|
||||
or None
|
||||
)
|
||||
resolved_api_key = (
|
||||
dynamic_api_key
|
||||
or litellm_params.api_key
|
||||
or litellm.api_key
|
||||
or litellm.openai_key
|
||||
or get_secret_str("OPENAI_API_KEY")
|
||||
)
|
||||
|
||||
await base_llm_http_handler.async_responses_websocket(
|
||||
model=model,
|
||||
websocket=websocket,
|
||||
logging_obj=litellm_logging_obj,
|
||||
responses_api_provider_config=responses_api_provider_config,
|
||||
api_base=resolved_api_base,
|
||||
api_key=resolved_api_key,
|
||||
timeout=timeout,
|
||||
user_api_key_dict=kwargs.get("user_api_key_dict"),
|
||||
litellm_metadata=_build_litellm_metadata_for_ws(kwargs),
|
||||
)
|
||||
|
|
|
|||
|
|
@ -3,7 +3,7 @@ import json
|
|||
import time
|
||||
import traceback
|
||||
from datetime import datetime
|
||||
from typing import Any, Dict, Optional
|
||||
from typing import TYPE_CHECKING, Any, Dict, List, Optional
|
||||
|
||||
import httpx
|
||||
|
||||
|
|
@ -682,3 +682,594 @@ class MockResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator):
|
|||
for c in getattr(out_item, "content", []):
|
||||
out += c.text
|
||||
return out
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# WebSocket mode streaming (bidirectional forwarding)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from websockets.asyncio.client import ClientConnection as _WsClientConnection
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.litellm_core_utils.thread_pool_executor import executor as _ws_executor
|
||||
|
||||
RESPONSES_WS_LOGGED_EVENT_TYPES = [
|
||||
"response.created",
|
||||
"response.completed",
|
||||
"response.failed",
|
||||
"response.incomplete",
|
||||
"error",
|
||||
]
|
||||
|
||||
|
||||
class ResponsesWebSocketStreaming:
|
||||
"""
|
||||
Manages bidirectional WebSocket forwarding for the Responses API
|
||||
WebSocket mode (wss://.../v1/responses).
|
||||
|
||||
Unlike the Realtime API, the Responses API WebSocket mode:
|
||||
- Uses response.create as the client-to-server event
|
||||
- Streams back the same events as the HTTP streaming Responses API
|
||||
- Supports previous_response_id for incremental continuation
|
||||
- Supports generate: false for warmup
|
||||
- One response at a time per connection (sequential, no multiplexing)
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
websocket: Any,
|
||||
backend_ws: Any,
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
user_api_key_dict: Optional[Any] = None,
|
||||
request_data: Optional[Dict] = None,
|
||||
):
|
||||
self.websocket = websocket
|
||||
self.backend_ws = backend_ws
|
||||
self.logging_obj = logging_obj
|
||||
self.user_api_key_dict = user_api_key_dict
|
||||
self.request_data: Dict = request_data or {}
|
||||
self.messages: list[Dict] = []
|
||||
self.input_messages: list[Dict[str, str]] = []
|
||||
|
||||
def _should_store_event(self, event_obj: dict) -> bool:
|
||||
return event_obj.get("type") in RESPONSES_WS_LOGGED_EVENT_TYPES
|
||||
|
||||
def _store_event(self, event: Any) -> None:
|
||||
if isinstance(event, bytes):
|
||||
event = event.decode("utf-8")
|
||||
if isinstance(event, str):
|
||||
try:
|
||||
event_obj = json.loads(event)
|
||||
except (json.JSONDecodeError, TypeError):
|
||||
return
|
||||
else:
|
||||
event_obj = event
|
||||
|
||||
if self._should_store_event(event_obj):
|
||||
self.messages.append(event_obj)
|
||||
|
||||
def _collect_input_from_client_event(self, message: Any) -> None:
|
||||
"""Extract user input content from response.create for logging."""
|
||||
try:
|
||||
if isinstance(message, str):
|
||||
msg_obj = json.loads(message)
|
||||
elif isinstance(message, dict):
|
||||
msg_obj = message
|
||||
else:
|
||||
return
|
||||
|
||||
if msg_obj.get("type") != "response.create":
|
||||
return
|
||||
|
||||
input_items = msg_obj.get("input", [])
|
||||
if isinstance(input_items, str):
|
||||
self.input_messages.append({"role": "user", "content": input_items})
|
||||
return
|
||||
|
||||
if isinstance(input_items, list):
|
||||
for item in input_items:
|
||||
if not isinstance(item, dict):
|
||||
continue
|
||||
if item.get("type") == "message" and item.get("role") == "user":
|
||||
content = item.get("content", [])
|
||||
if isinstance(content, str):
|
||||
self.input_messages.append(
|
||||
{"role": "user", "content": content}
|
||||
)
|
||||
elif isinstance(content, list):
|
||||
for c in content:
|
||||
if (
|
||||
isinstance(c, dict)
|
||||
and c.get("type") == "input_text"
|
||||
):
|
||||
text = c.get("text", "")
|
||||
if text:
|
||||
self.input_messages.append(
|
||||
{"role": "user", "content": text}
|
||||
)
|
||||
except (json.JSONDecodeError, AttributeError, TypeError):
|
||||
pass
|
||||
|
||||
def _store_input(self, message: Any) -> None:
|
||||
self._collect_input_from_client_event(message)
|
||||
if self.logging_obj:
|
||||
self.logging_obj.pre_call(input=message, api_key="")
|
||||
|
||||
async def _log_messages(self) -> None:
|
||||
if not self.logging_obj:
|
||||
return
|
||||
if self.input_messages:
|
||||
self.logging_obj.model_call_details["messages"] = self.input_messages
|
||||
if self.messages:
|
||||
asyncio.create_task(
|
||||
self.logging_obj.async_success_handler(self.messages)
|
||||
)
|
||||
_ws_executor.submit(self.logging_obj.success_handler, self.messages)
|
||||
|
||||
async def backend_to_client(self) -> None:
|
||||
"""Forward events from backend WebSocket to the client."""
|
||||
import websockets
|
||||
|
||||
try:
|
||||
while True:
|
||||
try:
|
||||
raw_response = await self.backend_ws.recv(decode=False) # type: ignore[union-attr]
|
||||
except TypeError:
|
||||
raw_response = await self.backend_ws.recv() # type: ignore[union-attr, assignment]
|
||||
|
||||
if isinstance(raw_response, bytes):
|
||||
response_str = raw_response.decode("utf-8")
|
||||
else:
|
||||
response_str = raw_response
|
||||
|
||||
self._store_event(response_str)
|
||||
await self.websocket.send_text(response_str)
|
||||
|
||||
except websockets.exceptions.ConnectionClosed as e: # type: ignore
|
||||
verbose_logger.debug(
|
||||
"Responses WS backend connection closed: %s", e
|
||||
)
|
||||
except Exception as e:
|
||||
verbose_logger.exception(
|
||||
"Error in responses WS backend_to_client: %s", e
|
||||
)
|
||||
finally:
|
||||
await self._log_messages()
|
||||
|
||||
async def client_to_backend(self) -> None:
|
||||
"""Forward response.create events from client to backend."""
|
||||
try:
|
||||
while True:
|
||||
message = await self.websocket.receive_text()
|
||||
|
||||
self._store_input(message)
|
||||
self._store_event(message)
|
||||
await self.backend_ws.send(message) # type: ignore[union-attr]
|
||||
|
||||
except Exception as e:
|
||||
verbose_logger.debug("Responses WS client_to_backend ended: %s", e)
|
||||
|
||||
async def bidirectional_forward(self) -> None:
|
||||
"""Run both forwarding directions concurrently."""
|
||||
forward_task = asyncio.create_task(self.backend_to_client())
|
||||
try:
|
||||
await self.client_to_backend()
|
||||
except Exception:
|
||||
pass
|
||||
finally:
|
||||
if not forward_task.done():
|
||||
forward_task.cancel()
|
||||
try:
|
||||
await forward_task
|
||||
except asyncio.CancelledError:
|
||||
pass
|
||||
try:
|
||||
await self.backend_ws.close()
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Managed WebSocket mode (HTTP-backed, provider-agnostic)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
_RESPONSE_CREATE_PARAMS = (
|
||||
"input",
|
||||
"model",
|
||||
"previous_response_id",
|
||||
"instructions",
|
||||
"max_output_tokens",
|
||||
"tools",
|
||||
"tool_choice",
|
||||
"temperature",
|
||||
"top_p",
|
||||
"store",
|
||||
"metadata",
|
||||
"truncation",
|
||||
"reasoning",
|
||||
"stream",
|
||||
"include",
|
||||
"parallel_tool_calls",
|
||||
"text",
|
||||
"user",
|
||||
"service_tier",
|
||||
"safety_identifier",
|
||||
"background",
|
||||
)
|
||||
|
||||
_MANAGED_WS_SKIP_KWARGS = frozenset(
|
||||
{
|
||||
"litellm_logging_obj",
|
||||
"litellm_call_id",
|
||||
"aresponses",
|
||||
"_aresponses_websocket",
|
||||
"user_api_key_dict",
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
class ManagedResponsesWebSocketHandler:
|
||||
"""
|
||||
Handles Responses API WebSocket mode for providers that do not expose a
|
||||
native ``wss://`` responses endpoint.
|
||||
|
||||
Instead of proxying to a provider WebSocket, this handler:
|
||||
- Listens for ``response.create`` events from the client
|
||||
- Makes HTTP streaming calls via ``litellm.aresponses(stream=True)``
|
||||
- Serialises and forwards every streaming event back over the WebSocket
|
||||
- Supports ``previous_response_id`` for multi-turn conversations via
|
||||
in-memory session tracking (avoids async DB-write timing issues)
|
||||
- Supports sequential requests over a single persistent connection
|
||||
|
||||
This makes every provider that LiteLLM can reach over HTTP available on
|
||||
the WebSocket transport without any provider-specific changes.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
websocket: Any,
|
||||
model: str,
|
||||
logging_obj: "LiteLLMLoggingObj",
|
||||
user_api_key_dict: Optional[Any] = None,
|
||||
litellm_metadata: Optional[Dict[str, Any]] = None,
|
||||
api_key: Optional[str] = None,
|
||||
api_base: Optional[str] = None,
|
||||
timeout: Optional[float] = None,
|
||||
custom_llm_provider: Optional[str] = None,
|
||||
**kwargs: Any,
|
||||
) -> None:
|
||||
self.websocket = websocket
|
||||
self.model = model
|
||||
self.logging_obj = logging_obj
|
||||
self.user_api_key_dict = user_api_key_dict
|
||||
self.litellm_metadata: Dict[str, Any] = litellm_metadata or {}
|
||||
self.api_key = api_key
|
||||
self.api_base = api_base
|
||||
self.timeout = timeout
|
||||
self.custom_llm_provider = custom_llm_provider
|
||||
# Carry through safe pass-through kwargs (e.g. extra_headers)
|
||||
self.extra_kwargs: Dict[str, Any] = {
|
||||
k: v for k, v in kwargs.items() if k not in _MANAGED_WS_SKIP_KWARGS
|
||||
}
|
||||
# In-memory session history: response_id → list of input+output messages.
|
||||
# Keyed by the DECODED (pre-encoding) response ID from response.completed.
|
||||
# This avoids the async DB-write race condition where spend logs haven't
|
||||
# been committed yet when the next response.create arrives.
|
||||
self._session_history: Dict[str, List[Dict[str, Any]]] = {}
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Internal helpers
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
@staticmethod
|
||||
def _serialize_chunk(chunk: Any) -> Optional[str]:
|
||||
"""Serialize a streaming chunk to a JSON string for WebSocket transmission."""
|
||||
try:
|
||||
if hasattr(chunk, "model_dump_json"):
|
||||
return chunk.model_dump_json(exclude_none=True)
|
||||
if hasattr(chunk, "model_dump"):
|
||||
return json.dumps(chunk.model_dump(exclude_none=True), default=str)
|
||||
if isinstance(chunk, dict):
|
||||
return json.dumps(chunk, default=str)
|
||||
return json.dumps(str(chunk))
|
||||
except Exception as exc:
|
||||
verbose_logger.debug("ManagedResponsesWS: failed to serialize chunk: %s", exc)
|
||||
return None
|
||||
|
||||
async def _send_error(self, message: str, error_type: str = "server_error") -> None:
|
||||
try:
|
||||
await self.websocket.send_text(
|
||||
json.dumps({"type": "error", "error": {"type": error_type, "message": message}})
|
||||
)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Core request handler
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
def _get_history_messages(self, previous_response_id: str) -> List[Dict[str, Any]]:
|
||||
"""
|
||||
Return accumulated message history for *previous_response_id*.
|
||||
|
||||
Checks the in-memory session store first (fast path, no DB round-trip).
|
||||
The key is the *decoded* response ID (the raw provider response ID before
|
||||
LiteLLM base64-encodes it into the ``resp_...`` format).
|
||||
"""
|
||||
from litellm.responses.utils import ResponsesAPIRequestUtils
|
||||
|
||||
decoded = ResponsesAPIRequestUtils._decode_responses_api_response_id(
|
||||
previous_response_id
|
||||
)
|
||||
raw_id = decoded.get("response_id", previous_response_id)
|
||||
return list(self._session_history.get(raw_id, []))
|
||||
|
||||
def _store_history(
|
||||
self,
|
||||
response_id: str,
|
||||
input_messages: List[Dict[str, Any]],
|
||||
output_messages: List[Dict[str, Any]],
|
||||
) -> None:
|
||||
"""
|
||||
Persist a turn's messages in the in-memory session store.
|
||||
|
||||
*response_id* is the raw (decoded) provider ID extracted from the
|
||||
``response.completed`` event so that the next turn can look it up via
|
||||
:meth:`_get_history_messages`.
|
||||
"""
|
||||
prior: List[Dict[str, Any]] = self._session_history.get(response_id, [])
|
||||
self._session_history[response_id] = prior + input_messages + output_messages
|
||||
|
||||
@staticmethod
|
||||
def _extract_response_id(completed_event: Dict[str, Any]) -> Optional[str]:
|
||||
"""
|
||||
Pull the raw (decoded) response ID out of a ``response.completed`` event.
|
||||
Returns *None* if the event doesn't contain a usable ID.
|
||||
"""
|
||||
from litellm.responses.utils import ResponsesAPIRequestUtils
|
||||
|
||||
resp_obj = completed_event.get("response", {})
|
||||
encoded_id: Optional[str] = resp_obj.get("id") if isinstance(resp_obj, dict) else None
|
||||
if not encoded_id:
|
||||
return None
|
||||
decoded = ResponsesAPIRequestUtils._decode_responses_api_response_id(encoded_id)
|
||||
return decoded.get("response_id", encoded_id)
|
||||
|
||||
@staticmethod
|
||||
def _extract_output_messages(completed_event: Dict[str, Any]) -> List[Dict[str, Any]]:
|
||||
"""
|
||||
Convert the output items in a ``response.completed`` event into
|
||||
chat-completion style messages suitable for the next turn's ``input``.
|
||||
"""
|
||||
resp_obj = completed_event.get("response", {})
|
||||
if not isinstance(resp_obj, dict):
|
||||
return []
|
||||
messages: List[Dict[str, Any]] = []
|
||||
for item in resp_obj.get("output", []) or []:
|
||||
if not isinstance(item, dict):
|
||||
continue
|
||||
item_type = item.get("type")
|
||||
role = item.get("role", "assistant")
|
||||
if item_type == "message":
|
||||
content_parts = item.get("content") or []
|
||||
text_parts = [
|
||||
p.get("text", "")
|
||||
for p in content_parts
|
||||
if isinstance(p, dict) and p.get("type") in ("output_text", "text")
|
||||
]
|
||||
text = "".join(text_parts)
|
||||
if text:
|
||||
messages.append({"type": "message", "role": role, "content": [{"type": "output_text", "text": text}]})
|
||||
elif item_type == "function_call":
|
||||
messages.append(item)
|
||||
return messages
|
||||
|
||||
@staticmethod
|
||||
def _input_to_messages(input_val: Any) -> List[Dict[str, Any]]:
|
||||
"""
|
||||
Normalise the ``input`` field of a ``response.create`` event to a list
|
||||
of Responses API message dicts.
|
||||
"""
|
||||
if isinstance(input_val, str):
|
||||
return [{"type": "message", "role": "user", "content": [{"type": "input_text", "text": input_val}]}]
|
||||
if isinstance(input_val, list):
|
||||
return [item for item in input_val if isinstance(item, dict)]
|
||||
return []
|
||||
|
||||
async def _process_response_create(self, raw_message: str) -> None:
|
||||
"""
|
||||
Parse one ``response.create`` event, call ``litellm.aresponses(stream=True)``,
|
||||
and forward every streaming event to the client.
|
||||
|
||||
Multi-turn support via in-memory session history
|
||||
------------------------------------------------
|
||||
When ``previous_response_id`` is present in the event:
|
||||
1. Look up the accumulated message history in ``self._session_history``
|
||||
(keyed by the decoded provider response ID).
|
||||
2. Prepend those messages to the current ``input`` so the model has full
|
||||
conversation context.
|
||||
3. After the stream completes, extract the new response ID and output
|
||||
messages from ``response.completed`` and store them in
|
||||
``self._session_history`` for the next turn.
|
||||
|
||||
This in-memory approach avoids the async DB-write race condition that
|
||||
occurs when spend logs haven't been committed by the time the second
|
||||
``response.create`` arrives over the same WebSocket connection.
|
||||
"""
|
||||
import litellm as _litellm
|
||||
|
||||
try:
|
||||
msg_obj = json.loads(raw_message)
|
||||
except json.JSONDecodeError:
|
||||
await self._send_error("Invalid JSON in response.create event", "invalid_request_error")
|
||||
return
|
||||
|
||||
if msg_obj.get("type") != "response.create":
|
||||
# Silently ignore non-response.create messages (e.g. warmup pings)
|
||||
return
|
||||
|
||||
# Support two wire formats:
|
||||
# Nested : {"type": "response.create", "response": {"input": [...], ...}}
|
||||
# Flat : {"type": "response.create", "input": [...], "model": "...", ...}
|
||||
nested = msg_obj.get("response")
|
||||
if isinstance(nested, dict) and nested:
|
||||
response_params: Dict[str, Any] = nested
|
||||
else:
|
||||
response_params = {k: v for k, v in msg_obj.items() if k != "type"}
|
||||
|
||||
# Build kwargs for aresponses from the response.create payload
|
||||
call_kwargs: Dict[str, Any] = {}
|
||||
for param in _RESPONSE_CREATE_PARAMS:
|
||||
if param in response_params and response_params[param] is not None:
|
||||
call_kwargs[param] = response_params[param]
|
||||
|
||||
# Always stream
|
||||
call_kwargs["stream"] = True
|
||||
|
||||
# Use the model from the event if provided, otherwise fall back to the
|
||||
# model supplied at WebSocket connect time.
|
||||
event_model = call_kwargs.pop("model", None)
|
||||
model = event_model or self.model
|
||||
|
||||
# ---- In-memory multi-turn: prepend history when previous_response_id set ----
|
||||
previous_response_id: Optional[str] = call_kwargs.pop("previous_response_id", None)
|
||||
current_input = call_kwargs.get("input")
|
||||
current_messages = self._input_to_messages(current_input)
|
||||
if previous_response_id:
|
||||
history = self._get_history_messages(previous_response_id)
|
||||
if history:
|
||||
# Prepend history; current messages are the new user turn
|
||||
call_kwargs["input"] = history + current_messages
|
||||
verbose_logger.debug(
|
||||
"ManagedResponsesWS: prepended %d history messages for previous_response_id=%s",
|
||||
len(history),
|
||||
previous_response_id,
|
||||
)
|
||||
else:
|
||||
verbose_logger.debug(
|
||||
"ManagedResponsesWS: no in-memory history for previous_response_id=%s; "
|
||||
"falling back to DB-based session reconstruction",
|
||||
previous_response_id,
|
||||
)
|
||||
# Fall back to DB-based session reconstruction (may work for
|
||||
# cross-connection multi-turn when spend logs are committed)
|
||||
call_kwargs["previous_response_id"] = previous_response_id
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
# Inject connection-level credentials and metadata.
|
||||
# Only propagate custom_llm_provider when the request is using the
|
||||
# same model as the WebSocket connection (i.e. no per-request model
|
||||
# override). If the payload specifies a different model, let litellm
|
||||
# re-resolve the provider from the model name so we don't accidentally
|
||||
# force the wrong backend.
|
||||
if self.api_key is not None:
|
||||
call_kwargs["api_key"] = self.api_key
|
||||
if self.api_base is not None:
|
||||
call_kwargs["api_base"] = self.api_base
|
||||
if self.timeout is not None:
|
||||
call_kwargs["timeout"] = self.timeout
|
||||
if self.custom_llm_provider is not None and not event_model:
|
||||
call_kwargs["custom_llm_provider"] = self.custom_llm_provider
|
||||
if self.litellm_metadata:
|
||||
call_kwargs["litellm_metadata"] = dict(self.litellm_metadata)
|
||||
|
||||
# Update proxy_server_request body so spend logs record the full request.
|
||||
proxy_server_request = (call_kwargs.get("litellm_metadata") or {}).get(
|
||||
"proxy_server_request"
|
||||
) or {}
|
||||
if isinstance(proxy_server_request, dict):
|
||||
body = dict(proxy_server_request.get("body") or {})
|
||||
body["input"] = call_kwargs.get("input")
|
||||
body["store"] = call_kwargs.get("store")
|
||||
body["model"] = model
|
||||
for k in ("tools", "tool_choice", "instructions", "metadata"):
|
||||
if k in call_kwargs and call_kwargs[k] is not None:
|
||||
body[k] = call_kwargs[k]
|
||||
proxy_server_request = dict(proxy_server_request)
|
||||
proxy_server_request["body"] = body
|
||||
if "litellm_metadata" not in call_kwargs:
|
||||
call_kwargs["litellm_metadata"] = {}
|
||||
call_kwargs["litellm_metadata"]["proxy_server_request"] = proxy_server_request
|
||||
call_kwargs.setdefault("litellm_params", {})
|
||||
call_kwargs["litellm_params"]["proxy_server_request"] = proxy_server_request
|
||||
|
||||
# Merge any safe pass-through kwargs (extra_headers, etc.)
|
||||
call_kwargs.update(self.extra_kwargs)
|
||||
|
||||
# Track the completed event to update in-memory history after the turn.
|
||||
completed_event: Optional[Dict[str, Any]] = None
|
||||
|
||||
try:
|
||||
stream_response = await _litellm.aresponses(model=model, **call_kwargs)
|
||||
|
||||
async for chunk in stream_response: # type: ignore[union-attr]
|
||||
if chunk is None:
|
||||
continue
|
||||
serialized = self._serialize_chunk(chunk)
|
||||
if serialized is not None:
|
||||
# Capture the completed event for history bookkeeping
|
||||
try:
|
||||
chunk_dict = json.loads(serialized) if isinstance(serialized, str) else {}
|
||||
if chunk_dict.get("type") == "response.completed":
|
||||
completed_event = chunk_dict
|
||||
except Exception:
|
||||
pass
|
||||
try:
|
||||
await self.websocket.send_text(serialized)
|
||||
except Exception as send_exc:
|
||||
verbose_logger.debug(
|
||||
"ManagedResponsesWS: error sending chunk to client: %s", send_exc
|
||||
)
|
||||
return # Client disconnected
|
||||
|
||||
except Exception as exc:
|
||||
verbose_logger.exception("ManagedResponsesWS: error processing response.create: %s", exc)
|
||||
await self._send_error(str(exc))
|
||||
return
|
||||
|
||||
# ---- Store this turn in in-memory history for future previous_response_id lookups ----
|
||||
if completed_event is not None:
|
||||
new_response_id = self._extract_response_id(completed_event)
|
||||
if new_response_id:
|
||||
output_msgs = self._extract_output_messages(completed_event)
|
||||
# Accumulate: history from previous turn + current input + new output
|
||||
prior_history: List[Dict[str, Any]] = []
|
||||
if previous_response_id:
|
||||
prior_history = self._get_history_messages(previous_response_id)
|
||||
self._store_history(
|
||||
new_response_id,
|
||||
prior_history + current_messages,
|
||||
output_msgs,
|
||||
)
|
||||
verbose_logger.debug(
|
||||
"ManagedResponsesWS: stored %d messages for response_id=%s",
|
||||
len(prior_history) + len(current_messages) + len(output_msgs),
|
||||
new_response_id,
|
||||
)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Main entry point
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
async def run(self) -> None:
|
||||
"""
|
||||
Main loop: accept ``response.create`` events sequentially and handle
|
||||
each one before waiting for the next message.
|
||||
"""
|
||||
try:
|
||||
while True:
|
||||
try:
|
||||
message = await self.websocket.receive_text()
|
||||
except Exception as exc:
|
||||
verbose_logger.debug(
|
||||
"ManagedResponsesWS: client disconnected: %s", exc
|
||||
)
|
||||
break
|
||||
|
||||
await self._process_response_create(message)
|
||||
|
||||
except Exception as exc:
|
||||
verbose_logger.exception("ManagedResponsesWS: unexpected error: %s", exc)
|
||||
await self._send_error(f"Internal server error: {exc}")
|
||||
|
|
|
|||
|
|
@ -885,6 +885,9 @@ class Router:
|
|||
self._arealtime = self.factory_function(
|
||||
litellm._arealtime, call_type="_arealtime"
|
||||
)
|
||||
self._aresponses_websocket = self.factory_function(
|
||||
litellm._aresponses_websocket, call_type="_aresponses_websocket"
|
||||
)
|
||||
self.acreate_fine_tuning_job = self.factory_function(
|
||||
litellm.acreate_fine_tuning_job, call_type="acreate_fine_tuning_job"
|
||||
)
|
||||
|
|
@ -1848,7 +1851,7 @@ class Router:
|
|||
finally:
|
||||
if hasattr(model_response, "close"):
|
||||
try:
|
||||
model_response.close()
|
||||
model_response.close() # type: ignore[reportAttributeAccessIssue]
|
||||
except BaseException as close_err:
|
||||
verbose_router_logger.debug(
|
||||
"stream_with_fallbacks: error closing model_response: %s",
|
||||
|
|
@ -4683,6 +4686,7 @@ class Router:
|
|||
"afile_delete",
|
||||
"afile_content",
|
||||
"_arealtime",
|
||||
"_aresponses_websocket",
|
||||
"acreate_fine_tuning_job",
|
||||
"acancel_fine_tuning_job",
|
||||
"alist_fine_tuning_jobs",
|
||||
|
|
@ -4855,6 +4859,7 @@ class Router:
|
|||
"anthropic_messages",
|
||||
"aresponses",
|
||||
"_arealtime",
|
||||
"_aresponses_websocket",
|
||||
"acreate_fine_tuning_job",
|
||||
"acancel_fine_tuning_job",
|
||||
"alist_fine_tuning_jobs",
|
||||
|
|
|
|||
|
|
@ -291,6 +291,7 @@ class CallTypes(str, Enum):
|
|||
search = "search"
|
||||
asearch = "asearch"
|
||||
arealtime = "_arealtime"
|
||||
aresponses_websocket = "_aresponses_websocket"
|
||||
create_batch = "create_batch"
|
||||
acreate_batch = "acreate_batch"
|
||||
aretrieve_batch = "aretrieve_batch"
|
||||
|
|
|
|||
|
|
@ -16,3 +16,4 @@ exclude = ["litellm/types/*", "litellm/__init__.py", "litellm/proxy/example_conf
|
|||
"litellm/proxy/utils.py" = ["F401", "PLR0915"]
|
||||
"litellm/proxy/guardrails/guardrail_hooks/litellm_content_filter/content_filter.py" = ["PLR0915"]
|
||||
"litellm/proxy/guardrails/guardrail_hooks/guardrail_benchmarks/test_eval.py" = ["PLR0915"]
|
||||
"litellm/responses/streaming_iterator.py" = ["PLR0915"]
|
||||
|
|
|
|||
|
|
@ -0,0 +1,239 @@
|
|||
"""
|
||||
E2E tests for OpenAI Responses API WebSocket mode through the LiteLLM proxy.
|
||||
|
||||
Connects to ws://0.0.0.0:4000/v1/responses, sends response.create events,
|
||||
and validates the streamed response events.
|
||||
|
||||
Requires:
|
||||
- Proxy running: python -m litellm.proxy.proxy_cli --config <config> --port 4000
|
||||
- Model configured in proxy (e.g. gpt-4o-mini)
|
||||
|
||||
See: https://developers.openai.com/api/docs/guides/websocket-mode/
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import os
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
|
||||
# ── Configuration ─────────────────────────────────────────────────────────────
|
||||
PROXY_BASE_URL = os.environ.get("LITELLM_PROXY_BASE_URL", "ws://0.0.0.0:4000")
|
||||
PROXY_MASTER_KEY = os.environ.get("LITELLM_PROXY_KEY", "sk-1234")
|
||||
PROXY_MODEL = os.environ.get("LITELLM_PROXY_RESPONSES_MODEL", "gpt-4o-mini")
|
||||
# ──────────────────────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def _generate_key() -> str:
|
||||
"""Generate a key for testing via proxy key/generate endpoint."""
|
||||
url = "http://0.0.0.0:4000/key/generate"
|
||||
headers = {
|
||||
"Authorization": f"Bearer {PROXY_MASTER_KEY}",
|
||||
"Content-Type": "application/json",
|
||||
}
|
||||
response = httpx.post(url, headers=headers, json={}, timeout=10)
|
||||
if response.status_code != 200:
|
||||
raise Exception(
|
||||
f"Key generation failed with status: {response.status_code}. "
|
||||
"Is the proxy running?"
|
||||
)
|
||||
return response.json()["key"]
|
||||
|
||||
|
||||
def _assert_basic_response(events: list[dict], label: str = "") -> None:
|
||||
"""Assert that events contain response.created, response.completed, and usage."""
|
||||
prefix = f"[{label}] " if label else ""
|
||||
types = [e.get("type") for e in events]
|
||||
assert len(events) > 0, f"{prefix}no events received"
|
||||
assert "response.created" in types, f"{prefix}missing response.created, got: {types}"
|
||||
assert "response.completed" in types, (
|
||||
f"{prefix}missing response.completed, got: {types}"
|
||||
)
|
||||
completed = next(e for e in events if e.get("type") == "response.completed")
|
||||
resp = completed.get("response", {})
|
||||
assert resp.get("status") == "completed", (
|
||||
f"{prefix}status != completed: {resp.get('status')}"
|
||||
)
|
||||
usage = resp.get("usage", {})
|
||||
assert usage.get("input_tokens", 0) > 0, f"{prefix}input_tokens=0"
|
||||
assert usage.get("output_tokens", 0) > 0, f"{prefix}output_tokens=0"
|
||||
streaming_types = {
|
||||
"response.output_item.added",
|
||||
"response.content_part.added",
|
||||
"response.output_text.delta",
|
||||
"response.output_item.done",
|
||||
}
|
||||
found = streaming_types & set(types)
|
||||
assert found, f"{prefix}no streaming delta events found, got: {types}"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_responses_websocket_proxy_basic():
|
||||
"""
|
||||
Sends a simple response.create event to the proxy WebSocket endpoint
|
||||
and validates response.created, response.completed, and streaming events.
|
||||
"""
|
||||
try:
|
||||
import websockets
|
||||
except ImportError:
|
||||
pytest.skip("websockets not installed")
|
||||
|
||||
try:
|
||||
key = _generate_key()
|
||||
except Exception as e:
|
||||
pytest.skip(
|
||||
f"Proxy not available or key generation failed: {e}. "
|
||||
"Start proxy: python -m litellm.proxy.proxy_cli --config <config> --port 4000"
|
||||
)
|
||||
|
||||
url = f"{PROXY_BASE_URL}/v1/responses?model={PROXY_MODEL}"
|
||||
headers = {"Authorization": f"Bearer {key}"}
|
||||
events: list[dict] = []
|
||||
|
||||
try:
|
||||
async with websockets.connect(
|
||||
url, additional_headers=headers, open_timeout=5
|
||||
) as ws:
|
||||
payload = {
|
||||
"type": "response.create",
|
||||
"model": PROXY_MODEL,
|
||||
"store": False,
|
||||
"input": [
|
||||
{
|
||||
"type": "message",
|
||||
"role": "user",
|
||||
"content": [
|
||||
{"type": "input_text", "text": "Say hello in one word."}
|
||||
],
|
||||
}
|
||||
],
|
||||
"tools": [],
|
||||
}
|
||||
await ws.send(json.dumps(payload))
|
||||
for _ in range(50):
|
||||
msg = await asyncio.wait_for(ws.recv(), timeout=15)
|
||||
event = json.loads(msg)
|
||||
events.append(event)
|
||||
if event.get("type") in (
|
||||
"response.completed",
|
||||
"response.failed",
|
||||
"error",
|
||||
):
|
||||
break
|
||||
except Exception as e:
|
||||
pytest.fail(
|
||||
f"WebSocket connection failed: {e}. "
|
||||
"Ensure proxy is running and model is configured."
|
||||
)
|
||||
|
||||
_assert_basic_response(events, "proxy-basic")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_responses_websocket_proxy_multi_turn():
|
||||
"""
|
||||
Sends two sequential response.create events with previous_response_id
|
||||
to validate multi-turn conversation over a single WebSocket.
|
||||
"""
|
||||
try:
|
||||
import websockets
|
||||
except ImportError:
|
||||
pytest.skip("websockets not installed")
|
||||
|
||||
try:
|
||||
key = _generate_key()
|
||||
except Exception as e:
|
||||
pytest.skip(
|
||||
f"Proxy not available or key generation failed: {e}. "
|
||||
"Start proxy: python -m litellm.proxy.proxy_cli --config <config> --port 4000"
|
||||
)
|
||||
|
||||
url = f"{PROXY_BASE_URL}/v1/responses?model={PROXY_MODEL}"
|
||||
headers = {"Authorization": f"Bearer {key}"}
|
||||
all_events: list[dict] = []
|
||||
completed: list[dict] = []
|
||||
first_id = None
|
||||
|
||||
try:
|
||||
async with websockets.connect(
|
||||
url, additional_headers=headers, open_timeout=5
|
||||
) as ws:
|
||||
# Turn 1
|
||||
await ws.send(
|
||||
json.dumps(
|
||||
{
|
||||
"type": "response.create",
|
||||
"model": PROXY_MODEL,
|
||||
"store": True,
|
||||
"input": [
|
||||
{
|
||||
"type": "message",
|
||||
"role": "user",
|
||||
"content": [
|
||||
{
|
||||
"type": "input_text",
|
||||
"text": "Remember the number 7. Just say OK.",
|
||||
}
|
||||
],
|
||||
}
|
||||
],
|
||||
}
|
||||
)
|
||||
)
|
||||
for _ in range(50):
|
||||
msg = await asyncio.wait_for(ws.recv(), timeout=15)
|
||||
event = json.loads(msg)
|
||||
all_events.append(event)
|
||||
if event.get("type") == "response.completed":
|
||||
completed.append(event)
|
||||
first_id = event.get("response", {}).get("id")
|
||||
break
|
||||
if event.get("type") in ("response.failed", "error"):
|
||||
break
|
||||
|
||||
assert first_id, "Turn 1 never completed"
|
||||
|
||||
# Turn 2
|
||||
await ws.send(
|
||||
json.dumps(
|
||||
{
|
||||
"type": "response.create",
|
||||
"model": PROXY_MODEL,
|
||||
"store": True,
|
||||
"previous_response_id": first_id,
|
||||
"input": [
|
||||
{
|
||||
"type": "message",
|
||||
"role": "user",
|
||||
"content": [
|
||||
{
|
||||
"type": "input_text",
|
||||
"text": "What number did I tell you to remember?",
|
||||
}
|
||||
],
|
||||
}
|
||||
],
|
||||
}
|
||||
)
|
||||
)
|
||||
for _ in range(50):
|
||||
msg = await asyncio.wait_for(ws.recv(), timeout=15)
|
||||
event = json.loads(msg)
|
||||
all_events.append(event)
|
||||
if event.get("type") == "response.completed":
|
||||
completed.append(event)
|
||||
break
|
||||
if event.get("type") in ("response.failed", "error"):
|
||||
break
|
||||
|
||||
except Exception as e:
|
||||
pytest.fail(
|
||||
f"WebSocket multi-turn failed: {e}. "
|
||||
"Ensure proxy is running and model is configured."
|
||||
)
|
||||
|
||||
assert len(completed) >= 2, (
|
||||
f"Expected 2 response.completed events, got {len(completed)}"
|
||||
)
|
||||
assert completed[1].get("response", {}).get("status") == "completed"
|
||||
Loading…
Add table
Reference in a new issue