mirror of
https://github.com/BerriAI/litellm.git
synced 2026-08-28 05:25:59 +00:00
fix(responses): add in-memory session tracking to ManagedResponsesWebSocketHandler for previous_response_id
Spend logs are written asynchronously in batches, so a DB lookup for a just-completed response's spend log races and returns zero rows. Replace the DB-only fallback with an in-memory session store (_session_history) that is populated after each response.completed event and consulted at the start of the next response.create. This makes same-connection multi-turn reliable without any timing dependency on the DB write queue. Made-with: Cursor
This commit is contained in:
parent
c1136348f0
commit
cc650f4865
2 changed files with 411 additions and 1 deletions
|
|
@ -3,7 +3,7 @@ import json
|
|||
import time
|
||||
import traceback
|
||||
from datetime import datetime
|
||||
from typing import TYPE_CHECKING, Any, Dict, Optional
|
||||
from typing import TYPE_CHECKING, Any, Dict, List, Optional
|
||||
|
||||
import httpx
|
||||
|
||||
|
|
@ -839,3 +839,412 @@ class ResponsesWebSocketStreaming:
|
|||
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}")
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue