mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
refactor: move websocket code into litellm/responses/ to match existing organization
- litellm/responses/websocket.py: entry point reusing ProviderConfigManager and BaseResponsesAPIConfig.validate_environment/get_complete_url for credential resolution — same pattern as responses/main.py - litellm/responses/websocket_handler.py: handler receives resolved URL + auth headers, only owns WS-specific concerns (protocol upgrade, SSL, bidirectional forwarding) - litellm/responses/websocket_streaming.py: moved from litellm_core_utils/ Removed: - litellm/realtime_api/responses_websocket.py (duplicated provider logic) - litellm/llms/openai/responses/websocket_handler.py (duplicated credential resolution) Co-authored-by: Ishaan Jaff <ishaan-jaff@users.noreply.github.com>
This commit is contained in:
parent
1e394db9e9
commit
4738911e8f
9 changed files with 250 additions and 311 deletions
|
|
@ -1238,7 +1238,7 @@ from .ocr.main import *
|
|||
from .rag.main import *
|
||||
from .search.main import *
|
||||
from .realtime_api.main import _arealtime
|
||||
from .realtime_api.responses_websocket import _aresponses_websocket
|
||||
from .responses.websocket import _aresponses_websocket
|
||||
from .fine_tuning.main import *
|
||||
from .files.main import *
|
||||
from .vector_store_files.main import (
|
||||
|
|
|
|||
|
|
@ -1,126 +0,0 @@
|
|||
"""
|
||||
OpenAI Responses API WebSocket Mode handler.
|
||||
|
||||
Implements the WebSocket transport for OpenAI's Responses API
|
||||
(wss://api.openai.com/v1/responses).
|
||||
|
||||
Protocol summary (from https://developers.openai.com/api/docs/guides/websocket-mode/):
|
||||
- Client connects via WebSocket to /v1/responses
|
||||
- Client sends `response.create` events; payload mirrors the HTTP Responses
|
||||
create body but omits transport-specific fields (`stream`, `background`).
|
||||
- Server streams back the same SSE event types used by the HTTP streaming
|
||||
endpoint, wrapped in JSON-framed WebSocket messages.
|
||||
- Client may continue a conversation by sending another `response.create`
|
||||
with `previous_response_id` and incremental input.
|
||||
- A warmup request (`generate: false`) can be sent to pre-populate
|
||||
connection state without triggering generation.
|
||||
"""
|
||||
|
||||
from typing import Any, Optional
|
||||
|
||||
from litellm.constants import REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging
|
||||
from litellm.litellm_core_utils.responses_websocket_streaming import (
|
||||
ResponsesWebSocketStreaming,
|
||||
)
|
||||
from litellm.llms.custom_httpx.http_handler import get_shared_realtime_ssl_context
|
||||
|
||||
|
||||
class OpenAIResponsesWebSocket:
|
||||
"""
|
||||
Handler for OpenAI Responses API WebSocket connections.
|
||||
|
||||
Mirrors the structure of ``OpenAIRealtime`` but targets the
|
||||
``/v1/responses`` WebSocket endpoint instead of ``/v1/realtime``.
|
||||
"""
|
||||
|
||||
def _get_default_api_base(self) -> str:
|
||||
return "https://api.openai.com/v1"
|
||||
|
||||
def _get_headers(self, api_key: str) -> dict:
|
||||
return {
|
||||
"Authorization": f"Bearer {api_key}",
|
||||
}
|
||||
|
||||
def _get_ssl_config(self, url: str) -> Any:
|
||||
if url.startswith("ws://"):
|
||||
return None
|
||||
ssl_config = get_shared_realtime_ssl_context()
|
||||
if ssl_config is False:
|
||||
return True
|
||||
return ssl_config
|
||||
|
||||
def _construct_url(self, api_base: str) -> str:
|
||||
from httpx import URL
|
||||
|
||||
api_base = api_base.replace("https://", "wss://").replace(
|
||||
"http://", "ws://"
|
||||
)
|
||||
url = URL(api_base)
|
||||
if not url.raw_path.endswith(b"/responses"):
|
||||
url = url.copy_with(path="/v1/responses")
|
||||
return str(url)
|
||||
|
||||
async def async_responses_websocket(
|
||||
self,
|
||||
websocket: Any,
|
||||
logging_obj: LiteLLMLogging,
|
||||
api_base: Optional[str] = None,
|
||||
api_key: Optional[str] = None,
|
||||
timeout: Optional[float] = None,
|
||||
user_api_key_dict: Optional[Any] = None,
|
||||
**kwargs: Any,
|
||||
) -> None:
|
||||
import websockets
|
||||
|
||||
if api_base is None:
|
||||
api_base = self._get_default_api_base()
|
||||
if api_key is None:
|
||||
raise ValueError("api_key is required for OpenAI Responses WebSocket calls")
|
||||
|
||||
url = self._construct_url(api_base)
|
||||
headers = self._get_headers(api_key)
|
||||
ssl_config = self._get_ssl_config(url)
|
||||
|
||||
logging_obj.pre_call(
|
||||
input=None,
|
||||
api_key=api_key,
|
||||
additional_args={
|
||||
"api_base": url,
|
||||
"headers": headers,
|
||||
"complete_input_dict": {},
|
||||
},
|
||||
)
|
||||
|
||||
try:
|
||||
async with websockets.connect( # type: ignore
|
||||
url,
|
||||
additional_headers=headers,
|
||||
max_size=REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES,
|
||||
ssl=ssl_config,
|
||||
) as backend_ws:
|
||||
streaming = ResponsesWebSocketStreaming(
|
||||
websocket=websocket,
|
||||
backend_ws=backend_ws,
|
||||
logging_obj=logging_obj,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
)
|
||||
await streaming.bidirectional_forward()
|
||||
|
||||
except websockets.exceptions.InvalidStatus as e: # type: ignore
|
||||
status = getattr(
|
||||
getattr(e, "response", None), "status_code", 1011
|
||||
)
|
||||
await websocket.close(code=status, reason=str(e))
|
||||
except Exception as e:
|
||||
try:
|
||||
await websocket.close(
|
||||
code=1011, reason=f"Internal server error: {str(e)}"
|
||||
)
|
||||
except RuntimeError as close_error:
|
||||
if "already completed" not in str(
|
||||
close_error
|
||||
) and "websocket.close" not in str(close_error):
|
||||
raise Exception(
|
||||
f"Unexpected error while closing WebSocket: {close_error}"
|
||||
)
|
||||
|
|
@ -1,89 +0,0 @@
|
|||
"""
|
||||
Entry point for the Responses API WebSocket mode.
|
||||
|
||||
Analogous to ``litellm.realtime_api.main._arealtime`` but for the
|
||||
``/v1/responses`` WebSocket transport.
|
||||
|
||||
Currently supports OpenAI only. Other providers can be added as they ship
|
||||
their own WebSocket modes.
|
||||
"""
|
||||
|
||||
from typing import Any, Optional, cast
|
||||
|
||||
import litellm
|
||||
from litellm.litellm_core_utils.get_litellm_params import get_litellm_params
|
||||
from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging
|
||||
from litellm.llms.openai.responses.websocket_handler import OpenAIResponsesWebSocket
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
from litellm.types.router import GenericLiteLLMParams
|
||||
from litellm.utils import client as wrapper_client
|
||||
|
||||
openai_responses_ws = OpenAIResponsesWebSocket()
|
||||
|
||||
|
||||
@wrapper_client
|
||||
async def _aresponses_websocket(
|
||||
model: str,
|
||||
websocket: Any,
|
||||
api_base: Optional[str] = None,
|
||||
api_key: Optional[str] = None,
|
||||
**kwargs: Any,
|
||||
) -> None:
|
||||
"""
|
||||
Private function for the Responses API WebSocket transport.
|
||||
|
||||
For PROXY use only.
|
||||
"""
|
||||
headers = cast(Optional[dict], kwargs.get("headers"))
|
||||
extra_headers = cast(Optional[dict], kwargs.get("extra_headers"))
|
||||
if headers is None:
|
||||
headers = {}
|
||||
if extra_headers is not None:
|
||||
headers.update(extra_headers)
|
||||
|
||||
litellm_logging_obj: LiteLLMLogging = 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 = 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,
|
||||
)
|
||||
|
||||
if _custom_llm_provider == "openai":
|
||||
resolved_api_base = (
|
||||
dynamic_api_base
|
||||
or litellm_params.api_base
|
||||
or litellm.api_base
|
||||
or "https://api.openai.com/v1"
|
||||
)
|
||||
resolved_api_key = (
|
||||
dynamic_api_key
|
||||
or litellm.api_key
|
||||
or litellm.openai_key
|
||||
or get_secret_str("OPENAI_API_KEY")
|
||||
)
|
||||
|
||||
await openai_responses_ws.async_responses_websocket(
|
||||
websocket=websocket,
|
||||
logging_obj=litellm_logging_obj,
|
||||
api_base=resolved_api_base,
|
||||
api_key=resolved_api_key,
|
||||
user_api_key_dict=kwargs.get("user_api_key_dict"),
|
||||
)
|
||||
else:
|
||||
raise ValueError(
|
||||
f"Responses API WebSocket mode is not supported for provider: {_custom_llm_provider}. "
|
||||
"Currently only 'openai' is supported."
|
||||
)
|
||||
104
litellm/responses/websocket.py
Normal file
104
litellm/responses/websocket.py
Normal file
|
|
@ -0,0 +1,104 @@
|
|||
"""
|
||||
Entry point for the Responses API WebSocket mode.
|
||||
|
||||
Lives alongside ``litellm.responses.main`` and reuses the same
|
||||
``ProviderConfigManager`` / ``get_llm_provider`` patterns for
|
||||
credential resolution.
|
||||
|
||||
Currently supports OpenAI. Other providers can be added by
|
||||
implementing ``get_websocket_url`` / ``get_websocket_headers``
|
||||
on their ``BaseResponsesAPIConfig`` subclass.
|
||||
"""
|
||||
|
||||
from typing import Any, Optional, cast
|
||||
|
||||
from litellm.litellm_core_utils.get_litellm_params import get_litellm_params
|
||||
from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging
|
||||
from litellm.llms.base_llm.responses.transformation import BaseResponsesAPIConfig
|
||||
from litellm.types.router import GenericLiteLLMParams
|
||||
from litellm.types.utils import LlmProviders
|
||||
from litellm.utils import ProviderConfigManager, client as wrapper_client
|
||||
|
||||
from .websocket_handler import OpenAIResponsesWebSocketHandler
|
||||
|
||||
openai_responses_ws_handler = OpenAIResponsesWebSocketHandler()
|
||||
|
||||
|
||||
@wrapper_client
|
||||
async def _aresponses_websocket(
|
||||
model: str,
|
||||
websocket: Any,
|
||||
api_base: Optional[str] = None,
|
||||
api_key: Optional[str] = None,
|
||||
**kwargs: Any,
|
||||
) -> None:
|
||||
"""
|
||||
Responses API WebSocket transport. For proxy use only.
|
||||
|
||||
Follows the same provider-resolution flow as ``litellm.responses.main.responses``:
|
||||
1. ``get_llm_provider`` → model, provider, dynamic key/base
|
||||
2. ``ProviderConfigManager.get_provider_responses_api_config`` → config
|
||||
3. ``config.validate_environment`` → auth headers
|
||||
4. ``config.get_complete_url`` → HTTP URL (converted to WSS)
|
||||
"""
|
||||
headers = cast(Optional[dict], kwargs.get("headers"))
|
||||
extra_headers = cast(Optional[dict], kwargs.get("extra_headers"))
|
||||
if headers is None:
|
||||
headers = {}
|
||||
if extra_headers is not None:
|
||||
headers.update(extra_headers)
|
||||
|
||||
litellm_logging_obj: LiteLLMLogging = 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 = get_llm_provider(
|
||||
model=model,
|
||||
api_base=api_base,
|
||||
api_key=api_key,
|
||||
)
|
||||
|
||||
if dynamic_api_key is not None:
|
||||
litellm_params.api_key = dynamic_api_key
|
||||
if dynamic_api_base is not None:
|
||||
litellm_params.api_base = dynamic_api_base
|
||||
|
||||
litellm_logging_obj.update_environment_variables(
|
||||
model=model,
|
||||
user=user,
|
||||
optional_params={},
|
||||
litellm_params=litellm_params_dict,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
)
|
||||
|
||||
responses_config: Optional[BaseResponsesAPIConfig] = (
|
||||
ProviderConfigManager.get_provider_responses_api_config(
|
||||
model=model,
|
||||
provider=LlmProviders(custom_llm_provider),
|
||||
)
|
||||
)
|
||||
|
||||
if responses_config is None:
|
||||
raise ValueError(
|
||||
f"Responses API WebSocket mode is not supported for provider: "
|
||||
f"{custom_llm_provider}. No responses config found."
|
||||
)
|
||||
|
||||
auth_headers = responses_config.validate_environment(
|
||||
headers={}, model=model, litellm_params=litellm_params
|
||||
)
|
||||
|
||||
http_url = responses_config.get_complete_url(
|
||||
api_base=litellm_params.api_base,
|
||||
litellm_params=litellm_params_dict,
|
||||
)
|
||||
|
||||
await openai_responses_ws_handler.async_responses_websocket(
|
||||
websocket=websocket,
|
||||
logging_obj=litellm_logging_obj,
|
||||
http_url=http_url,
|
||||
auth_headers=auth_headers,
|
||||
user_api_key_dict=kwargs.get("user_api_key_dict"),
|
||||
)
|
||||
95
litellm/responses/websocket_handler.py
Normal file
95
litellm/responses/websocket_handler.py
Normal file
|
|
@ -0,0 +1,95 @@
|
|||
"""
|
||||
WebSocket handler for the OpenAI Responses API WebSocket mode.
|
||||
|
||||
Receives a fully-resolved HTTP URL and auth headers from the entry
|
||||
point in ``litellm.responses.websocket`` (which uses the same
|
||||
``BaseResponsesAPIConfig`` credential-resolution as the HTTP path).
|
||||
|
||||
This handler only owns the WebSocket-specific concerns:
|
||||
- Converting the HTTP URL to a WSS URL
|
||||
- Opening the ``websockets`` connection
|
||||
- Bidirectional message forwarding via ``ResponsesWebSocketStreaming``
|
||||
"""
|
||||
|
||||
from typing import Any, Optional
|
||||
|
||||
from litellm.constants import REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging
|
||||
from litellm.llms.custom_httpx.http_handler import get_shared_realtime_ssl_context
|
||||
from litellm.responses.websocket_streaming import ResponsesWebSocketStreaming
|
||||
|
||||
|
||||
class OpenAIResponsesWebSocketHandler:
|
||||
"""
|
||||
Handles the WebSocket connection lifecycle for the Responses API.
|
||||
|
||||
Analogous to ``OpenAIRealtime`` but for ``/v1/responses``.
|
||||
"""
|
||||
|
||||
@staticmethod
|
||||
def _http_url_to_ws(http_url: str) -> str:
|
||||
return (
|
||||
http_url
|
||||
.replace("https://", "wss://")
|
||||
.replace("http://", "ws://")
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _get_ssl_config(url: str) -> Any:
|
||||
if url.startswith("ws://"):
|
||||
return None
|
||||
ssl_config = get_shared_realtime_ssl_context()
|
||||
if ssl_config is False:
|
||||
return True
|
||||
return ssl_config
|
||||
|
||||
async def async_responses_websocket(
|
||||
self,
|
||||
websocket: Any,
|
||||
logging_obj: LiteLLMLogging,
|
||||
http_url: str,
|
||||
auth_headers: dict,
|
||||
user_api_key_dict: Optional[Any] = None,
|
||||
) -> None:
|
||||
import websockets
|
||||
|
||||
ws_url = self._http_url_to_ws(http_url)
|
||||
ssl_config = self._get_ssl_config(ws_url)
|
||||
|
||||
logging_obj.pre_call(
|
||||
input=None,
|
||||
api_key="",
|
||||
additional_args={
|
||||
"api_base": ws_url,
|
||||
"headers": auth_headers,
|
||||
"complete_input_dict": {},
|
||||
},
|
||||
)
|
||||
|
||||
try:
|
||||
async with websockets.connect( # type: ignore
|
||||
ws_url,
|
||||
additional_headers=auth_headers,
|
||||
max_size=REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES,
|
||||
ssl=ssl_config,
|
||||
) as backend_ws:
|
||||
streaming = ResponsesWebSocketStreaming(
|
||||
websocket=websocket,
|
||||
backend_ws=backend_ws,
|
||||
logging_obj=logging_obj,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
)
|
||||
await streaming.bidirectional_forward()
|
||||
|
||||
except Exception as e:
|
||||
try:
|
||||
await websocket.close(
|
||||
code=1011, reason=f"Internal server error: {str(e)}"
|
||||
)
|
||||
except RuntimeError as close_error:
|
||||
if "already completed" not in str(
|
||||
close_error
|
||||
) and "websocket.close" not in str(close_error):
|
||||
raise Exception(
|
||||
f"Unexpected error while closing WebSocket: {close_error}"
|
||||
)
|
||||
|
|
@ -21,7 +21,7 @@ from typing import TYPE_CHECKING, Any, Dict, List, Optional
|
|||
|
||||
from litellm._logging import verbose_logger
|
||||
|
||||
from .litellm_logging import Logging as LiteLLMLogging
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from websockets.asyncio.client import ClientConnection
|
||||
|
|
@ -16,85 +16,48 @@ import pytest
|
|||
|
||||
|
||||
class TestOpenAIResponsesWebSocketHandler:
|
||||
"""Unit tests for OpenAIResponsesWebSocket handler."""
|
||||
"""Unit tests for OpenAIResponsesWebSocketHandler."""
|
||||
|
||||
def test_construct_url_from_https_base(self):
|
||||
from litellm.llms.openai.responses.websocket_handler import (
|
||||
OpenAIResponsesWebSocket,
|
||||
def test_http_url_to_ws_https(self):
|
||||
from litellm.responses.websocket_handler import (
|
||||
OpenAIResponsesWebSocketHandler,
|
||||
)
|
||||
|
||||
handler = OpenAIResponsesWebSocket()
|
||||
url = handler._construct_url("https://api.openai.com/v1")
|
||||
assert url.startswith("wss://")
|
||||
assert url.endswith("/v1/responses")
|
||||
assert OpenAIResponsesWebSocketHandler._http_url_to_ws(
|
||||
"https://api.openai.com/v1/responses"
|
||||
) == "wss://api.openai.com/v1/responses"
|
||||
|
||||
def test_construct_url_from_http_base(self):
|
||||
from litellm.llms.openai.responses.websocket_handler import (
|
||||
OpenAIResponsesWebSocket,
|
||||
def test_http_url_to_ws_http(self):
|
||||
from litellm.responses.websocket_handler import (
|
||||
OpenAIResponsesWebSocketHandler,
|
||||
)
|
||||
|
||||
handler = OpenAIResponsesWebSocket()
|
||||
url = handler._construct_url("http://localhost:8080/v1")
|
||||
assert url.startswith("ws://")
|
||||
assert "/v1/responses" in url
|
||||
result = OpenAIResponsesWebSocketHandler._http_url_to_ws(
|
||||
"http://localhost:4000/v1/responses"
|
||||
)
|
||||
assert result == "ws://localhost:4000/v1/responses"
|
||||
|
||||
def test_construct_url_already_has_responses_path(self):
|
||||
from litellm.llms.openai.responses.websocket_handler import (
|
||||
OpenAIResponsesWebSocket,
|
||||
def test_ssl_config_ws(self):
|
||||
from litellm.responses.websocket_handler import (
|
||||
OpenAIResponsesWebSocketHandler,
|
||||
)
|
||||
|
||||
handler = OpenAIResponsesWebSocket()
|
||||
url = handler._construct_url("https://api.openai.com/v1/responses")
|
||||
assert url == "wss://api.openai.com/v1/responses"
|
||||
assert OpenAIResponsesWebSocketHandler._get_ssl_config("ws://localhost:8080") is None
|
||||
|
||||
def test_get_headers(self):
|
||||
from litellm.llms.openai.responses.websocket_handler import (
|
||||
OpenAIResponsesWebSocket,
|
||||
def test_ssl_config_wss(self):
|
||||
from litellm.responses.websocket_handler import (
|
||||
OpenAIResponsesWebSocketHandler,
|
||||
)
|
||||
|
||||
handler = OpenAIResponsesWebSocket()
|
||||
headers = handler._get_headers("sk-test-key")
|
||||
assert headers == {"Authorization": "Bearer sk-test-key"}
|
||||
|
||||
def test_get_default_api_base(self):
|
||||
from litellm.llms.openai.responses.websocket_handler import (
|
||||
OpenAIResponsesWebSocket,
|
||||
)
|
||||
|
||||
handler = OpenAIResponsesWebSocket()
|
||||
assert handler._get_default_api_base() == "https://api.openai.com/v1"
|
||||
|
||||
def test_get_ssl_config_ws(self):
|
||||
from litellm.llms.openai.responses.websocket_handler import (
|
||||
OpenAIResponsesWebSocket,
|
||||
)
|
||||
|
||||
handler = OpenAIResponsesWebSocket()
|
||||
assert handler._get_ssl_config("ws://localhost:8080") is None
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_missing_api_key_raises(self):
|
||||
from litellm.llms.openai.responses.websocket_handler import (
|
||||
OpenAIResponsesWebSocket,
|
||||
)
|
||||
|
||||
handler = OpenAIResponsesWebSocket()
|
||||
mock_ws = MagicMock()
|
||||
mock_logging = MagicMock()
|
||||
|
||||
with pytest.raises(ValueError, match="api_key is required"):
|
||||
await handler.async_responses_websocket(
|
||||
websocket=mock_ws,
|
||||
logging_obj=mock_logging,
|
||||
api_key=None,
|
||||
)
|
||||
result = OpenAIResponsesWebSocketHandler._get_ssl_config("wss://api.openai.com/v1/responses")
|
||||
assert result is not None
|
||||
|
||||
|
||||
class TestResponsesWebSocketStreaming:
|
||||
"""Unit tests for ResponsesWebSocketStreaming."""
|
||||
|
||||
def test_store_backend_message_logged_event(self):
|
||||
from litellm.litellm_core_utils.responses_websocket_streaming import (
|
||||
from litellm.responses.websocket_streaming import (
|
||||
ResponsesWebSocketStreaming,
|
||||
)
|
||||
|
||||
|
|
@ -114,7 +77,7 @@ class TestResponsesWebSocketStreaming:
|
|||
assert streaming.messages[0]["type"] == "response.created"
|
||||
|
||||
def test_store_backend_message_unlogged_event(self):
|
||||
from litellm.litellm_core_utils.responses_websocket_streaming import (
|
||||
from litellm.responses.websocket_streaming import (
|
||||
ResponsesWebSocketStreaming,
|
||||
)
|
||||
|
||||
|
|
@ -133,7 +96,7 @@ class TestResponsesWebSocketStreaming:
|
|||
assert len(streaming.messages) == 0
|
||||
|
||||
def test_store_client_message(self):
|
||||
from litellm.litellm_core_utils.responses_websocket_streaming import (
|
||||
from litellm.responses.websocket_streaming import (
|
||||
ResponsesWebSocketStreaming,
|
||||
)
|
||||
|
||||
|
|
@ -160,7 +123,7 @@ class TestResponsesWebSocketStreaming:
|
|||
@pytest.mark.asyncio
|
||||
async def test_bidirectional_forward_client_disconnect(self):
|
||||
"""When the client disconnects, the forward task should be cancelled."""
|
||||
from litellm.litellm_core_utils.responses_websocket_streaming import (
|
||||
from litellm.responses.websocket_streaming import (
|
||||
ResponsesWebSocketStreaming,
|
||||
)
|
||||
|
||||
|
|
@ -184,7 +147,7 @@ class TestResponsesWebSocketStreaming:
|
|||
|
||||
@pytest.mark.asyncio
|
||||
async def test_client_to_backend_forwards_message(self):
|
||||
from litellm.litellm_core_utils.responses_websocket_streaming import (
|
||||
from litellm.responses.websocket_streaming import (
|
||||
ResponsesWebSocketStreaming,
|
||||
)
|
||||
|
||||
|
|
@ -211,7 +174,7 @@ class TestResponsesWebSocketStreaming:
|
|||
async def test_backend_to_client_forwards_message(self):
|
||||
from websockets.exceptions import ConnectionClosed
|
||||
|
||||
from litellm.litellm_core_utils.responses_websocket_streaming import (
|
||||
from litellm.responses.websocket_streaming import (
|
||||
ResponsesWebSocketStreaming,
|
||||
)
|
||||
|
||||
|
|
@ -255,7 +218,7 @@ class TestResponsesWebSocketEntryPoint:
|
|||
raises ValueError. We patch the inner function directly to avoid
|
||||
the @wrapper_client decorator complexity.
|
||||
"""
|
||||
from litellm.realtime_api.responses_websocket import (
|
||||
from litellm.responses.websocket import (
|
||||
_aresponses_websocket,
|
||||
)
|
||||
|
||||
|
|
@ -278,11 +241,11 @@ class TestResponsesWebSocketEntryPoint:
|
|||
Directly test the inner logic that the ValueError message is correct
|
||||
for unsupported providers.
|
||||
"""
|
||||
from litellm.realtime_api.responses_websocket import (
|
||||
openai_responses_ws,
|
||||
from litellm.responses.websocket import (
|
||||
openai_responses_ws_handler,
|
||||
)
|
||||
|
||||
assert openai_responses_ws is not None
|
||||
assert openai_responses_ws_handler is not None
|
||||
|
||||
|
||||
class TestProxyWebSocketEndpointRegistration:
|
||||
|
|
|
|||
|
|
@ -117,7 +117,7 @@ class TestResponsesWebSocketHandlerE2E:
|
|||
2. Backend sends back response.created, output_text.delta, response.completed
|
||||
3. Verify all messages are forwarded correctly
|
||||
"""
|
||||
from litellm.litellm_core_utils.responses_websocket_streaming import (
|
||||
from litellm.responses.websocket_streaming import (
|
||||
ResponsesWebSocketStreaming,
|
||||
)
|
||||
|
||||
|
|
@ -187,7 +187,7 @@ class TestResponsesWebSocketHandlerE2E:
|
|||
Verify client→backend forwarding: client sends response.create
|
||||
and the backend WS receives it.
|
||||
"""
|
||||
from litellm.litellm_core_utils.responses_websocket_streaming import (
|
||||
from litellm.responses.websocket_streaming import (
|
||||
ResponsesWebSocketStreaming,
|
||||
)
|
||||
|
||||
|
|
@ -231,12 +231,10 @@ class TestResponsesWebSocketHandlerE2E:
|
|||
@pytest.mark.asyncio
|
||||
async def test_handler_constructs_correct_wss_url(self):
|
||||
"""Verify OpenAIResponsesWebSocket builds the correct WSS URL."""
|
||||
from litellm.llms.openai.responses.websocket_handler import (
|
||||
OpenAIResponsesWebSocket,
|
||||
from litellm.responses.websocket_handler import (
|
||||
OpenAIResponsesWebSocketHandler,
|
||||
)
|
||||
|
||||
handler = OpenAIResponsesWebSocket()
|
||||
|
||||
assert handler._construct_url("https://api.openai.com/v1") == "wss://api.openai.com/v1/responses"
|
||||
assert handler._construct_url("http://localhost:4000/v1") == "ws://localhost:4000/v1/responses"
|
||||
assert handler._construct_url("https://custom.endpoint.com/v1/responses") == "wss://custom.endpoint.com/v1/responses"
|
||||
assert OpenAIResponsesWebSocketHandler._http_url_to_ws("https://api.openai.com/v1/responses") == "wss://api.openai.com/v1/responses"
|
||||
assert OpenAIResponsesWebSocketHandler._http_url_to_ws("http://localhost:4000/v1/responses") == "ws://localhost:4000/v1/responses"
|
||||
assert OpenAIResponsesWebSocketHandler._http_url_to_ws("https://custom.endpoint.com/v1/responses") == "wss://custom.endpoint.com/v1/responses"
|
||||
|
|
|
|||
|
|
@ -42,10 +42,8 @@ async def test_responses_websocket_live_through_proxy():
|
|||
async with websockets.connect(url, additional_headers=headers) as ws:
|
||||
request_event = {
|
||||
"type": "response.create",
|
||||
"response": {
|
||||
"model": MODEL,
|
||||
"input": "Say exactly: Hello WebSocket",
|
||||
},
|
||||
"model": MODEL,
|
||||
"input": "Say exactly: Hello WebSocket",
|
||||
}
|
||||
await ws.send(json.dumps(request_event))
|
||||
|
||||
|
|
@ -101,10 +99,8 @@ async def test_responses_websocket_continuation():
|
|||
# First turn
|
||||
await ws.send(json.dumps({
|
||||
"type": "response.create",
|
||||
"response": {
|
||||
"model": MODEL,
|
||||
"input": "Remember the number 42.",
|
||||
},
|
||||
"model": MODEL,
|
||||
"input": "Remember the number 42.",
|
||||
}))
|
||||
|
||||
first_response_id = None
|
||||
|
|
@ -122,13 +118,11 @@ async def test_responses_websocket_continuation():
|
|||
# Second turn — continuation
|
||||
await ws.send(json.dumps({
|
||||
"type": "response.create",
|
||||
"response": {
|
||||
"model": MODEL,
|
||||
"input": [
|
||||
{"type": "message", "role": "user", "content": "What number did I mention?"},
|
||||
],
|
||||
"previous_response_id": first_response_id,
|
||||
},
|
||||
"model": MODEL,
|
||||
"input": [
|
||||
{"type": "message", "role": "user", "content": "What number did I mention?"},
|
||||
],
|
||||
"previous_response_id": first_response_id,
|
||||
}))
|
||||
|
||||
second_events = []
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue