Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/response_polling/background_streaming.py: 21%
199 statements
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-10 12:01 +0000
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-10 12:01 +0000
1"""
2Background Streaming Task for Polling Via Cache Feature
4Handles streaming responses from LLM providers and updates Redis cache
5with partial results for polling.
7Follows OpenAI Response Streaming format:
8https://platform.openai.com/docs/api-reference/responses-streaming
9"""
11import asyncio
12import json
13from collections.abc import Callable, Mapping, Sequence
14from typing import TYPE_CHECKING, Final, TypeAlias
16from fastapi import Request, Response
17from fastapi.responses import StreamingResponse
18from starlette.types import Message
19from typing_extensions import ReadOnly, TypedDict
21from litellm._logging import verbose_proxy_logger
22from litellm.proxy.auth.user_api_key_auth import UserAPIKeyAuth
23from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing
24from litellm.proxy.response_polling.polling_handler import ResponsePollingHandler
25from litellm.types.llms.openai import ResponsesAPIStatus
27if TYPE_CHECKING: 27 ↛ 28line 27 didn't jump to line 28 because the condition on line 27 was never true
28 from litellm.proxy.proxy_server import ProxyConfig
29 from litellm.proxy.utils import ProxyLogging
30 from litellm.router import Router
33_JsonDict: TypeAlias = dict[str, object]
34_JsonList: TypeAlias = list[object]
37class _OutputItem(TypedDict, total=False):
38 id: ReadOnly[str]
39 content: ReadOnly[Sequence[object]]
42class _TerminalResponse(TypedDict, total=False):
43 status: ReadOnly[ResponsesAPIStatus]
44 error: ReadOnly[_JsonDict]
45 usage: ReadOnly[_JsonDict]
46 reasoning: ReadOnly[_JsonDict]
47 tool_choice: ReadOnly[object]
48 tools: ReadOnly[_JsonList]
49 model: ReadOnly[str]
50 instructions: ReadOnly[str]
51 temperature: ReadOnly[float]
52 top_p: ReadOnly[float]
53 max_output_tokens: ReadOnly[int]
54 previous_response_id: ReadOnly[str]
55 text: ReadOnly[_JsonDict]
56 truncation: ReadOnly[str]
57 parallel_tool_calls: ReadOnly[bool]
58 user: ReadOnly[str]
59 store: ReadOnly[bool]
60 incomplete_details: ReadOnly[_JsonDict]
61 output: ReadOnly[Sequence[_OutputItem]]
64class _StreamEvent(TypedDict, total=False):
65 type: ReadOnly[str]
66 item: ReadOnly[_OutputItem]
67 item_id: ReadOnly[str]
68 content_index: ReadOnly[int]
69 delta: ReadOnly[str]
70 part: ReadOnly[object]
71 response: ReadOnly[_TerminalResponse]
74class _StreamEventParser:
75 parse: Callable[[str], _StreamEvent] = staticmethod(json.loads)
78def _sse_frame_data(frame: str) -> str | None:
79 return next((line[6:].strip() for line in frame.splitlines() if line.startswith("data: ")), None)
82async def _never_receive() -> Message:
83 await asyncio.Event().wait()
84 raise AssertionError("unreachable")
87def detach_request_from_client(request: Request) -> Request:
88 """Same scope (headers, parsed body, auth) but a receive() that never yields http.disconnect.
90 The polling client closes its connection right after getting the polling id, so the
91 upstream call must not be cancelled by the client-disconnect guards.
92 """
93 return Request(request.scope, _never_receive)
96async def background_streaming_task(
97 polling_id: str,
98 data: dict[str, object],
99 polling_handler: ResponsePollingHandler,
100 request: Request,
101 fastapi_response: Response,
102 user_api_key_dict: UserAPIKeyAuth,
103 general_settings: dict[str, object],
104 llm_router: "Router | None",
105 proxy_config: "ProxyConfig",
106 proxy_logging_obj: "ProxyLogging",
107 select_data_generator: Callable[..., object] | None,
108 user_model: str | None,
109 user_temperature: float | None,
110 user_request_timeout: float | None,
111 user_max_tokens: int | None,
112 user_api_base: str | None,
113 version: str | None,
114):
115 """
116 Background task to stream response and update cache
118 Follows OpenAI Response Streaming format:
119 https://platform.openai.com/docs/api-reference/responses-streaming
121 Processes streaming events and builds Response object:
122 https://platform.openai.com/docs/api-reference/responses/object
123 """
125 try:
126 verbose_proxy_logger.info("Starting background streaming for %s", polling_id)
128 # Update status to in_progress (OpenAI format)
129 await polling_handler.update_state(
130 polling_id=polling_id,
131 status="in_progress",
132 )
134 # Force streaming mode and remove background flag
135 data["stream"] = True
136 data.pop("background", None)
138 # Create processor
139 processor: Final = ProxyBaseLLMRequestProcessing(data=data)
141 # Make streaming request.
142 # Pre-call checks (rate limits, guardrails, budget) were already run
143 # before polling ID creation, so skip them here to avoid double-counting.
144 response: Final[StreamingResponse] = await processor.base_process_llm_request(
145 request=detach_request_from_client(request),
146 fastapi_response=fastapi_response,
147 user_api_key_dict=user_api_key_dict,
148 route_type="aresponses",
149 proxy_logging_obj=proxy_logging_obj,
150 llm_router=llm_router,
151 general_settings=general_settings,
152 proxy_config=proxy_config,
153 select_data_generator=select_data_generator,
154 model=None,
155 user_model=user_model,
156 user_temperature=user_temperature,
157 user_request_timeout=user_request_timeout,
158 user_max_tokens=user_max_tokens,
159 user_api_base=user_api_base,
160 version=version,
161 skip_pre_call_logic=True,
162 )
164 # Process streaming response following OpenAI events format
165 # https://platform.openai.com/docs/api-reference/responses-streaming
166 output_items: Final = dict[str, _OutputItem]()
167 accumulated_text: Final = dict[tuple[str, int], str]()
169 # ResponsesAPIResponse fields to extract from response.completed
170 usage_data = None
171 reasoning_data = None
172 tool_choice_data = None
173 tools_data = None
174 model_data = None
175 instructions_data = None
176 temperature_data = None
177 top_p_data = None
178 max_output_tokens_data = None
179 previous_response_id_data = None
180 text_data = None
181 truncation_data = None
182 parallel_tool_calls_data = None
183 user_data = None
184 store_data = None
185 incomplete_details_data = None
187 state_dirty = False # Track if state needs to be synced
188 last_update_time = asyncio.get_event_loop().time()
189 UPDATE_INTERVAL: Final = 0.150 # 150ms batching interval
191 # Track the terminal event from the stream (may not be "completed")
192 terminal_status: ResponsesAPIStatus | None = (
193 None # Will be set by response.completed/failed/incomplete/cancelled
194 )
195 terminal_error = None
196 _event_to_status: Final[Mapping[str, ResponsesAPIStatus]] = {
197 "response.completed": "completed",
198 "response.failed": "failed",
199 "response.incomplete": "incomplete",
200 "response.cancelled": "cancelled",
201 }
203 async def flush_state_if_needed(force: bool = False) -> None:
204 """Flush accumulated state to Redis if interval elapsed or forced"""
205 nonlocal state_dirty, last_update_time
207 current_time: Final = asyncio.get_event_loop().time()
208 if state_dirty and (force or (current_time - last_update_time) >= UPDATE_INTERVAL):
209 # Convert output_items dict to list for update
210 output_list: Final = list(output_items.values())
211 await polling_handler.update_state(
212 polling_id=polling_id,
213 output=output_list,
214 )
215 state_dirty = False
216 last_update_time = current_time
218 # Handle StreamingResponse
219 if not hasattr(response, "body_iterator"):
220 verbose_proxy_logger.warning(
221 "background_streaming_task: response for %s has no body_iterator; this may indicate a misconfiguration or provider error",
222 polling_id,
223 )
225 if hasattr(response, "body_iterator"):
226 async for chunk in response.body_iterator:
227 # Parse chunk
228 if isinstance(chunk, bytes):
229 chunk = chunk.decode("utf-8")
231 if isinstance(chunk, str) and (chunk_data := _sse_frame_data(chunk)) is not None:
232 if chunk_data == "[DONE]":
233 break
235 try:
236 event: _StreamEvent = _StreamEventParser.parse(chunk_data)
237 event_type = event.get("type", "")
239 # Process different event types based on OpenAI streaming spec
240 if event_type == "response.output_item.added":
241 # New output item added
242 item = event.get("item", {})
243 item_id = item.get("id")
244 if item_id:
245 output_items[item_id] = item
246 state_dirty = True
248 elif event_type == "response.content_part.added":
249 # Content part added to an output item
250 item_id = event.get("item_id")
251 content_part = event.get("part", {})
253 if item_id and item_id in output_items:
254 # Update the output item with new content
255 added_item = output_items[item_id]
256 output_items[item_id] = {
257 **added_item,
258 "content": (*added_item.get("content", ()), content_part),
259 }
260 state_dirty = True
262 elif event_type == "response.output_text.delta":
263 # Text delta - accumulate text content
264 # https://platform.openai.com/docs/api-reference/responses-streaming/response-text-delta
265 item_id = event.get("item_id")
266 content_index = event.get("content_index", 0)
267 delta = event.get("delta", "")
269 if item_id and item_id in output_items:
270 # Accumulate text delta
271 key = (item_id, content_index)
272 if key not in accumulated_text:
273 accumulated_text[key] = ""
274 accumulated_text[key] += delta
276 # Update the content in output_items
277 delta_item = output_items[item_id]
278 if "content" in delta_item:
279 content_list = delta_item["content"]
280 if content_index < len(content_list):
281 content_entry = content_list[content_index]
282 if isinstance(content_entry, dict):
283 content_entry["text"] = accumulated_text[key]
284 state_dirty = True
286 elif event_type == "response.content_part.done":
287 # Content part completed
288 item_id = event.get("item_id")
289 content_part = event.get("part", {})
290 content_index = event.get("content_index", 0)
292 if item_id and item_id in output_items:
293 # Update with final content from event
294 done_item = output_items[item_id]
295 if "content" in done_item:
296 content_list = done_item["content"]
297 if content_index < len(content_list):
298 output_items[item_id] = {
299 **done_item,
300 "content": tuple(
301 content_part if part_index == content_index else existing_part
302 for part_index, existing_part in enumerate(content_list)
303 ),
304 }
305 state_dirty = True
307 elif event_type == "response.output_item.done":
308 # Output item completed - use final item data
309 item = event.get("item", {})
310 item_id = item.get("id")
311 if item_id:
312 output_items[item_id] = item
313 state_dirty = True
315 elif event_type == "response.in_progress":
316 # Response is now in progress
317 # https://platform.openai.com/docs/api-reference/responses-streaming/response-in-progress
318 await polling_handler.update_state(
319 polling_id=polling_id,
320 status="in_progress",
321 )
323 elif event_type in (
324 "response.completed",
325 "response.failed",
326 "response.incomplete",
327 "response.cancelled",
328 ):
329 # Terminal event - extract all ResponsesAPIResponse fields
330 # https://platform.openai.com/docs/api-reference/responses-streaming
331 response_data = event.get("response", {})
332 terminal_status = response_data.get(
333 "status",
334 _event_to_status.get(event_type, "completed"),
335 )
337 # Extract error for failed and incomplete responses
338 if event_type == "response.failed" or event_type == "response.incomplete":
339 terminal_error = response_data.get("error")
341 # Core response fields
342 usage_data = response_data.get("usage")
343 reasoning_data = response_data.get("reasoning")
344 tool_choice_data = response_data.get("tool_choice")
345 tools_data = response_data.get("tools")
347 # Additional ResponsesAPIResponse fields
348 model_data = response_data.get("model")
349 instructions_data = response_data.get("instructions")
350 temperature_data = response_data.get("temperature")
351 top_p_data = response_data.get("top_p")
352 max_output_tokens_data = response_data.get("max_output_tokens")
353 previous_response_id_data = response_data.get("previous_response_id")
354 text_data = response_data.get("text")
355 truncation_data = response_data.get("truncation")
356 parallel_tool_calls_data = response_data.get("parallel_tool_calls")
357 user_data = response_data.get("user")
358 store_data = response_data.get("store")
359 incomplete_details_data = response_data.get("incomplete_details")
361 # Also update output from final response if available
362 if "output" in response_data:
363 final_output = response_data.get("output", [])
364 for item in final_output:
365 item_id = item.get("id")
366 if item_id:
367 output_items[item_id] = item
368 state_dirty = True
370 # Flush state to Redis if interval elapsed
371 await flush_state_if_needed()
373 except json.JSONDecodeError as e:
374 verbose_proxy_logger.warning("Failed to parse streaming chunk: %s", e)
376 # Final flush to ensure all accumulated state is saved
377 await flush_state_if_needed(force=True)
379 # Use the terminal status from the stream, default to "completed"
380 final_status: Final = terminal_status or "completed"
382 await polling_handler.update_state(
383 polling_id=polling_id,
384 status=final_status,
385 usage=usage_data,
386 error=terminal_error,
387 reasoning=reasoning_data,
388 tool_choice=tool_choice_data,
389 tools=tools_data,
390 model=model_data,
391 instructions=instructions_data,
392 temperature=temperature_data,
393 top_p=top_p_data,
394 max_output_tokens=max_output_tokens_data,
395 previous_response_id=previous_response_id_data,
396 text=text_data,
397 truncation=truncation_data,
398 parallel_tool_calls=parallel_tool_calls_data,
399 user=user_data,
400 store=store_data,
401 incomplete_details=incomplete_details_data,
402 )
404 verbose_proxy_logger.info(
405 "Finished background streaming for %s, status=%s, error=%s, incomplete_details=%s, output_items=%s",
406 polling_id,
407 final_status,
408 terminal_error,
409 incomplete_details_data,
410 len(output_items),
411 )
413 except Exception as e:
414 verbose_proxy_logger.error("Error in background streaming task for %s: %s", polling_id, e)
415 import traceback
417 verbose_proxy_logger.error(traceback.format_exc())
419 await polling_handler.update_state(
420 polling_id=polling_id,
421 status="failed",
422 error={
423 "type": "internal_error",
424 "message": str(e),
425 "code": "background_streaming_error",
426 },
427 )