mirror of
https://github.com/tiennm99/litellm.git
synced 2026-07-30 00:21:04 +00:00
Add support for responses websocket for all providers
This commit is contained in:
@@ -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:
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user