From eec6f6ee697a386d76890d03f665f5e2b68688c3 Mon Sep 17 00:00:00 2001 From: Sameer Kankute Date: Wed, 4 Mar 2026 18:24:50 +0530 Subject: [PATCH] Add support for responses websocket for all providers --- litellm/llms/custom_httpx/llm_http_handler.py | 35 +- litellm/responses/main.py | 10 +- litellm/responses/streaming_iterator.py | 384 +++++++++--------- 3 files changed, 233 insertions(+), 196 deletions(-) diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index 29d494dbb5..b6fcf853ab 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -4737,20 +4737,46 @@ class BaseLLMHTTPHandler: model: str, websocket: Any, logging_obj: LiteLLMLoggingObj, - responses_api_provider_config: BaseResponsesAPIConfig, + responses_api_provider_config: Optional[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, + custom_llm_provider: Optional[str] = None, + **kwargs: Any, ): """ 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. + For providers with native websocket support (OpenAI, Azure): + - Opens a persistent WebSocket to the provider's /v1/responses endpoint + - Proxies response.create events bidirectionally for lower-latency agentic workflows + + For providers without native websocket support (all others): + - Uses ManagedResponsesWebSocketHandler which makes HTTP streaming calls + - Forwards events over the websocket connection """ + if responses_api_provider_config is None or not responses_api_provider_config.supports_native_websocket(): + from litellm.responses.streaming_iterator import ( + ManagedResponsesWebSocketHandler, + ) + + handler = ManagedResponsesWebSocketHandler( + websocket=websocket, + model=model, + logging_obj=logging_obj, + user_api_key_dict=user_api_key_dict, + litellm_metadata=litellm_metadata, + api_key=api_key, + api_base=api_base, + timeout=timeout, + custom_llm_provider=custom_llm_provider, + **kwargs, + ) + await handler.run() + return + import websockets from websockets.asyncio.client import ClientConnection @@ -4767,7 +4793,6 @@ class BaseLLMHTTPHandler: api_base=api_base, litellm_params={}, ) - # /responses -> wss:// URL ws_url = http_url.replace("https://", "wss://").replace("http://", "ws://") try: diff --git a/litellm/responses/main.py b/litellm/responses/main.py index 42f9d0d778..9c397aaaae 100644 --- a/litellm/responses/main.py +++ b/litellm/responses/main.py @@ -1729,11 +1729,6 @@ async def _aresponses_websocket( ) ) - 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 @@ -1748,6 +1743,9 @@ async def _aresponses_websocket( or get_secret_str("OPENAI_API_KEY") ) + # Extract params that we're passing explicitly to avoid duplicates in **kwargs + remaining_kwargs = {k: v for k, v in kwargs.items() if k not in {"user_api_key_dict", "litellm_metadata"}} + await base_llm_http_handler.async_responses_websocket( model=model, websocket=websocket, @@ -1758,4 +1756,6 @@ async def _aresponses_websocket( timeout=timeout, user_api_key_dict=kwargs.get("user_api_key_dict"), litellm_metadata=_build_litellm_metadata_for_ws(kwargs), + custom_llm_provider=_custom_llm_provider, + **remaining_kwargs, ) diff --git a/litellm/responses/streaming_iterator.py b/litellm/responses/streaming_iterator.py index 09d770b0fa..00392fc3cc 100644 --- a/litellm/responses/streaming_iterator.py +++ b/litellm/responses/streaming_iterator.py @@ -26,6 +26,7 @@ from litellm.types.llms.openai import ( OutputTextDeltaEvent, ResponseAPIUsage, ResponseCompletedEvent, + ResponsesAPIRequestParams, ResponsesAPIResponse, ResponsesAPIStreamEvents, ResponsesAPIStreamingResponse, @@ -874,37 +875,17 @@ class ResponsesWebSocketStreaming: # 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", +_RESPONSE_CREATE_PARAMS: frozenset = ( + ResponsesAPIRequestParams.__required_keys__ | ResponsesAPIRequestParams.__optional_keys__ ) -_MANAGED_WS_SKIP_KWARGS = frozenset( +_MANAGED_WS_SKIP_KWARGS: frozenset = frozenset( { - "litellm_logging_obj", - "litellm_call_id", - "aresponses", - "_aresponses_websocket", - "user_api_key_dict", + "litellm_logging_obj", + "litellm_call_id", + "aresponses", + "_aresponses_websocket", + "user_api_key_dict", } ) @@ -952,7 +933,7 @@ class ManagedResponsesWebSocketHandler: 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. + # In-memory session history: response_id → full accumulated message list. # 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. @@ -985,41 +966,27 @@ class ManagedResponsesWebSocketHandler: 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: + def _store_history(self, response_id: str, messages: List[Dict[str, Any]]) -> None: """ - Persist a turn's messages in the in-memory session store. + Store the complete accumulated message history for *response_id*. - *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`. + Replaces any prior value — callers are responsible for passing the full + history (prior turns + current input + new output). """ - prior: List[Dict[str, Any]] = self._session_history.get(response_id, []) - self._session_history[response_id] = prior + input_messages + output_messages + self._session_history[response_id] = messages @staticmethod def _extract_response_id(completed_event: Dict[str, Any]) -> Optional[str]: @@ -1027,8 +994,6 @@ class ManagedResponsesWebSocketHandler: 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: @@ -1040,7 +1005,7 @@ class ManagedResponsesWebSocketHandler: 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``. + Responses API message dicts suitable for the next turn's ``input``. """ resp_obj = completed_event.get("response", {}) if not isinstance(resp_obj, dict): @@ -1077,6 +1042,169 @@ class ManagedResponsesWebSocketHandler: return [item for item in input_val if isinstance(item, dict)] return [] + # ------------------------------------------------------------------ + # _process_response_create sub-methods + # ------------------------------------------------------------------ + + async def _parse_message(self, raw_message: str) -> Optional[Dict[str, Any]]: + """Parse raw WS text; return the message dict or None (JSON error / ignored type).""" + try: + msg_obj = json.loads(raw_message) + except json.JSONDecodeError: + await self._send_error("Invalid JSON in response.create event", "invalid_request_error") + return None + if msg_obj.get("type") != "response.create": + # Silently ignore non-response.create messages (e.g. warmup pings) + return None + return msg_obj + + @staticmethod + def _build_base_call_kwargs(msg_obj: Dict[str, Any]) -> Dict[str, Any]: + """ + Extract Responses API params from the event, handling both wire formats: + Nested: {"type": "response.create", "response": {"input": [...], ...}} + Flat: {"type": "response.create", "input": [...], "model": "...", ...} + """ + nested = msg_obj.get("response") + response_params: Dict[str, Any] = ( + nested + if isinstance(nested, dict) and nested + else {k: v for k, v in msg_obj.items() if k != "type"} + ) + return { + param: response_params[param] + for param in _RESPONSE_CREATE_PARAMS + if param in response_params and response_params[param] is not None + } + + def _apply_history( + self, + call_kwargs: Dict[str, Any], + previous_response_id: Optional[str], + current_messages: List[Dict[str, Any]], + prior_history: List[Dict[str, Any]], + ) -> None: + """Prepend in-memory turn history, or fall back to DB-based reconstruction.""" + if not previous_response_id: + return + if prior_history: + call_kwargs["input"] = prior_history + current_messages + verbose_logger.debug( + "ManagedResponsesWS: prepended %d history messages for previous_response_id=%s", + len(prior_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 + + def _inject_credentials( + self, call_kwargs: Dict[str, Any], event_model: Optional[str] + ) -> None: + """Inject connection-level credentials and metadata into call_kwargs.""" + 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 + # Only propagate custom_llm_provider when no per-request model override exists. + # If the payload specifies a different model, let litellm re-resolve the + # provider so we don't accidentally force the wrong backend. + 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) + + @staticmethod + def _update_proxy_request(call_kwargs: Dict[str, Any], model: str) -> None: + """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 not isinstance(proxy_server_request, dict): + return + 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 = {**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 + + async def _stream_and_forward( + self, model: str, call_kwargs: Dict[str, Any] + ) -> Optional[Dict[str, Any]]: + """ + Stream ``litellm.aresponses`` and forward every chunk over the WebSocket. + + Captures the ``response.completed`` event type from the chunk object + directly (before serialization) to avoid a redundant JSON round-trip on + every chunk. Returns the completed event dict, or ``None``. + """ + completed_event: Optional[Dict[str, Any]] = None + stream_response = await litellm.aresponses(model=model, **call_kwargs) + async for chunk in stream_response: # type: ignore[union-attr] + if chunk is None: + continue + # Read type from the object before serializing to avoid double JSON parse + chunk_type = getattr(chunk, "type", None) or ( + chunk.get("type") if isinstance(chunk, dict) else None + ) + serialized = self._serialize_chunk(chunk) + if serialized is None: + continue + if chunk_type == "response.completed" and completed_event is None: + try: + completed_event = json.loads(serialized) + 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 completed_event # Client disconnected + return completed_event + + def _save_turn_history( + self, + completed_event: Optional[Dict[str, Any]], + prior_history: List[Dict[str, Any]], + current_messages: List[Dict[str, Any]], + ) -> None: + """Store this turn in in-memory history for future previous_response_id lookups.""" + if completed_event is None: + return + new_response_id = self._extract_response_id(completed_event) + if not new_response_id: + return + output_msgs = self._extract_output_messages(completed_event) + all_messages = prior_history + current_messages + output_msgs + self._store_history(new_response_id, all_messages) + verbose_logger.debug( + "ManagedResponsesWS: stored %d messages for response_id=%s", + len(all_messages), + new_response_id, + ) + + # ------------------------------------------------------------------ + # Core request handler + # ------------------------------------------------------------------ + async def _process_response_create(self, raw_message: str) -> None: """ Parse one ``response.create`` event, call ``litellm.aresponses(stream=True)``, @@ -1097,157 +1225,41 @@ class ManagedResponsesWebSocketHandler: 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") + msg_obj = await self._parse_message(raw_message) + if msg_obj is None: 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 = self._build_base_call_kwargs(msg_obj) 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) + event_model: Optional[str] = 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 - # --------------------------------------------------------------------------- + current_messages = self._input_to_messages(call_kwargs.get("input")) - # 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) + # Fetch history once; reused in both _apply_history and _save_turn_history + prior_history = ( + self._get_history_messages(previous_response_id) + if previous_response_id + else [] + ) - # 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.) + self._apply_history(call_kwargs, previous_response_id, current_messages, prior_history) + self._inject_credentials(call_kwargs, event_model) + self._update_proxy_request(call_kwargs, model) 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 - + completed_event = await self._stream_and_forward(model, call_kwargs) except Exception as exc: - verbose_logger.exception("ManagedResponsesWS: error processing response.create: %s", 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, - ) - # --------------------------------------------------------------------------- + self._save_turn_history(completed_event, prior_history, current_messages) # ------------------------------------------------------------------ # Main entry point