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

1""" 

2Background Streaming Task for Polling Via Cache Feature 

3 

4Handles streaming responses from LLM providers and updates Redis cache 

5with partial results for polling. 

6 

7Follows OpenAI Response Streaming format: 

8https://platform.openai.com/docs/api-reference/responses-streaming 

9""" 

10 

11import asyncio 

12import json 

13from collections.abc import Callable, Mapping, Sequence 

14from typing import TYPE_CHECKING, Final, TypeAlias 

15 

16from fastapi import Request, Response 

17from fastapi.responses import StreamingResponse 

18from starlette.types import Message 

19from typing_extensions import ReadOnly, TypedDict 

20 

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 

26 

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 

31 

32 

33_JsonDict: TypeAlias = dict[str, object] 

34_JsonList: TypeAlias = list[object] 

35 

36 

37class _OutputItem(TypedDict, total=False): 

38 id: ReadOnly[str] 

39 content: ReadOnly[Sequence[object]] 

40 

41 

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]] 

62 

63 

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] 

72 

73 

74class _StreamEventParser: 

75 parse: Callable[[str], _StreamEvent] = staticmethod(json.loads) 

76 

77 

78def _sse_frame_data(frame: str) -> str | None: 

79 return next((line[6:].strip() for line in frame.splitlines() if line.startswith("data: ")), None) 

80 

81 

82async def _never_receive() -> Message: 

83 await asyncio.Event().wait() 

84 raise AssertionError("unreachable") 

85 

86 

87def detach_request_from_client(request: Request) -> Request: 

88 """Same scope (headers, parsed body, auth) but a receive() that never yields http.disconnect. 

89 

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) 

94 

95 

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 

117 

118 Follows OpenAI Response Streaming format: 

119 https://platform.openai.com/docs/api-reference/responses-streaming 

120 

121 Processes streaming events and builds Response object: 

122 https://platform.openai.com/docs/api-reference/responses/object 

123 """ 

124 

125 try: 

126 verbose_proxy_logger.info("Starting background streaming for %s", polling_id) 

127 

128 # Update status to in_progress (OpenAI format) 

129 await polling_handler.update_state( 

130 polling_id=polling_id, 

131 status="in_progress", 

132 ) 

133 

134 # Force streaming mode and remove background flag 

135 data["stream"] = True 

136 data.pop("background", None) 

137 

138 # Create processor 

139 processor: Final = ProxyBaseLLMRequestProcessing(data=data) 

140 

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 ) 

163 

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]() 

168 

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 

186 

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 

190 

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 } 

202 

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 

206 

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 

217 

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 ) 

224 

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") 

230 

231 if isinstance(chunk, str) and (chunk_data := _sse_frame_data(chunk)) is not None: 

232 if chunk_data == "[DONE]": 

233 break 

234 

235 try: 

236 event: _StreamEvent = _StreamEventParser.parse(chunk_data) 

237 event_type = event.get("type", "") 

238 

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 

247 

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", {}) 

252 

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 

261 

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", "") 

268 

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 

275 

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 

285 

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) 

291 

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 

306 

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 

314 

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 ) 

322 

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 ) 

336 

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") 

340 

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") 

346 

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") 

360 

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 

369 

370 # Flush state to Redis if interval elapsed 

371 await flush_state_if_needed() 

372 

373 except json.JSONDecodeError as e: 

374 verbose_proxy_logger.warning("Failed to parse streaming chunk: %s", e) 

375 

376 # Final flush to ensure all accumulated state is saved 

377 await flush_state_if_needed(force=True) 

378 

379 # Use the terminal status from the stream, default to "completed" 

380 final_status: Final = terminal_status or "completed" 

381 

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 ) 

403 

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 ) 

412 

413 except Exception as e: 

414 verbose_proxy_logger.error("Error in background streaming task for %s: %s", polling_id, e) 

415 import traceback 

416 

417 verbose_proxy_logger.error(traceback.format_exc()) 

418 

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 )