mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
chore(typing): wip snapshot of basedpyright Any cleanup (streaming_iterator)
Intermediate checkpoint while a background pass fixes LIT001 (mutable-annotation) fallout in streaming_iterator.py introduced by the Any-cleanup itself. Not yet passing make pre-commit; follow-up commit will finish this.
This commit is contained in:
parent
5bd9e12c85
commit
4051c6c4ff
3 changed files with 258 additions and 197 deletions
|
|
@ -14,8 +14,10 @@ from typing import (
|
|||
Dict,
|
||||
Iterator,
|
||||
List,
|
||||
Mapping,
|
||||
NoReturn,
|
||||
Optional,
|
||||
Sequence,
|
||||
Union,
|
||||
cast,
|
||||
)
|
||||
|
|
@ -104,13 +106,13 @@ def _json_loads_object(raw: Union[str, bytes]) -> object:
|
|||
return json.loads(raw) # any-ok: json.loads is the untyped-JSON boundary; callers narrow via _require_*
|
||||
|
||||
|
||||
def _require_dict(value: object) -> dict[str, object]:
|
||||
def _require_dict(value: object) -> Mapping[str, object]:
|
||||
if isinstance(value, dict):
|
||||
return value
|
||||
raise ValueError(f"Expected a JSON object, got: {value!r}")
|
||||
|
||||
|
||||
def _require_list(value: object) -> list[object]:
|
||||
def _require_list(value: object) -> Sequence[object]:
|
||||
if isinstance(value, list):
|
||||
return value
|
||||
raise ValueError(f"Expected a JSON array, got: {value!r}")
|
||||
|
|
@ -697,17 +699,17 @@ class CustomStreamWrapper:
|
|||
except Exception as e:
|
||||
raise e
|
||||
|
||||
def model_response_creator(self, chunk: Optional[dict[str, object]] = None, hidden_params: Optional[dict] = None):
|
||||
def model_response_creator(
|
||||
self,
|
||||
chunk: Optional[Mapping[str, object]] = None,
|
||||
hidden_params: Optional[Mapping[str, object]] = None,
|
||||
):
|
||||
_model = self._cached_model_name
|
||||
_logging_obj_llm_provider = self._cached_logging_llm_provider
|
||||
|
||||
if chunk is None:
|
||||
args: dict[str, object] = {"model": _model}
|
||||
else:
|
||||
chunk.pop("model", None)
|
||||
args = {"model": _model}
|
||||
if chunk:
|
||||
args.update({k: v for k, v in chunk.items() if k != "stream"})
|
||||
args: dict[str, object] = {"model": _model}
|
||||
if chunk:
|
||||
args.update({k: v for k, v in chunk.items() if k not in ("model", "stream")})
|
||||
|
||||
model_response = ModelResponseStream.model_validate(args)
|
||||
if self.response_id is not None:
|
||||
|
|
|
|||
|
|
@ -161,6 +161,7 @@ from .http_handler import get_shared_realtime_ssl_context
|
|||
|
||||
if TYPE_CHECKING:
|
||||
from aiohttp import ClientSession
|
||||
from starlette.websockets import WebSocket
|
||||
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj
|
||||
from litellm.llms.anthropic.experimental_pass_through.messages.streaming_iterator import (
|
||||
|
|
@ -6137,7 +6138,7 @@ class BaseLLMHTTPHandler:
|
|||
async def async_responses_websocket(
|
||||
self,
|
||||
model: str,
|
||||
websocket: Any,
|
||||
websocket: "WebSocket",
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
responses_api_provider_config: Optional[BaseResponsesAPIConfig],
|
||||
api_base: Optional[str] = None,
|
||||
|
|
@ -6300,7 +6301,10 @@ class BaseLLMHTTPHandler:
|
|||
|
||||
except websockets.exceptions.InvalidStatusCode as e: # type: ignore
|
||||
verbose_logger.exception(f"Error connecting to responses WS backend: {e}")
|
||||
await websocket.close(code=e.status_code, reason=_redact_string(str(e)))
|
||||
await websocket.close(
|
||||
code=e.status_code if isinstance(e.status_code, int) else 1011,
|
||||
reason=_redact_string(str(e)),
|
||||
)
|
||||
except Exception as e:
|
||||
verbose_logger.exception(f"Error in responses WS: {e}")
|
||||
try:
|
||||
|
|
|
|||
|
|
@ -8,10 +8,21 @@ import uuid
|
|||
from datetime import datetime
|
||||
from functools import lru_cache
|
||||
from types import MappingProxyType
|
||||
from typing import Any, Dict, List, Literal, Mapping, Optional
|
||||
from typing import (
|
||||
TYPE_CHECKING,
|
||||
Any,
|
||||
Dict,
|
||||
List,
|
||||
Literal,
|
||||
Mapping,
|
||||
Optional,
|
||||
Protocol,
|
||||
Union,
|
||||
)
|
||||
|
||||
import httpx
|
||||
from openai._streaming import SSEDecoder
|
||||
from pydantic import BaseModel, TypeAdapter
|
||||
|
||||
import litellm
|
||||
from litellm.constants import (
|
||||
|
|
@ -29,10 +40,56 @@ from litellm.litellm_core_utils.llm_response_utils.response_metadata import (
|
|||
from litellm.litellm_core_utils.thread_pool_executor import executor
|
||||
from litellm.llms.base_llm.responses.transformation import BaseResponsesAPIConfig
|
||||
from litellm.responses.utils import ResponseAPILoggingUtils, ResponsesAPIRequestUtils
|
||||
from litellm.types.llms.openai import ResponsesAPIStreamEvents
|
||||
from litellm.types.guardrails import PresidioPerRequestConfig
|
||||
from litellm.types.llms.openai import (
|
||||
PART_UNION_TYPES,
|
||||
ResponsesAPIResponse,
|
||||
ResponsesAPIStreamEvents,
|
||||
ResponsesAPIStreamingResponse,
|
||||
)
|
||||
from litellm.types.utils import CallTypes
|
||||
from litellm.utils import async_post_call_success_deployment_hook
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from starlette.websockets import WebSocket
|
||||
from websockets.asyncio.client import ClientConnection
|
||||
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
|
||||
_JSON_OBJECT_ADAPTER: TypeAdapter[Dict[str, object]] = TypeAdapter(Dict[str, object])
|
||||
_DUMP_ADAPTER: TypeAdapter[Dict[str, Any]] = TypeAdapter(Dict[str, Any])
|
||||
|
||||
|
||||
class PIIMaskingGuardrail(Protocol):
|
||||
"""Structural type for guardrails that expose the Presidio PII masking interface.
|
||||
|
||||
Any guardrail implementing this shape works here (duck typing), not just
|
||||
the Presidio guardrail hook, to avoid a layering violation (SDK importing
|
||||
from the proxy-only guardrails package).
|
||||
"""
|
||||
|
||||
output_parse_pii: bool
|
||||
apply_to_output: bool
|
||||
|
||||
def get_presidio_settings_from_request_data(
|
||||
self, data: Dict[str, object]
|
||||
) -> Optional[PresidioPerRequestConfig]: ...
|
||||
|
||||
async def check_pii(
|
||||
self,
|
||||
text: str,
|
||||
output_parse_pii: bool,
|
||||
presidio_config: Optional[PresidioPerRequestConfig],
|
||||
request_data: Dict[str, object],
|
||||
) -> str: ...
|
||||
|
||||
|
||||
def _call_unmask_pii_text(guardrail: PIIMaskingGuardrail, text: str, pii_tokens: Dict[str, str]) -> str:
|
||||
# any-ok: _unmask_pii_text is a private helper on the concrete guardrail class;
|
||||
# dispatched dynamically so PIIMaskingGuardrail need not declare private methods.
|
||||
unmask = getattr(guardrail, "_unmask_pii_text")
|
||||
return unmask(text, pii_tokens)
|
||||
|
||||
|
||||
@lru_cache(maxsize=1)
|
||||
def _get_openai_response_types():
|
||||
|
|
@ -130,7 +187,7 @@ class BaseResponsesAPIStreamingIterator:
|
|||
self.logging_obj = logging_obj
|
||||
self.finished = False
|
||||
self.responses_api_provider_config = responses_api_provider_config
|
||||
self.completed_response: Optional[Any] = None
|
||||
self.completed_response: Optional[ResponsesAPIStreamingResponse] = None
|
||||
self.start_time = getattr(logging_obj, "start_time", datetime.now())
|
||||
self._failure_handled = False # Track if failure handler has been called
|
||||
self._yielded_first_chunk = False
|
||||
|
|
@ -175,7 +232,7 @@ class BaseResponsesAPIStreamingIterator:
|
|||
llm_provider=self.custom_llm_provider or "",
|
||||
)
|
||||
|
||||
def _process_chunk(self, chunk) -> Optional[Any]:
|
||||
def _process_chunk(self, chunk: str) -> Optional[ResponsesAPIStreamingResponse]:
|
||||
"""Process a single chunk of data from the stream"""
|
||||
if not chunk:
|
||||
return None
|
||||
|
|
@ -298,9 +355,9 @@ class BaseResponsesAPIStreamingIterator:
|
|||
self.completed_response = openai_responses_api_chunk
|
||||
# Add cost to usage object if include_cost_in_streaming_usage is True
|
||||
if litellm.include_cost_in_streaming_usage and self.logging_obj is not None:
|
||||
response_obj: Optional[Any] = getattr(openai_responses_api_chunk, "response", None)
|
||||
response_obj = getattr(openai_responses_api_chunk, "response", None)
|
||||
if response_obj:
|
||||
usage_obj: Optional[Any] = getattr(response_obj, "usage", None)
|
||||
usage_obj = getattr(response_obj, "usage", None)
|
||||
if usage_obj is not None:
|
||||
try:
|
||||
cost: Optional[float] = self.logging_obj._response_cost_calculator(
|
||||
|
|
@ -403,7 +460,7 @@ class BaseResponsesAPIStreamingIterator:
|
|||
)
|
||||
self._handle_failure(exception)
|
||||
|
||||
def _record_failed_response_usage(self, response_obj: Optional[Any]) -> None:
|
||||
def _record_failed_response_usage(self, response_obj: Optional[ResponsesAPIResponse]) -> None:
|
||||
if response_obj is None or self.logging_obj is None:
|
||||
return
|
||||
usage_obj = getattr(response_obj, "usage", None)
|
||||
|
|
@ -453,14 +510,13 @@ class BaseResponsesAPIStreamingIterator:
|
|||
is_pre_first_chunk=not self._yielded_first_chunk,
|
||||
)
|
||||
|
||||
def _get_completed_response_object(self) -> Optional[Any]:
|
||||
openai_types = _get_openai_response_types()
|
||||
def _get_completed_response_object(self) -> Optional[ResponsesAPIResponse]:
|
||||
completed_response = self.completed_response
|
||||
if isinstance(completed_response, openai_types.ResponsesAPIResponse):
|
||||
if isinstance(completed_response, ResponsesAPIResponse):
|
||||
return completed_response
|
||||
|
||||
response_obj = getattr(completed_response, "response", None)
|
||||
if isinstance(response_obj, openai_types.ResponsesAPIResponse):
|
||||
if isinstance(response_obj, ResponsesAPIResponse):
|
||||
return response_obj
|
||||
|
||||
return None
|
||||
|
|
@ -529,7 +585,9 @@ class BaseResponsesAPIStreamingIterator:
|
|||
|
||||
self._completed_response_cached = True
|
||||
|
||||
async def _call_post_streaming_deployment_hook(self, chunk):
|
||||
async def _call_post_streaming_deployment_hook(
|
||||
self, chunk: ResponsesAPIStreamingResponse
|
||||
) -> ResponsesAPIStreamingResponse:
|
||||
"""
|
||||
Allow callbacks to modify streaming chunks before returning (parity with chat).
|
||||
"""
|
||||
|
|
@ -547,13 +605,13 @@ class BaseResponsesAPIStreamingIterator:
|
|||
except Exception:
|
||||
typed_call_type = None
|
||||
|
||||
request_data = self.request_data or getattr(self.logging_obj, "model_call_details", {})
|
||||
callbacks = getattr(litellm, "callbacks", None) or []
|
||||
request_data = self.request_data or self.logging_obj.model_call_details
|
||||
hooks_ran = False
|
||||
for callback in callbacks:
|
||||
if hasattr(callback, "async_post_call_streaming_deployment_hook"):
|
||||
for callback in litellm.callbacks:
|
||||
hook = getattr(callback, "async_post_call_streaming_deployment_hook", None)
|
||||
if hook is not None:
|
||||
hooks_ran = True
|
||||
result = await callback.async_post_call_streaming_deployment_hook(
|
||||
result = await hook(
|
||||
request_data=request_data,
|
||||
response_chunk=chunk,
|
||||
call_type=typed_call_type,
|
||||
|
|
@ -566,7 +624,9 @@ class BaseResponsesAPIStreamingIterator:
|
|||
except Exception:
|
||||
return chunk
|
||||
|
||||
async def call_post_streaming_hooks_for_testing(self, chunk):
|
||||
async def call_post_streaming_hooks_for_testing(
|
||||
self, chunk: ResponsesAPIStreamingResponse
|
||||
) -> ResponsesAPIStreamingResponse:
|
||||
"""
|
||||
Helper to invoke streaming deployment hooks explicitly (used in tests).
|
||||
"""
|
||||
|
|
@ -583,15 +643,12 @@ class BaseResponsesAPIStreamingIterator:
|
|||
if isinstance(self.request_data, dict):
|
||||
request_payload.update(self.request_data)
|
||||
try:
|
||||
if hasattr(self.logging_obj, "model_call_details"):
|
||||
request_payload.update(self.logging_obj.model_call_details)
|
||||
request_payload.update(self.logging_obj.model_call_details)
|
||||
except Exception:
|
||||
pass
|
||||
if "litellm_params" not in request_payload:
|
||||
try:
|
||||
request_payload["litellm_params"] = getattr(self.logging_obj, "model_call_details", {}).get(
|
||||
"litellm_params", {}
|
||||
)
|
||||
request_payload["litellm_params"] = self.logging_obj.model_call_details.get("litellm_params", {})
|
||||
except Exception:
|
||||
request_payload["litellm_params"] = {}
|
||||
|
||||
|
|
@ -668,14 +725,16 @@ class BaseResponsesAPIStreamingIterator:
|
|||
pass
|
||||
|
||||
|
||||
async def call_post_streaming_hooks_for_testing(iterator, chunk):
|
||||
async def call_post_streaming_hooks_for_testing(
|
||||
iterator: object, chunk: ResponsesAPIStreamingResponse
|
||||
) -> ResponsesAPIStreamingResponse:
|
||||
"""
|
||||
Module-level helper for tests to ensure hooks can be invoked even if the iterator is wrapped.
|
||||
"""
|
||||
hook_fn = getattr(iterator, "_call_post_streaming_deployment_hook", None)
|
||||
if hook_fn is None:
|
||||
return chunk
|
||||
return await hook_fn(chunk)
|
||||
return await hook_fn(chunk) # any-ok: test helper must dispatch onto arbitrary wrapped iterator doubles
|
||||
|
||||
|
||||
class ResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator):
|
||||
|
|
@ -709,7 +768,7 @@ class ResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator):
|
|||
def __aiter__(self):
|
||||
return self
|
||||
|
||||
async def __anext__(self) -> Any:
|
||||
async def __anext__(self) -> ResponsesAPIStreamingResponse:
|
||||
try:
|
||||
self._check_max_streaming_duration()
|
||||
while True:
|
||||
|
|
@ -791,7 +850,7 @@ class SyncResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator):
|
|||
def __iter__(self):
|
||||
return self
|
||||
|
||||
def __next__(self):
|
||||
def __next__(self) -> ResponsesAPIStreamingResponse:
|
||||
try:
|
||||
self._check_max_streaming_duration()
|
||||
while True:
|
||||
|
|
@ -882,10 +941,10 @@ class MockResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator):
|
|||
|
||||
def _set_events_from_response(
|
||||
self,
|
||||
transformed: Any,
|
||||
transformed: ResponsesAPIResponse,
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
) -> None:
|
||||
self._events = _build_synthetic_response_events(
|
||||
self._events: List[ResponsesAPIStreamingResponse] = _build_synthetic_response_events(
|
||||
transformed=transformed,
|
||||
logging_obj=logging_obj,
|
||||
chunk_size=self.CHUNK_SIZE,
|
||||
|
|
@ -896,13 +955,12 @@ class MockResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator):
|
|||
def __aiter__(self):
|
||||
return self
|
||||
|
||||
async def __anext__(self) -> Any:
|
||||
async def __anext__(self) -> ResponsesAPIStreamingResponse:
|
||||
if self._idx >= len(self._events):
|
||||
raise StopAsyncIteration
|
||||
evt = self._events[self._idx]
|
||||
self._idx += 1
|
||||
openai_types = _get_openai_response_types()
|
||||
if getattr(evt, "type", None) == openai_types.ResponsesAPIStreamEvents.RESPONSE_COMPLETED:
|
||||
if getattr(evt, "type", None) == ResponsesAPIStreamEvents.RESPONSE_COMPLETED:
|
||||
self.completed_response = evt
|
||||
self._log_completed_response(is_async=True)
|
||||
return evt
|
||||
|
|
@ -910,13 +968,12 @@ class MockResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator):
|
|||
def __iter__(self):
|
||||
return self
|
||||
|
||||
def __next__(self) -> Any:
|
||||
def __next__(self) -> ResponsesAPIStreamingResponse:
|
||||
if self._idx >= len(self._events):
|
||||
raise StopIteration
|
||||
evt = self._events[self._idx]
|
||||
self._idx += 1
|
||||
openai_types = _get_openai_response_types()
|
||||
if getattr(evt, "type", None) == openai_types.ResponsesAPIStreamEvents.RESPONSE_COMPLETED:
|
||||
if getattr(evt, "type", None) == ResponsesAPIStreamEvents.RESPONSE_COMPLETED:
|
||||
self.completed_response = evt
|
||||
self._log_completed_response(is_async=False)
|
||||
return evt
|
||||
|
|
@ -925,7 +982,7 @@ class MockResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator):
|
|||
class CachedResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator):
|
||||
def __init__(
|
||||
self,
|
||||
response: Any,
|
||||
response: ResponsesAPIResponse,
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
request_data: Optional[Dict[str, Any]] = None,
|
||||
call_type: Optional[str] = None,
|
||||
|
|
@ -933,7 +990,7 @@ class CachedResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator):
|
|||
BaseResponsesAPIStreamingIterator.__init__(
|
||||
self,
|
||||
response=httpx.Response(200),
|
||||
model=getattr(response, "model", ""),
|
||||
model=response.model or "",
|
||||
responses_api_provider_config=None,
|
||||
logging_obj=logging_obj,
|
||||
litellm_metadata=None,
|
||||
|
|
@ -943,13 +1000,13 @@ class CachedResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator):
|
|||
)
|
||||
self._completed_response_cache_hit = True
|
||||
self._persist_completed_response_before_logging = False
|
||||
self._events: List[Any] = []
|
||||
self._events: List[ResponsesAPIStreamingResponse] = []
|
||||
self._idx = 0
|
||||
self._set_events_from_response(transformed=response, logging_obj=logging_obj)
|
||||
|
||||
def _set_events_from_response(
|
||||
self,
|
||||
transformed: Any,
|
||||
transformed: ResponsesAPIResponse,
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
) -> None:
|
||||
self._events = _build_synthetic_response_events(
|
||||
|
|
@ -963,13 +1020,12 @@ class CachedResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator):
|
|||
def __aiter__(self):
|
||||
return self
|
||||
|
||||
async def __anext__(self) -> Any:
|
||||
async def __anext__(self) -> ResponsesAPIStreamingResponse:
|
||||
if self._idx >= len(self._events):
|
||||
raise StopAsyncIteration
|
||||
evt = self._events[self._idx]
|
||||
self._idx += 1
|
||||
openai_types = _get_openai_response_types()
|
||||
if getattr(evt, "type", None) == openai_types.ResponsesAPIStreamEvents.RESPONSE_COMPLETED:
|
||||
if getattr(evt, "type", None) == ResponsesAPIStreamEvents.RESPONSE_COMPLETED:
|
||||
self.completed_response = evt
|
||||
self._log_completed_response(is_async=True)
|
||||
return evt
|
||||
|
|
@ -977,23 +1033,29 @@ class CachedResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator):
|
|||
def __iter__(self):
|
||||
return self
|
||||
|
||||
def __next__(self) -> Any:
|
||||
def __next__(self) -> ResponsesAPIStreamingResponse:
|
||||
if self._idx >= len(self._events):
|
||||
raise StopIteration
|
||||
evt = self._events[self._idx]
|
||||
self._idx += 1
|
||||
openai_types = _get_openai_response_types()
|
||||
if getattr(evt, "type", None) == openai_types.ResponsesAPIStreamEvents.RESPONSE_COMPLETED:
|
||||
if getattr(evt, "type", None) == ResponsesAPIStreamEvents.RESPONSE_COMPLETED:
|
||||
self.completed_response = evt
|
||||
self._log_completed_response(is_async=False)
|
||||
return evt
|
||||
|
||||
|
||||
def _dump_response_object(obj: Any) -> Dict[str, Any]:
|
||||
if hasattr(obj, "model_dump"):
|
||||
def _dump_response_object(obj: object) -> Dict[str, Any]:
|
||||
"""Normalize a Responses API output item to a plain dict.
|
||||
|
||||
Returns ``Dict[str, Any]`` (not ``object``) because callers splat this
|
||||
payload into strongly-typed ``BaseLiteLLMOpenAIResponseObject`` subclass
|
||||
constructors (e.g. ``logprobs=payload.get("logprobs")``); an ``object``
|
||||
value type would fail those field-typed constructor calls.
|
||||
"""
|
||||
if isinstance(obj, BaseModel):
|
||||
return obj.model_dump()
|
||||
if isinstance(obj, dict):
|
||||
return obj
|
||||
return _DUMP_ADAPTER.validate_python(obj)
|
||||
return {}
|
||||
|
||||
|
||||
|
|
@ -1002,8 +1064,8 @@ def _build_response_status_event(
|
|||
"response.created",
|
||||
"response.in_progress",
|
||||
],
|
||||
transformed: Any,
|
||||
) -> Any:
|
||||
transformed: ResponsesAPIResponse,
|
||||
) -> ResponsesAPIStreamingResponse:
|
||||
openai_types = _get_openai_response_types()
|
||||
in_progress_response = transformed.model_copy(
|
||||
deep=True,
|
||||
|
|
@ -1020,10 +1082,10 @@ def _build_content_part_done_event(
|
|||
output_index: int,
|
||||
content_index: int,
|
||||
part_payload: Dict[str, Any],
|
||||
) -> Optional[Any]:
|
||||
) -> Optional[ResponsesAPIStreamingResponse]:
|
||||
openai_types = _get_openai_response_types()
|
||||
part_type = part_payload.get("type")
|
||||
part: Any
|
||||
part: PART_UNION_TYPES
|
||||
if part_type == "output_text":
|
||||
annotations = [
|
||||
openai_types.BaseLiteLLMOpenAIResponseObject(**annotation)
|
||||
|
|
@ -1059,7 +1121,7 @@ def _build_content_part_done_event(
|
|||
|
||||
def _add_text_like_part_events(
|
||||
*,
|
||||
events: List[Any],
|
||||
events: List[ResponsesAPIStreamingResponse],
|
||||
item_id: str,
|
||||
output_index: int,
|
||||
content_index: int,
|
||||
|
|
@ -1125,13 +1187,13 @@ def _add_text_like_part_events(
|
|||
|
||||
def _build_synthetic_response_events(
|
||||
*,
|
||||
transformed: Any,
|
||||
transformed: ResponsesAPIResponse,
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
chunk_size: int,
|
||||
) -> List[Any]:
|
||||
) -> List[ResponsesAPIStreamingResponse]:
|
||||
openai_types = _get_openai_response_types()
|
||||
if litellm.include_cost_in_streaming_usage and logging_obj is not None:
|
||||
usage_obj: Optional[Any] = getattr(transformed, "usage", None)
|
||||
usage_obj = transformed.usage
|
||||
if usage_obj is not None:
|
||||
try:
|
||||
cost: Optional[float] = logging_obj._response_cost_calculator(result=transformed)
|
||||
|
|
@ -1140,13 +1202,13 @@ def _build_synthetic_response_events(
|
|||
except Exception:
|
||||
pass
|
||||
|
||||
events: List[Any] = [
|
||||
events: List[ResponsesAPIStreamingResponse] = [
|
||||
_build_response_status_event(openai_types.ResponsesAPIStreamEvents.RESPONSE_CREATED, transformed),
|
||||
_build_response_status_event(openai_types.ResponsesAPIStreamEvents.RESPONSE_IN_PROGRESS, transformed),
|
||||
]
|
||||
|
||||
sequence_number = 0
|
||||
for output_index, output_item in enumerate(getattr(transformed, "output", []) or []):
|
||||
for output_index, output_item in enumerate(transformed.output):
|
||||
output_item_payload = _dump_response_object(output_item)
|
||||
item_id = str(output_item_payload.get("id") or transformed.id)
|
||||
item_type = output_item_payload.get("type")
|
||||
|
|
@ -1279,6 +1341,22 @@ RESPONSES_WS_LOGGED_EVENT_TYPES = [
|
|||
RESPONSES_WS_MASKABLE_TEXT_BLOCK_TYPES = frozenset({"input_text", "output_text", "text"})
|
||||
|
||||
|
||||
def _parse_json_object(raw: Union[str, bytes]) -> Optional[Dict[str, object]]:
|
||||
try:
|
||||
parsed = json.loads(raw)
|
||||
except (json.JSONDecodeError, TypeError):
|
||||
return None
|
||||
return _as_str_object_dict(parsed)
|
||||
|
||||
|
||||
def _as_str_object_dict(value: object) -> Optional[Dict[str, object]]:
|
||||
return _JSON_OBJECT_ADAPTER.validate_python(value) if isinstance(value, dict) else None
|
||||
|
||||
|
||||
def _as_object_list(value: object) -> List[object]:
|
||||
return value if isinstance(value, list) else []
|
||||
|
||||
|
||||
class ResponsesWebSocketStreaming:
|
||||
"""
|
||||
Manages bidirectional WebSocket forwarding for the Responses API
|
||||
|
|
@ -1294,55 +1372,47 @@ class ResponsesWebSocketStreaming:
|
|||
|
||||
def __init__(
|
||||
self,
|
||||
websocket: Any,
|
||||
backend_ws: Any,
|
||||
websocket: "WebSocket",
|
||||
backend_ws: "ClientConnection",
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
user_api_key_dict: Optional[Any] = None,
|
||||
request_data: Optional[Dict] = None,
|
||||
user_api_key_dict: Optional["UserAPIKeyAuth"] = None,
|
||||
request_data: Optional[Dict[str, Any]] = None,
|
||||
first_message: Optional[str] = None,
|
||||
guardrail_callbacks: Optional[List[Any]] = None,
|
||||
output_guardrail_callbacks: Optional[List[Any]] = None,
|
||||
guardrail_callbacks: Optional[List[PIIMaskingGuardrail]] = None,
|
||||
output_guardrail_callbacks: Optional[List[PIIMaskingGuardrail]] = None,
|
||||
authorized_model: Optional[str] = None,
|
||||
):
|
||||
self.websocket = websocket
|
||||
self.backend_ws = backend_ws
|
||||
self.logging_obj = logging_obj
|
||||
self.user_api_key_dict = user_api_key_dict
|
||||
self.request_data: Dict = request_data or {}
|
||||
self.messages: list[Dict] = []
|
||||
self.request_data: Dict[str, Any] = request_data or {}
|
||||
self.messages: List[Dict[str, object]] = []
|
||||
self.input_messages: list[Dict[str, str]] = []
|
||||
self.first_message = first_message
|
||||
self.guardrail_callbacks: List[Any] = guardrail_callbacks or []
|
||||
self.output_guardrail_callbacks: List[Any] = output_guardrail_callbacks or []
|
||||
self.guardrail_callbacks: List[PIIMaskingGuardrail] = guardrail_callbacks or []
|
||||
self.output_guardrail_callbacks: List[PIIMaskingGuardrail] = output_guardrail_callbacks or []
|
||||
# Model name authorized at connection time; enforced on every
|
||||
# response.create frame to prevent deployment-substitution attacks.
|
||||
self.authorized_model: Optional[str] = authorized_model
|
||||
|
||||
def _should_store_event(self, event_obj: dict) -> bool:
|
||||
def _should_store_event(self, event_obj: Dict[str, object]) -> bool:
|
||||
return event_obj.get("type") in RESPONSES_WS_LOGGED_EVENT_TYPES
|
||||
|
||||
def _store_event(self, event: Any) -> None:
|
||||
if isinstance(event, bytes):
|
||||
event = event.decode("utf-8")
|
||||
if isinstance(event, str):
|
||||
try:
|
||||
event_obj = json.loads(event)
|
||||
except (json.JSONDecodeError, TypeError):
|
||||
return
|
||||
else:
|
||||
event_obj = event
|
||||
def _store_event(self, event: Union[str, bytes]) -> None:
|
||||
decoded = event.decode("utf-8") if isinstance(event, bytes) else event
|
||||
event_obj = _parse_json_object(decoded)
|
||||
if event_obj is None:
|
||||
return
|
||||
|
||||
if self._should_store_event(event_obj):
|
||||
self.messages.append(event_obj)
|
||||
|
||||
def _collect_input_from_client_event(self, message: Any) -> None:
|
||||
def _collect_input_from_client_event(self, message: Union[str, Dict[str, object]]) -> None:
|
||||
"""Extract user input content from response.create for logging."""
|
||||
try:
|
||||
if isinstance(message, str):
|
||||
msg_obj = json.loads(message)
|
||||
elif isinstance(message, dict):
|
||||
msg_obj = message
|
||||
else:
|
||||
msg_obj = _parse_json_object(message) if isinstance(message, str) else message
|
||||
if msg_obj is None:
|
||||
return
|
||||
|
||||
if msg_obj.get("type") != "response.create":
|
||||
|
|
@ -1367,10 +1437,10 @@ class ResponsesWebSocketStreaming:
|
|||
text = c.get("text", "")
|
||||
if text:
|
||||
self.input_messages.append({"role": "user", "content": text})
|
||||
except (json.JSONDecodeError, AttributeError, TypeError):
|
||||
except (AttributeError, TypeError):
|
||||
pass
|
||||
|
||||
def _store_input(self, message: Any) -> None:
|
||||
def _store_input(self, message: str) -> None:
|
||||
self._collect_input_from_client_event(message)
|
||||
if self.logging_obj:
|
||||
self.logging_obj.pre_call(input=message, api_key="")
|
||||
|
|
@ -1390,9 +1460,9 @@ class ResponsesWebSocketStreaming:
|
|||
try:
|
||||
while True:
|
||||
try:
|
||||
raw_response = await self.backend_ws.recv(decode=False) # type: ignore[union-attr]
|
||||
raw_response = await self.backend_ws.recv(decode=False)
|
||||
except TypeError:
|
||||
raw_response = await self.backend_ws.recv() # type: ignore[union-attr, assignment]
|
||||
raw_response = await self.backend_ws.recv()
|
||||
|
||||
if isinstance(raw_response, bytes):
|
||||
response_str = raw_response.decode("utf-8")
|
||||
|
|
@ -1408,10 +1478,8 @@ class ResponsesWebSocketStreaming:
|
|||
# before response.completed arrives. The client receives only the
|
||||
# masked response.completed.
|
||||
if self.output_guardrail_callbacks:
|
||||
try:
|
||||
_evt_type = json.loads(response_str).get("type")
|
||||
except (json.JSONDecodeError, TypeError):
|
||||
_evt_type = None
|
||||
_parsed_for_type = _parse_json_object(response_str)
|
||||
_evt_type = _parsed_for_type.get("type") if _parsed_for_type is not None else None
|
||||
if _evt_type in self._DELTA_EVENT_TYPES or _evt_type in self._OUTPUT_DONE_EVENT_TYPES:
|
||||
continue
|
||||
|
||||
|
|
@ -1431,7 +1499,7 @@ class ResponsesWebSocketStreaming:
|
|||
finally:
|
||||
await self._log_messages()
|
||||
|
||||
def _enforce_authorized_model(self, msg_obj: dict) -> bool:
|
||||
def _enforce_authorized_model(self, msg_obj: Dict[str, object]) -> bool:
|
||||
"""
|
||||
Overwrite any ``model`` field in a ``response.create`` frame with the
|
||||
connection-authorized model to prevent deployment-substitution attacks.
|
||||
|
|
@ -1472,9 +1540,8 @@ class ResponsesWebSocketStreaming:
|
|||
|
||||
Non-``response.create`` messages are returned unchanged.
|
||||
"""
|
||||
try:
|
||||
msg_obj = json.loads(message)
|
||||
except (json.JSONDecodeError, TypeError):
|
||||
msg_obj = _parse_json_object(message)
|
||||
if msg_obj is None:
|
||||
return message
|
||||
|
||||
if msg_obj.get("type") != "response.create":
|
||||
|
|
@ -1497,10 +1564,10 @@ class ResponsesWebSocketStreaming:
|
|||
# nested: {"type": "response.create", "response": {"input": ..., "instructions": ...}}
|
||||
# Mask "input" and "instructions" in both shapes so PII is never
|
||||
# forwarded unmasked regardless of where the client places it.
|
||||
nested_response = msg_obj.get("response") if isinstance(msg_obj.get("response"), dict) else None
|
||||
nested_response = msg_obj.get("response")
|
||||
text_containers: list[tuple[dict, str]] = []
|
||||
for container in (msg_obj, nested_response):
|
||||
if container is None:
|
||||
if not isinstance(container, dict):
|
||||
continue
|
||||
if "input" in container:
|
||||
text_containers.append((container, "input"))
|
||||
|
|
@ -1535,13 +1602,14 @@ class ResponsesWebSocketStreaming:
|
|||
modified = True
|
||||
elif isinstance(value, list):
|
||||
for block in value:
|
||||
if (
|
||||
isinstance(block, dict)
|
||||
and block.get("type") in RESPONSES_WS_MASKABLE_TEXT_BLOCK_TYPES
|
||||
and isinstance(block.get("text"), str)
|
||||
):
|
||||
if not isinstance(block, dict):
|
||||
continue
|
||||
if block.get("type") not in RESPONSES_WS_MASKABLE_TEXT_BLOCK_TYPES:
|
||||
continue
|
||||
block_text = block.get("text")
|
||||
if isinstance(block_text, str):
|
||||
block["text"] = await cb.check_pii(
|
||||
text=block["text"],
|
||||
text=block_text,
|
||||
output_parse_pii=True,
|
||||
presidio_config=presidio_config,
|
||||
request_data=self.request_data,
|
||||
|
|
@ -1592,13 +1660,14 @@ class ResponsesWebSocketStreaming:
|
|||
if not self.guardrail_callbacks:
|
||||
return response_str
|
||||
|
||||
pii_tokens: Dict[str, str] = (self.request_data.get("metadata") or {}).get("pii_tokens", {})
|
||||
metadata = self.request_data.get("metadata")
|
||||
raw_pii_tokens = metadata.get("pii_tokens") if isinstance(metadata, dict) else None
|
||||
pii_tokens: Dict[str, str] = raw_pii_tokens if isinstance(raw_pii_tokens, dict) else {}
|
||||
if not pii_tokens:
|
||||
return response_str
|
||||
|
||||
try:
|
||||
evt_obj = json.loads(response_str)
|
||||
except (json.JSONDecodeError, TypeError):
|
||||
evt_obj = _parse_json_object(response_str)
|
||||
if evt_obj is None:
|
||||
return response_str
|
||||
|
||||
cb = self.guardrail_callbacks[0]
|
||||
|
|
@ -1606,21 +1675,18 @@ class ResponsesWebSocketStreaming:
|
|||
|
||||
if event_type == "response.completed":
|
||||
modified = False
|
||||
response_obj = evt_obj.get("response") or {}
|
||||
response_obj = evt_obj.get("response")
|
||||
if not isinstance(response_obj, dict):
|
||||
return response_str
|
||||
for output_item in response_obj.get("output") or []:
|
||||
if not isinstance(output_item, dict):
|
||||
continue
|
||||
content = output_item.get("content") or []
|
||||
if not isinstance(content, list):
|
||||
continue
|
||||
for content_block in content:
|
||||
for content_block in output_item.get("content") or []:
|
||||
if not isinstance(content_block, dict):
|
||||
continue
|
||||
text = content_block.get("text")
|
||||
if isinstance(text, str):
|
||||
unmasked = cb._unmask_pii_text(text, pii_tokens)
|
||||
unmasked = _call_unmask_pii_text(cb, text, pii_tokens)
|
||||
if unmasked != text:
|
||||
content_block["text"] = unmasked
|
||||
modified = True
|
||||
|
|
@ -1629,7 +1695,7 @@ class ResponsesWebSocketStreaming:
|
|||
if event_type in self._DELTA_EVENT_TYPES:
|
||||
delta = evt_obj.get("delta")
|
||||
if isinstance(delta, str):
|
||||
unmasked = cb._unmask_pii_text(delta, pii_tokens)
|
||||
unmasked = _call_unmask_pii_text(cb, delta, pii_tokens)
|
||||
if unmasked != delta:
|
||||
evt_obj["delta"] = unmasked
|
||||
return json.dumps(evt_obj)
|
||||
|
|
@ -1651,9 +1717,8 @@ class ResponsesWebSocketStreaming:
|
|||
if not self.output_guardrail_callbacks:
|
||||
return response_str
|
||||
|
||||
try:
|
||||
evt_obj = json.loads(response_str)
|
||||
except (json.JSONDecodeError, TypeError):
|
||||
evt_obj = _parse_json_object(response_str)
|
||||
if evt_obj is None:
|
||||
return response_str
|
||||
|
||||
if evt_obj.get("type") != "response.completed":
|
||||
|
|
@ -1662,7 +1727,7 @@ class ResponsesWebSocketStreaming:
|
|||
modified = False
|
||||
for cb in self.output_guardrail_callbacks:
|
||||
presidio_config = cb.get_presidio_settings_from_request_data(self.request_data)
|
||||
response_obj = evt_obj.get("response") or {}
|
||||
response_obj = evt_obj.get("response")
|
||||
if not isinstance(response_obj, dict):
|
||||
continue
|
||||
for output_item in response_obj.get("output") or []:
|
||||
|
|
@ -1679,26 +1744,21 @@ class ResponsesWebSocketStreaming:
|
|||
if masked_args != arguments:
|
||||
output_item["arguments"] = masked_args
|
||||
modified = True
|
||||
summary = output_item.get("summary") or []
|
||||
if isinstance(summary, list):
|
||||
for summary_block in summary:
|
||||
if not isinstance(summary_block, dict):
|
||||
continue
|
||||
summary_text = summary_block.get("text")
|
||||
if isinstance(summary_text, str):
|
||||
masked_summary = await cb.check_pii(
|
||||
text=summary_text,
|
||||
output_parse_pii=False,
|
||||
presidio_config=presidio_config,
|
||||
request_data=self.request_data,
|
||||
)
|
||||
if masked_summary != summary_text:
|
||||
summary_block["text"] = masked_summary
|
||||
modified = True
|
||||
content = output_item.get("content") or []
|
||||
if not isinstance(content, list):
|
||||
continue
|
||||
for content_block in content:
|
||||
for summary_block in output_item.get("summary") or []:
|
||||
if not isinstance(summary_block, dict):
|
||||
continue
|
||||
summary_text = summary_block.get("text")
|
||||
if isinstance(summary_text, str):
|
||||
masked_summary = await cb.check_pii(
|
||||
text=summary_text,
|
||||
output_parse_pii=False,
|
||||
presidio_config=presidio_config,
|
||||
request_data=self.request_data,
|
||||
)
|
||||
if masked_summary != summary_text:
|
||||
summary_block["text"] = masked_summary
|
||||
modified = True
|
||||
for content_block in output_item.get("content") or []:
|
||||
if not isinstance(content_block, dict):
|
||||
continue
|
||||
text = content_block.get("text")
|
||||
|
|
@ -1722,14 +1782,14 @@ class ResponsesWebSocketStreaming:
|
|||
masked_first = await self._mask_response_create(self.first_message)
|
||||
self._store_input(masked_first)
|
||||
self._store_event(masked_first)
|
||||
await self.backend_ws.send(masked_first) # type: ignore[union-attr]
|
||||
await self.backend_ws.send(masked_first)
|
||||
|
||||
while True:
|
||||
message = await self.websocket.receive_text()
|
||||
masked = await self._mask_response_create(message)
|
||||
self._store_input(masked)
|
||||
self._store_event(masked)
|
||||
await self.backend_ws.send(masked) # type: ignore[union-attr]
|
||||
await self.backend_ws.send(masked)
|
||||
|
||||
except Exception as e:
|
||||
verbose_logger.debug("Responses WS client_to_backend ended: %s", e)
|
||||
|
|
@ -1795,10 +1855,10 @@ class ManagedResponsesWebSocketHandler:
|
|||
|
||||
def __init__(
|
||||
self,
|
||||
websocket: Any,
|
||||
websocket: "WebSocket",
|
||||
model: str,
|
||||
logging_obj: "LiteLLMLoggingObj",
|
||||
user_api_key_dict: Optional[Any] = None,
|
||||
user_api_key_dict: Optional["UserAPIKeyAuth"] = None,
|
||||
litellm_metadata: Optional[Dict[str, Any]] = None,
|
||||
api_key: Optional[str] = None,
|
||||
api_base: Optional[str] = None,
|
||||
|
|
@ -1834,13 +1894,11 @@ class ManagedResponsesWebSocketHandler:
|
|||
# ------------------------------------------------------------------
|
||||
|
||||
@staticmethod
|
||||
def _serialize_chunk(chunk: Any) -> Optional[str]:
|
||||
def _serialize_chunk(chunk: object) -> Optional[str]:
|
||||
"""Serialize a streaming chunk to a JSON string for WebSocket transmission."""
|
||||
try:
|
||||
if hasattr(chunk, "model_dump_json"):
|
||||
if isinstance(chunk, BaseModel):
|
||||
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))
|
||||
|
|
@ -1877,41 +1935,41 @@ class ManagedResponsesWebSocketHandler:
|
|||
self._session_history[response_id] = messages
|
||||
|
||||
@staticmethod
|
||||
def _extract_response_id(completed_event: Dict[str, Any]) -> Optional[str]:
|
||||
def _extract_response_id(completed_event: Dict[str, object]) -> 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.
|
||||
"""
|
||||
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:
|
||||
resp_obj = _as_str_object_dict(completed_event.get("response"))
|
||||
encoded_id = resp_obj.get("id") if resp_obj is not None else None
|
||||
if not isinstance(encoded_id, str) or 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],
|
||||
completed_event: Dict[str, object],
|
||||
) -> List[Dict[str, Any]]:
|
||||
"""
|
||||
Convert the output items in a ``response.completed`` event into
|
||||
Responses API message dicts suitable for the next turn's ``input``.
|
||||
"""
|
||||
resp_obj = completed_event.get("response", {})
|
||||
if not isinstance(resp_obj, dict):
|
||||
resp_obj = _as_str_object_dict(completed_event.get("response"))
|
||||
if resp_obj is None:
|
||||
return []
|
||||
messages: List[Dict[str, Any]] = []
|
||||
for item in resp_obj.get("output", []) or []:
|
||||
if not isinstance(item, dict):
|
||||
for raw_item in _as_object_list(resp_obj.get("output")):
|
||||
item = _as_str_object_dict(raw_item)
|
||||
if item is None:
|
||||
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")
|
||||
str(p.get("text") or "")
|
||||
for p in (_as_str_object_dict(raw_p) for raw_p in _as_object_list(item.get("content")))
|
||||
if p is not None and p.get("type") in ("output_text", "text")
|
||||
]
|
||||
text = "".join(text_parts)
|
||||
if text:
|
||||
|
|
@ -1927,7 +1985,7 @@ class ManagedResponsesWebSocketHandler:
|
|||
return messages
|
||||
|
||||
@staticmethod
|
||||
def _input_to_messages(input_val: Any) -> List[Dict[str, Any]]:
|
||||
def _input_to_messages(input_val: object) -> List[Dict[str, object]]:
|
||||
"""
|
||||
Normalise the ``input`` field of a ``response.create`` event to a list
|
||||
of Responses API message dicts.
|
||||
|
|
@ -1940,31 +1998,30 @@ class ManagedResponsesWebSocketHandler:
|
|||
"content": [{"type": "input_text", "text": input_val}],
|
||||
}
|
||||
]
|
||||
if isinstance(input_val, list):
|
||||
return [item for item in input_val if isinstance(item, dict)]
|
||||
return []
|
||||
return [item for item in (_as_str_object_dict(v) for v in _as_object_list(input_val)) if item is not None]
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# _process_response_create sub-methods
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
async def _parse_message(self, raw_message: str) -> Optional[Dict[str, Any]]:
|
||||
async def _parse_message(self, raw_message: str) -> Optional[Dict[str, object]]:
|
||||
"""Parse raw WS text; return the message dict or None (JSON error / ignored type)."""
|
||||
try:
|
||||
msg_obj = json.loads(raw_message)
|
||||
parsed = 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":
|
||||
msg_obj = _as_str_object_dict(parsed)
|
||||
if msg_obj is None or msg_obj.get("type") != "response.create":
|
||||
# Silently ignore non-response.create messages (e.g. warmup pings)
|
||||
return None
|
||||
return msg_obj
|
||||
|
||||
@staticmethod
|
||||
def _is_warmup_frame(msg_obj: Dict[str, Any]) -> bool:
|
||||
def _is_warmup_frame(msg_obj: Dict[str, object]) -> bool:
|
||||
"""Return True for a response.create whose generate flag is false."""
|
||||
nested = msg_obj.get("response")
|
||||
source = nested if isinstance(nested, dict) and nested else msg_obj
|
||||
nested = _as_str_object_dict(msg_obj.get("response"))
|
||||
source = nested if nested else msg_obj
|
||||
return source.get("generate") is False
|
||||
|
||||
@staticmethod
|
||||
|
|
@ -1977,13 +2034,13 @@ class ManagedResponsesWebSocketHandler:
|
|||
return str(raw_id).startswith(_WARMUP_RESPONSE_ID_PREFIX)
|
||||
|
||||
@staticmethod
|
||||
def _warmup_source_params(msg_obj: Dict[str, Any]) -> Dict[str, Any]:
|
||||
nested = msg_obj.get("response")
|
||||
if isinstance(nested, dict) and nested:
|
||||
def _warmup_source_params(msg_obj: Dict[str, object]) -> Dict[str, object]:
|
||||
nested = _as_str_object_dict(msg_obj.get("response"))
|
||||
if nested:
|
||||
return nested
|
||||
return {k: v for k, v in msg_obj.items() if k != "type"}
|
||||
|
||||
def _build_warmup_response(self, msg_obj: Dict[str, Any]) -> Dict[str, Any]:
|
||||
def _build_warmup_response(self, msg_obj: Dict[str, object]) -> Dict[str, Any]:
|
||||
"""Build a minimal completed Responses API object for a warmup ack."""
|
||||
source = self._warmup_source_params(msg_obj)
|
||||
wire_model = source.get("model") or self.model_group or self.model
|
||||
|
|
@ -2001,7 +2058,7 @@ class ManagedResponsesWebSocketHandler:
|
|||
},
|
||||
}
|
||||
|
||||
async def _send_warmup_ack(self, msg_obj: Dict[str, Any]) -> None:
|
||||
async def _send_warmup_ack(self, msg_obj: Dict[str, object]) -> None:
|
||||
"""
|
||||
Acknowledge a generate=false prewarm without calling the provider.
|
||||
|
||||
|
|
@ -2024,16 +2081,14 @@ class ManagedResponsesWebSocketHandler:
|
|||
await self.websocket.send_text(serialized)
|
||||
|
||||
@staticmethod
|
||||
def _build_base_call_kwargs(msg_obj: Dict[str, Any]) -> Dict[str, Any]:
|
||||
def _build_base_call_kwargs(msg_obj: Dict[str, object]) -> 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"}
|
||||
)
|
||||
nested = _as_str_object_dict(msg_obj.get("response"))
|
||||
response_params: Dict[str, object] = nested if 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
|
||||
|
|
@ -2133,7 +2188,7 @@ class ManagedResponsesWebSocketHandler:
|
|||
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]]:
|
||||
async def _stream_and_forward(self, model: str, call_kwargs: Dict[str, Any]) -> Optional[Dict[str, object]]:
|
||||
"""
|
||||
Stream ``litellm.aresponses`` and forward every chunk over the WebSocket.
|
||||
|
||||
|
|
@ -2141,7 +2196,7 @@ class ManagedResponsesWebSocketHandler:
|
|||
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
|
||||
completed_event: Optional[Dict[str, object]] = None
|
||||
stream_response = await litellm.aresponses(model=model, **call_kwargs)
|
||||
async for chunk in stream_response: # type: ignore[union-attr]
|
||||
if chunk is None:
|
||||
|
|
@ -2153,7 +2208,7 @@ class ManagedResponsesWebSocketHandler:
|
|||
continue
|
||||
if chunk_type == "response.completed" and completed_event is None:
|
||||
try:
|
||||
completed_event = json.loads(serialized)
|
||||
completed_event = _as_str_object_dict(json.loads(serialized))
|
||||
except Exception:
|
||||
pass
|
||||
try:
|
||||
|
|
@ -2165,7 +2220,7 @@ class ManagedResponsesWebSocketHandler:
|
|||
|
||||
def _save_turn_history(
|
||||
self,
|
||||
completed_event: Optional[Dict[str, Any]],
|
||||
completed_event: Optional[Dict[str, object]],
|
||||
prior_history: List[Dict[str, Any]],
|
||||
current_messages: List[Dict[str, Any]],
|
||||
) -> None:
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue