Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/common_utils/sse_keepalive.py: 30%

95 statements  

« prev     ^ index     » next       coverage.py v7.15.2, created at 2026-10-10 12:01 +0000

1import asyncio 

2import contextlib 

3import math 

4from collections.abc import AsyncGenerator, Iterable, Mapping 

5from typing import Final 

6 

7import anyio 

8 

9from litellm.constants import STREAM_SSE_KEEPALIVE_PING_CHUNK 

10 

11ANTHROPIC_PING_SSE_CHUNK: Final = STREAM_SSE_KEEPALIVE_PING_CHUNK 

12SSE_COMMENT_PING: Final = ": ping\n\n" 

13SSE_COMMENT_PING_BYTES: Final = SSE_COMMENT_PING.encode() 

14# The byte form of proxy_server._SSE_FRAME_DELIMITERS, CR-only included: SSE 

15# terminates a line with CRLF, LF or CR, so a blank line is any of these three. 

16_SSE_FRAME_DELIMITERS: Final = (b"\r\n\r\n", b"\n\n", b"\r\r") 

17_SSE_DELIMITER_LOOKBACK: Final = max(len(delimiter) for delimiter in _SSE_FRAME_DELIMITERS) 

18_STREAM_START_TAIL: Final = b"\n\n" 

19_SSE_MEDIA_TYPE: Final = "text/event-stream" 

20 

21 

22def coerce_keepalive_interval(ping_interval_seconds: float | str | None) -> float | None: 

23 if ping_interval_seconds is None: 23 ↛ 25line 23 didn't jump to line 25 because the condition on line 23 was always true

24 return None 

25 try: 

26 interval: Final = float(ping_interval_seconds) 

27 except (TypeError, ValueError): 

28 return None 

29 if not math.isfinite(interval) or interval <= 0: 

30 return None 

31 return interval 

32 

33 

34def keepalive_ping_has_fired(elapsed_seconds: float, ping_interval_seconds: float | str | None) -> bool: 

35 """Whether a keepalive ping has already gone out, which flushes the response headers. 

36 

37 A caller that discovers a failure after that point cannot raise its way to the client, since 

38 the status line is already on the wire. With pings disabled nothing flushes early, so a raise 

39 still carries its real status. 

40 """ 

41 interval: Final = coerce_keepalive_interval(ping_interval_seconds) 

42 return interval is not None and elapsed_seconds >= interval 

43 

44 

45def wrap_sse_stream_with_keepalive_pings( 

46 stream: AsyncGenerator[str, None], 

47 ping_interval_seconds: float | str | None, 

48 ping_chunk: str = ANTHROPIC_PING_SSE_CHUNK, 

49) -> AsyncGenerator[str, None]: 

50 """Fill idle gaps in an SSE stream, including the one before its first chunk. 

51 

52 ``ping_chunk`` is what gets written into those gaps. It defaults to Anthropic's 

53 own ``ping`` event because that is the protocol the first caller speaks; a 

54 stream carrying anything else wants ``SSE_COMMENT_PING``, which is a comment 

55 every conformant SSE client discards rather than a frame it has to understand. 

56 """ 

57 interval: Final = coerce_keepalive_interval(ping_interval_seconds) 

58 if interval is None: 58 ↛ 60line 58 didn't jump to line 60 because the condition on line 58 was always true

59 return stream 

60 return _keepalive_ping_stream(stream=stream, ping_interval_seconds=interval, ping_chunk=ping_chunk) 

61 

62 

63async def _keepalive_ping_stream( 

64 stream: AsyncGenerator[str, None], 

65 ping_interval_seconds: float, 

66 ping_chunk: str, 

67) -> AsyncGenerator[str, None]: 

68 pending = asyncio.ensure_future(stream.__anext__()) 

69 try: 

70 while True: 

71 await asyncio.wait({pending}, timeout=ping_interval_seconds) 

72 if not pending.done(): 

73 yield ping_chunk 

74 continue 

75 try: 

76 yield pending.result() 

77 except StopAsyncIteration: 

78 return 

79 pending = asyncio.ensure_future(stream.__anext__()) 

80 finally: 

81 pending.cancel() 

82 with anyio.CancelScope(shield=True): 

83 with contextlib.suppress(BaseException): 

84 await pending 

85 await stream.aclose() 

86 

87 

88def is_sse_content_type(content_type: str | None) -> bool: 

89 return content_type is not None and content_type.split(";", 1)[0].strip().lower() == _SSE_MEDIA_TYPE 

90 

91 

92def split_complete_sse_frames(pending: bytes) -> tuple[bytes, bytes]: 

93 """Split buffered SSE bytes into ``(complete_frames, unterminated_tail)``.""" 

94 boundary_end: Final = max( 

95 (pending.rfind(delimiter) + len(delimiter) for delimiter in _SSE_FRAME_DELIMITERS if delimiter in pending), 

96 default=0, 

97 ) 

98 if boundary_end == 0: 

99 return b"", pending 

100 return pending[:boundary_end], pending[boundary_end:] 

101 

102 

103def wrap_passthrough_sse_bytes_with_keepalive_pings( 

104 stream: AsyncGenerator[bytes, None], 

105 ping_interval_seconds: float | str | None, 

106 upstream_headers: Mapping[str, str], 

107) -> AsyncGenerator[bytes, None]: 

108 """Fill upstream silence on a byte-relaying passthrough stream with SSE comments. 

109 

110 Passthrough routes relay upstream bytes verbatim, so a model that thinks for 

111 longer than an intermediary's idle read timeout has its connection dropped 

112 before the first token. Only streams the upstream itself declares as 

113 ``text/event-stream`` are wrapped: a comment spliced into a binary transport 

114 (AWS event streams on ``/bedrock``, protobuf, NDJSON) would corrupt it. 

115 """ 

116 interval: Final = coerce_keepalive_interval(ping_interval_seconds) 

117 if interval is None or not is_sse_content_type(upstream_headers.get("content-type")): 

118 return stream 

119 return _keepalive_ping_byte_stream(stream=stream, ping_interval_seconds=interval) 

120 

121 

122async def _keepalive_ping_byte_stream( 

123 stream: AsyncGenerator[bytes, None], 

124 ping_interval_seconds: float, 

125) -> AsyncGenerator[bytes, None]: 

126 pending = asyncio.ensure_future(stream.__anext__()) 

127 # The tail of the bytes relayed so far, long enough to hold any delimiter. 

128 # Seeded as a delimiter because a stream starts at a frame boundary, and kept 

129 # across chunks because a delimiter can be split between two transport reads, 

130 # which testing only the latest chunk would miss for the rest of the stream. 

131 recent_tail = _STREAM_START_TAIL # rebind-ok: rolling window over the relayed bytes 

132 try: 

133 while True: 

134 await asyncio.wait((pending,), timeout=ping_interval_seconds) 

135 if not pending.done(): 

136 # The relayed chunks are raw transport reads, not whole SSE 

137 # frames, so an upstream that stalls halfway through a frame 

138 # must not have a comment spliced into it. 

139 if recent_tail.endswith(_SSE_FRAME_DELIMITERS): 

140 yield SSE_COMMENT_PING_BYTES 

141 continue 

142 try: 

143 chunk: bytes = pending.result() 

144 except StopAsyncIteration: 

145 return 

146 if chunk: 

147 recent_tail = (recent_tail + chunk)[-_SSE_DELIMITER_LOOKBACK:] 

148 yield chunk 

149 pending = asyncio.ensure_future(stream.__anext__()) 

150 finally: 

151 pending.cancel() 

152 with anyio.CancelScope(shield=True): 

153 with contextlib.suppress(BaseException): 

154 await pending 

155 await stream.aclose() 

156 

157 

158def resolve_ttft_keepalive_interval( 

159 deployments: Iterable[Mapping[str, object]], 

160 global_interval: float | str | None, 

161) -> float | None: 

162 """The keepalive interval to use before the upstream has answered at all. 

163 

164 No deployment has served the request yet, so a per-deployment 

165 ``keepalive_seconds`` is only trusted when every candidate under the requested 

166 model carries the same one, which is how the mid-stream engine treats its own 

167 model_name fallback. Otherwise the operator's global default applies. 

168 

169 An explicit ``0`` survives as a disable, since coercion rejects it: that keeps 

170 an operator's documented hard disable working on this path too, rather than 

171 letting the global switch a deployment back on behind their back. 

172 

173 A client-supplied value is deliberately not consulted. Opening the response 

174 early is an operator decision, and a request must not be able to enable it for 

175 a deployment that never did. 

176 """ 

177 configured: Final = frozenset(_keepalive_param(deployment) for deployment in deployments) 

178 agreed: Final = next(iter(configured)) if len(configured) == 1 else None 

179 return coerce_keepalive_interval(global_interval if agreed is None else agreed) 

180 

181 

182def _keepalive_param(deployment: Mapping[str, object]) -> float | str | None: 

183 params: Final = deployment.get("litellm_params") 

184 if not isinstance(params, Mapping): 

185 return None 

186 value: Final = params.get("keepalive_seconds") 

187 return value if isinstance(value, (int, float, str)) else None