mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-12 23:01:41 +00:00
chore(typing): clear basedpyright Any errors in responses streaming and websearch interception
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
f6b9518ddb
commit
23bd48f357
2 changed files with 151 additions and 70 deletions
|
|
@ -10,7 +10,10 @@ import asyncio
|
|||
import math
|
||||
import uuid
|
||||
from collections.abc import AsyncIterator, Mapping, Sequence
|
||||
from typing import TYPE_CHECKING, Any, Final, cast
|
||||
from typing import TYPE_CHECKING, Any, Final, TypeVar, cast
|
||||
|
||||
from pydantic import TypeAdapter, ValidationError
|
||||
from typing_extensions import TypeIs
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_logger
|
||||
|
|
@ -72,6 +75,43 @@ WEBSEARCH_EMIT_NATIVE_BLOCKS_KEY: Final = "_websearch_interception_emit_native_b
|
|||
# ``web_search_tool_result`` blocks to inject into the final response.
|
||||
WEBSEARCH_NATIVE_BLOCKS_METADATA_KEY: Final = "websearch_native_blocks"
|
||||
|
||||
_ResponseT = TypeVar("_ResponseT")
|
||||
|
||||
_CONTENT_ATTR: Final = "content"
|
||||
|
||||
_OBJECT_ITEMS_ADAPTER: Final[TypeAdapter[tuple[object, ...]]] = TypeAdapter(tuple[object, ...])
|
||||
|
||||
_NATIVE_BLOCKS_ADAPTER: Final[TypeAdapter[tuple[Mapping[str, object], ...]]] = TypeAdapter(
|
||||
tuple[Mapping[str, object], ...]
|
||||
)
|
||||
|
||||
_WEBSEARCH_CONFIG_ADAPTER: Final[TypeAdapter[WebSearchInterceptionConfig]] = TypeAdapter(WebSearchInterceptionConfig)
|
||||
|
||||
|
||||
def _is_json_object(value: object) -> TypeIs[dict[str, object]]: # guard-ok: trivial isinstance; JSON keys are str
|
||||
return isinstance(value, dict)
|
||||
|
||||
|
||||
def _parse_object_items(value: object) -> tuple[object, ...]:
|
||||
try:
|
||||
return _OBJECT_ITEMS_ADAPTER.validate_python(value)
|
||||
except ValidationError:
|
||||
return ()
|
||||
|
||||
|
||||
def _parse_native_blocks(value: object) -> tuple[Mapping[str, object], ...]:
|
||||
try:
|
||||
return _NATIVE_BLOCKS_ADAPTER.validate_python(value)
|
||||
except ValidationError:
|
||||
return ()
|
||||
|
||||
|
||||
def _parse_websearch_config(value: object) -> WebSearchInterceptionConfig:
|
||||
try:
|
||||
return _WEBSEARCH_CONFIG_ADAPTER.validate_python(value)
|
||||
except ValidationError:
|
||||
return {}
|
||||
|
||||
|
||||
class WebSearchInterceptionLogger(CustomLogger):
|
||||
"""
|
||||
|
|
@ -833,7 +873,7 @@ class WebSearchInterceptionLogger(CustomLogger):
|
|||
Anthropic-native clients (Claude Desktop, the Anthropic SDK) can
|
||||
render citations / sources alongside the model's textual reply.
|
||||
"""
|
||||
native_blocks: Final = plan.metadata.get(WEBSEARCH_NATIVE_BLOCKS_METADATA_KEY)
|
||||
native_blocks: Final = _parse_native_blocks(plan.metadata.get(WEBSEARCH_NATIVE_BLOCKS_METADATA_KEY))
|
||||
if not native_blocks:
|
||||
return response
|
||||
return self._inject_native_blocks(response, native_blocks)
|
||||
|
|
@ -883,17 +923,20 @@ class WebSearchInterceptionLogger(CustomLogger):
|
|||
)
|
||||
|
||||
@staticmethod
|
||||
def _inject_native_blocks(response: Any, native_blocks: Sequence[Mapping[str, object]]) -> Any:
|
||||
def _inject_native_blocks(
|
||||
response: _ResponseT,
|
||||
native_blocks: Sequence[Mapping[str, object]],
|
||||
) -> _ResponseT:
|
||||
"""Prepend native blocks to response content, dict or object form."""
|
||||
if not native_blocks:
|
||||
return response
|
||||
if isinstance(response, dict):
|
||||
existing = response.get("content") or []
|
||||
response["content"] = list(native_blocks) + list(existing)
|
||||
if _is_json_object(response):
|
||||
existing_items: Final = _parse_object_items(response.get(_CONTENT_ATTR))
|
||||
response[_CONTENT_ATTR] = [*native_blocks, *existing_items]
|
||||
return response
|
||||
existing = getattr(response, "content", None) or []
|
||||
existing_attr_items: Final = _parse_object_items(getattr(response, _CONTENT_ATTR, None))
|
||||
try:
|
||||
response.content = list(native_blocks) + list(existing)
|
||||
setattr(response, _CONTENT_ATTR, [*native_blocks, *existing_attr_items])
|
||||
except (AttributeError, TypeError):
|
||||
# Object refused write — fall through and leave the response
|
||||
# untouched rather than crash the request.
|
||||
|
|
@ -1675,8 +1718,8 @@ class WebSearchInterceptionLogger(CustomLogger):
|
|||
|
||||
@staticmethod
|
||||
def initialize_from_proxy_config(
|
||||
litellm_settings: dict[str, Any],
|
||||
callback_specific_params: dict[str, Any],
|
||||
litellm_settings: Mapping[str, object],
|
||||
callback_specific_params: Mapping[str, object],
|
||||
) -> "WebSearchInterceptionLogger":
|
||||
"""
|
||||
Static method to initialize WebSearchInterceptionLogger from proxy config.
|
||||
|
|
@ -1698,16 +1741,11 @@ class WebSearchInterceptionLogger(CustomLogger):
|
|||
)
|
||||
"""
|
||||
# Get websearch_interception_params from litellm_settings or callback_specific_params
|
||||
websearch_params: WebSearchInterceptionConfig = {}
|
||||
if "websearch_interception_params" in litellm_settings:
|
||||
websearch_params = litellm_settings["websearch_interception_params"]
|
||||
elif "websearch_interception" in callback_specific_params and isinstance(
|
||||
callback_specific_params["websearch_interception"], dict
|
||||
):
|
||||
websearch_params = cast(
|
||||
WebSearchInterceptionConfig,
|
||||
callback_specific_params["websearch_interception"],
|
||||
)
|
||||
raw_params: Final = (
|
||||
litellm_settings["websearch_interception_params"]
|
||||
if "websearch_interception_params" in litellm_settings
|
||||
else callback_specific_params.get("websearch_interception")
|
||||
)
|
||||
|
||||
# Use classmethod to initialize from config
|
||||
return WebSearchInterceptionLogger.from_config_yaml(websearch_params)
|
||||
return WebSearchInterceptionLogger.from_config_yaml(_parse_websearch_config(raw_params))
|
||||
|
|
|
|||
|
|
@ -9,10 +9,11 @@ from collections.abc import Awaitable, Callable, Mapping
|
|||
from datetime import datetime
|
||||
from functools import lru_cache
|
||||
from types import MappingProxyType
|
||||
from typing import TYPE_CHECKING, Any, Final, Literal
|
||||
from typing import TYPE_CHECKING, Any, Final, Literal, Protocol, runtime_checkable
|
||||
|
||||
import httpx
|
||||
from openai._streaming import SSEDecoder
|
||||
from pydantic import TypeAdapter, ValidationError
|
||||
from typing_extensions import TypeIs
|
||||
|
||||
import litellm
|
||||
|
|
@ -33,6 +34,7 @@ from litellm.llms.base_llm.responses.transformation import BaseResponsesAPIConfi
|
|||
from litellm.responses.utils import ResponseAPILoggingUtils, ResponsesAPIRequestUtils
|
||||
from litellm.types.llms.openai import (
|
||||
PART_UNION_TYPES,
|
||||
OpenAIChatCompletionLogprobsContent,
|
||||
ResponseAPIUsage,
|
||||
ResponsesAPIResponse,
|
||||
ResponsesAPIStreamEvents,
|
||||
|
|
@ -57,6 +59,21 @@ def _get_openai_response_types():
|
|||
return openai_types
|
||||
|
||||
|
||||
@runtime_checkable
|
||||
class _SupportsModelDump(Protocol):
|
||||
def model_dump(self) -> Mapping[str, object]: ...
|
||||
|
||||
|
||||
@runtime_checkable
|
||||
class _SupportsModelDumpJson(Protocol):
|
||||
def model_dump_json(self, *, exclude_none: bool = ...) -> str: ...
|
||||
|
||||
|
||||
@runtime_checkable
|
||||
class _SupportsModelDumpExcludeNone(Protocol):
|
||||
def model_dump(self, *, exclude_none: bool = ...) -> Mapping[str, object]: ...
|
||||
|
||||
|
||||
def _is_json_object(value: object) -> TypeIs[dict[str, object]]: # guard-ok: trivial isinstance; JSON keys are str
|
||||
return isinstance(value, dict)
|
||||
|
||||
|
|
@ -69,7 +86,30 @@ def _is_str_mapping(value: object) -> TypeIs[dict[str, str]]: # guard-ok: verif
|
|||
return _is_json_object(value) and all(isinstance(item, str) for item in value.values())
|
||||
|
||||
|
||||
def _model_id_from_metadata(litellm_metadata: dict[str, object] | None) -> str | None:
|
||||
def _json_array_items(value: object) -> tuple[object, ...]:
|
||||
return tuple(value) if _is_json_array(value) else ()
|
||||
|
||||
|
||||
def _json_object_items(value: object) -> tuple[dict[str, object], ...]:
|
||||
return tuple(item for item in _json_array_items(value) if _is_json_object(item))
|
||||
|
||||
|
||||
_LOGPROBS_ADAPTER: Final[TypeAdapter[list[OpenAIChatCompletionLogprobsContent] | None]] = TypeAdapter(
|
||||
list[OpenAIChatCompletionLogprobsContent] | None
|
||||
)
|
||||
|
||||
_OBJECT_ADAPTER: Final[TypeAdapter[object]] = TypeAdapter(object)
|
||||
|
||||
_JSON_OBJECT_ADAPTER: Final[TypeAdapter[dict[str, object]]] = TypeAdapter(dict[str, object])
|
||||
|
||||
_UNMASK_PII_TEXT_ATTR: Final = "_unmask_pii_text"
|
||||
|
||||
_UnmaskPiiText = Callable[[str, dict[str, str]], str]
|
||||
|
||||
_UNMASK_PII_TEXT_ADAPTER: Final[TypeAdapter[_UnmaskPiiText]] = TypeAdapter(_UnmaskPiiText)
|
||||
|
||||
|
||||
def _model_id_from_metadata(litellm_metadata: Mapping[str, object] | None) -> str | None:
|
||||
model_info: Final = litellm_metadata.get("model_info") if litellm_metadata else None
|
||||
model_id: Final = model_info.get("id") if _is_json_object(model_info) else None
|
||||
return model_id if isinstance(model_id, str) else None
|
||||
|
|
@ -228,10 +268,10 @@ class BaseResponsesAPIStreamingIterator:
|
|||
|
||||
try:
|
||||
# Parse the JSON chunk
|
||||
parsed_chunk: Final = json.loads(chunk)
|
||||
parsed_chunk: Final = _OBJECT_ADAPTER.validate_python(json.loads(chunk))
|
||||
|
||||
# Format as ResponsesAPIStreamingResponse
|
||||
if isinstance(parsed_chunk, dict):
|
||||
if _is_json_object(parsed_chunk):
|
||||
if self.responses_api_provider_config is None:
|
||||
raise ValueError("responses_api_provider_config is required to process live streaming chunks")
|
||||
openai_responses_api_chunk: Final = self.responses_api_provider_config.transform_streaming_response(
|
||||
|
|
@ -1029,8 +1069,8 @@ class CachedResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator):
|
|||
return evt
|
||||
|
||||
|
||||
def _dump_response_object(obj: Any) -> dict[str, Any]:
|
||||
if hasattr(obj, "model_dump"):
|
||||
def _dump_response_object(obj: object) -> Mapping[str, object]:
|
||||
if isinstance(obj, _SupportsModelDump):
|
||||
return obj.model_dump()
|
||||
if _is_json_object(obj):
|
||||
return obj
|
||||
|
|
@ -1059,7 +1099,7 @@ def _build_content_part_done_event(
|
|||
item_id: str,
|
||||
output_index: int,
|
||||
content_index: int,
|
||||
part_payload: dict[str, Any],
|
||||
part_payload: Mapping[str, object],
|
||||
) -> ResponsesAPIStreamingResponse | None:
|
||||
openai_types: Final = _get_openai_response_types()
|
||||
part_type: Final = part_payload.get("type")
|
||||
|
|
@ -1067,13 +1107,13 @@ def _build_content_part_done_event(
|
|||
if part_type == "output_text":
|
||||
annotations: Final = [
|
||||
openai_types.BaseLiteLLMOpenAIResponseObject(**annotation)
|
||||
for annotation in part_payload.get("annotations", []) or []
|
||||
for annotation in _json_object_items(part_payload.get("annotations"))
|
||||
]
|
||||
part = openai_types.ContentPartDonePartOutputText(
|
||||
type="output_text",
|
||||
text=str(part_payload.get("text") or ""),
|
||||
annotations=annotations,
|
||||
logprobs=part_payload.get("logprobs"),
|
||||
logprobs=_LOGPROBS_ADAPTER.validate_python(part_payload.get("logprobs")),
|
||||
)
|
||||
elif part_type == "refusal":
|
||||
part = openai_types.ContentPartDonePartRefusal(
|
||||
|
|
@ -1103,7 +1143,7 @@ def _add_text_like_part_events(
|
|||
item_id: str,
|
||||
output_index: int,
|
||||
content_index: int,
|
||||
part_payload: dict[str, Any],
|
||||
part_payload: Mapping[str, object],
|
||||
chunk_size: int,
|
||||
) -> None:
|
||||
openai_types: Final = _get_openai_response_types()
|
||||
|
|
@ -1120,7 +1160,7 @@ def _add_text_like_part_events(
|
|||
delta=text[i : i + chunk_size],
|
||||
)
|
||||
)
|
||||
for annotation_index, annotation in enumerate(part_payload.get("annotations", []) or []):
|
||||
for annotation_index, annotation in enumerate(_json_object_items(part_payload.get("annotations"))):
|
||||
events.append(
|
||||
openai_types.OutputTextAnnotationAddedEvent(
|
||||
type=openai_types.ResponsesAPIStreamEvents.OUTPUT_TEXT_ANNOTATION_ADDED,
|
||||
|
|
@ -1186,7 +1226,7 @@ def _build_synthetic_response_events(
|
|||
]
|
||||
|
||||
sequence_number = 0
|
||||
for output_index, output_item in enumerate(getattr(transformed, "output", []) or []):
|
||||
for output_index, output_item in enumerate(_json_array_items(getattr(transformed, "output", None))):
|
||||
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")
|
||||
|
|
@ -1200,7 +1240,7 @@ def _build_synthetic_response_events(
|
|||
)
|
||||
|
||||
if item_type == "message":
|
||||
for content_index, part in enumerate(output_item_payload.get("content", []) or []):
|
||||
for content_index, part in enumerate(_json_array_items(output_item_payload.get("content"))):
|
||||
part_payload = _dump_response_object(part)
|
||||
events.append(
|
||||
openai_types.ContentPartAddedEvent(
|
||||
|
|
@ -1247,7 +1287,7 @@ def _build_synthetic_response_events(
|
|||
)
|
||||
)
|
||||
elif item_type == "reasoning":
|
||||
for summary_index, summary in enumerate(output_item_payload.get("summary", []) or []):
|
||||
for summary_index, summary in enumerate(_json_array_items(output_item_payload.get("summary"))):
|
||||
summary_payload = _dump_response_object(summary)
|
||||
summary_text = str(summary_payload.get("text") or "")
|
||||
for i in range(0, len(summary_text), chunk_size):
|
||||
|
|
@ -1340,7 +1380,7 @@ class ResponsesWebSocketStreaming:
|
|||
user_api_key_dict: UserAPIKeyAuth | None = None,
|
||||
request_data: dict[str, object] | None = None,
|
||||
first_message: str | None = None,
|
||||
guardrail_callbacks: list[Any] | None = None,
|
||||
guardrail_callbacks: list[PresidioGuardrailCallback] | None = None,
|
||||
output_guardrail_callbacks: list[PresidioGuardrailCallback] | None = None,
|
||||
authorized_model: str | None = None,
|
||||
):
|
||||
|
|
@ -1352,7 +1392,7 @@ class ResponsesWebSocketStreaming:
|
|||
self.messages: list[dict[str, object]] = []
|
||||
self.input_messages: list[dict[str, object]] = []
|
||||
self.first_message = first_message
|
||||
self.guardrail_callbacks: list[Any] = guardrail_callbacks or []
|
||||
self.guardrail_callbacks: list[PresidioGuardrailCallback] = guardrail_callbacks or []
|
||||
self.output_guardrail_callbacks: list[PresidioGuardrailCallback] = output_guardrail_callbacks or []
|
||||
# Model name authorized at connection time; enforced on every
|
||||
# response.create frame to prevent deployment-substitution attacks.
|
||||
|
|
@ -1362,24 +1402,22 @@ class ResponsesWebSocketStreaming:
|
|||
return event_obj.get("type") in RESPONSES_WS_LOGGED_EVENT_TYPES
|
||||
|
||||
def _store_event(self, event: str | bytes | dict[str, object]) -> None:
|
||||
if isinstance(event, bytes):
|
||||
event = event.decode("utf-8")
|
||||
if isinstance(event, str):
|
||||
if isinstance(event, str | bytes):
|
||||
try:
|
||||
event_obj = json.loads(event)
|
||||
except (json.JSONDecodeError, TypeError):
|
||||
parsed_event: dict[str, object] = _JSON_OBJECT_ADAPTER.validate_json(event)
|
||||
except ValidationError:
|
||||
return
|
||||
else:
|
||||
event_obj = event
|
||||
parsed_event = event
|
||||
|
||||
if self._should_store_event(event_obj):
|
||||
self.messages.append(event_obj)
|
||||
if self._should_store_event(parsed_event):
|
||||
self.messages.append(parsed_event)
|
||||
|
||||
def _collect_input_from_client_event(self, message: object) -> None:
|
||||
"""Extract user input content from response.create for logging."""
|
||||
try:
|
||||
if isinstance(message, str):
|
||||
msg_obj = json.loads(message)
|
||||
msg_obj: dict[str, object] = _JSON_OBJECT_ADAPTER.validate_json(message)
|
||||
elif _is_json_object(message):
|
||||
msg_obj = message
|
||||
else:
|
||||
|
|
@ -1407,7 +1445,7 @@ class ResponsesWebSocketStreaming:
|
|||
text = c.get("text", "")
|
||||
if text:
|
||||
self.input_messages.append({"role": "user", "content": text})
|
||||
except (json.JSONDecodeError, AttributeError, TypeError):
|
||||
except (ValidationError, AttributeError, TypeError):
|
||||
pass
|
||||
|
||||
def _store_input(self, message: object) -> None:
|
||||
|
|
@ -1449,8 +1487,8 @@ class ResponsesWebSocketStreaming:
|
|||
# masked response.completed.
|
||||
if self.output_guardrail_callbacks:
|
||||
try:
|
||||
_evt_type = json.loads(response_str).get("type")
|
||||
except (json.JSONDecodeError, TypeError):
|
||||
_evt_type = _JSON_OBJECT_ADAPTER.validate_json(response_str).get("type")
|
||||
except ValidationError:
|
||||
_evt_type = None
|
||||
if _evt_type in self._DELTA_EVENT_TYPES or _evt_type in self._OUTPUT_DONE_EVENT_TYPES:
|
||||
continue
|
||||
|
|
@ -1513,8 +1551,8 @@ class ResponsesWebSocketStreaming:
|
|||
Non-``response.create`` messages are returned unchanged.
|
||||
"""
|
||||
try:
|
||||
msg_obj: Final = json.loads(message)
|
||||
except (json.JSONDecodeError, TypeError):
|
||||
msg_obj: Final = _JSON_OBJECT_ADAPTER.validate_json(message)
|
||||
except ValidationError:
|
||||
return message
|
||||
|
||||
if msg_obj.get("type") != "response.create":
|
||||
|
|
@ -1641,11 +1679,12 @@ class ResponsesWebSocketStreaming:
|
|||
return response_str
|
||||
|
||||
try:
|
||||
evt_obj: Final = json.loads(response_str)
|
||||
except (json.JSONDecodeError, TypeError):
|
||||
evt_obj: Final = _JSON_OBJECT_ADAPTER.validate_json(response_str)
|
||||
except ValidationError:
|
||||
return response_str
|
||||
|
||||
cb: Final = self.guardrail_callbacks[0]
|
||||
unmask_pii_text: Final = _UNMASK_PII_TEXT_ADAPTER.validate_python(getattr(cb, _UNMASK_PII_TEXT_ATTR))
|
||||
event_type: Final = evt_obj.get("type")
|
||||
|
||||
if event_type == "response.completed":
|
||||
|
|
@ -1665,7 +1704,7 @@ class ResponsesWebSocketStreaming:
|
|||
continue
|
||||
text = content_block.get("text")
|
||||
if isinstance(text, str):
|
||||
unmasked = cb._unmask_pii_text(text, pii_tokens)
|
||||
unmasked = unmask_pii_text(text, pii_tokens)
|
||||
if unmasked != text:
|
||||
content_block["text"] = unmasked
|
||||
modified = True
|
||||
|
|
@ -1674,7 +1713,7 @@ class ResponsesWebSocketStreaming:
|
|||
if event_type in self._DELTA_EVENT_TYPES:
|
||||
delta: Final = evt_obj.get("delta")
|
||||
if isinstance(delta, str):
|
||||
unmasked = cb._unmask_pii_text(delta, pii_tokens)
|
||||
unmasked = unmask_pii_text(delta, pii_tokens)
|
||||
if unmasked != delta:
|
||||
evt_obj["delta"] = unmasked
|
||||
return json.dumps(evt_obj)
|
||||
|
|
@ -1697,8 +1736,8 @@ class ResponsesWebSocketStreaming:
|
|||
return response_str
|
||||
|
||||
try:
|
||||
evt_obj: Final[Mapping[str, object]] = json.loads(response_str)
|
||||
except (json.JSONDecodeError, TypeError):
|
||||
evt_obj: Final[Mapping[str, object]] = _JSON_OBJECT_ADAPTER.validate_json(response_str)
|
||||
except ValidationError:
|
||||
return response_str
|
||||
|
||||
if evt_obj.get("type") != "response.completed":
|
||||
|
|
@ -1845,7 +1884,7 @@ class ManagedResponsesWebSocketHandler:
|
|||
model: str,
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
user_api_key_dict: UserAPIKeyAuth | None = None,
|
||||
litellm_metadata: dict[str, Any] | None = None,
|
||||
litellm_metadata: dict[str, object] | None = None,
|
||||
api_key: str | None = None,
|
||||
api_base: str | None = None,
|
||||
timeout: float | None = None,
|
||||
|
|
@ -1857,10 +1896,11 @@ class ManagedResponsesWebSocketHandler:
|
|||
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.model_group: str | None = self.litellm_metadata.get("model_group") or self.litellm_metadata.get(
|
||||
self.litellm_metadata: dict[str, object] = litellm_metadata or {}
|
||||
raw_model_group: Final = self.litellm_metadata.get("model_group") or self.litellm_metadata.get(
|
||||
"deployment_model_name"
|
||||
)
|
||||
self.model_group: str | None = raw_model_group if isinstance(raw_model_group, str) else None
|
||||
self.api_key = api_key
|
||||
self.api_base = api_base
|
||||
self.timeout = timeout
|
||||
|
|
@ -1880,14 +1920,14 @@ class ManagedResponsesWebSocketHandler:
|
|||
# ------------------------------------------------------------------
|
||||
|
||||
@staticmethod
|
||||
def _serialize_chunk(chunk: Any) -> str | None:
|
||||
def _serialize_chunk(chunk: object) -> str | None:
|
||||
"""Serialize a streaming chunk to a JSON string for WebSocket transmission."""
|
||||
try:
|
||||
if hasattr(chunk, "model_dump_json"):
|
||||
if isinstance(chunk, _SupportsModelDumpJson):
|
||||
return chunk.model_dump_json(exclude_none=True)
|
||||
if hasattr(chunk, "model_dump"):
|
||||
if isinstance(chunk, _SupportsModelDumpExcludeNone):
|
||||
return json.dumps(chunk.model_dump(exclude_none=True), default=str)
|
||||
if isinstance(chunk, dict):
|
||||
if _is_json_object(chunk):
|
||||
return json.dumps(chunk, default=str)
|
||||
return json.dumps(str(chunk))
|
||||
except Exception as exc:
|
||||
|
|
@ -1998,8 +2038,8 @@ class ManagedResponsesWebSocketHandler:
|
|||
async def _parse_message(self, raw_message: str) -> dict[str, object] | None:
|
||||
"""Parse raw WS text; return the message dict or None (JSON error / ignored type)."""
|
||||
try:
|
||||
msg_obj: Final = json.loads(raw_message)
|
||||
except json.JSONDecodeError:
|
||||
msg_obj: Final = _JSON_OBJECT_ADAPTER.validate_json(raw_message)
|
||||
except ValidationError:
|
||||
await self._send_error("Invalid JSON in response.create event", "invalid_request_error")
|
||||
return None
|
||||
if msg_obj.get("type") != "response.create":
|
||||
|
|
@ -2196,14 +2236,17 @@ class ManagedResponsesWebSocketHandler:
|
|||
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)
|
||||
chunk_obj = _OBJECT_ADAPTER.validate_python(chunk)
|
||||
chunk_type = getattr(chunk_obj, "type", None) or (
|
||||
chunk_obj.get("type") if _is_json_object(chunk_obj) else None
|
||||
)
|
||||
serialized = self._serialize_chunk(chunk_obj)
|
||||
if serialized is None:
|
||||
continue
|
||||
if chunk_type == "response.completed" and completed_event is None:
|
||||
try:
|
||||
completed_event = json.loads(serialized)
|
||||
except Exception:
|
||||
completed_event = _JSON_OBJECT_ADAPTER.validate_json(serialized)
|
||||
except ValidationError:
|
||||
pass
|
||||
try:
|
||||
await self.websocket.send_text(serialized)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue