Coverage for src/lexigram/web/sse/backpressure.py: 21%

133 statements  

« prev     ^ index     » next       coverage.py v7.15.4, created at 2026-08-25 04:37 +0800

1"""SSE improvements with backpressure and retry support. 

2 

3Adds backpressure handling, Last-Event-ID resume, and connection events. 

4""" 

5 

6from __future__ import annotations 

7 

8import asyncio 

9from typing import TYPE_CHECKING, Any 

10 

11from starlette.responses import StreamingResponse 

12 

13if TYPE_CHECKING: 

14 from collections.abc import AsyncIterator, Callable 

15 

16 

17class SSEBackpressureHandler: 

18 """Handles backpressure for SSE streams.""" 

19 

20 def __init__(self, max_buffer_size: int = 100): 

21 self.max_buffer_size = max_buffer_size 

22 self._queue: asyncio.Queue | None = None 

23 self._closed = False 

24 

25 @property 

26 def queue(self) -> asyncio.Queue: 

27 if self._queue is None: 

28 self._queue = asyncio.Queue(maxsize=self.max_buffer_size) 

29 return self._queue 

30 

31 async def send(self, event: str, data: str, event_id: str | None = None) -> bool: 

32 """Send an event. Returns False if buffer is full (backpressure).""" 

33 if self._closed: 

34 return False 

35 

36 message = {"event": event, "data": data} 

37 if event_id: 

38 message["id"] = event_id 

39 

40 try: 

41 self.queue.put_nowait(message) 

42 return True 

43 except asyncio.QueueFull: 

44 # Buffer full - apply backpressure 

45 return False 

46 

47 async def wait_for_space(self, timeout: float = 5.0) -> bool: 

48 """Wait for space in the buffer.""" 

49 try: 

50 await asyncio.wait_for( 

51 self.queue.join(), 

52 timeout=timeout, 

53 ) 

54 return True 

55 except TimeoutError: 

56 return False 

57 

58 def close(self) -> None: 

59 """Close the handler.""" 

60 self._closed = True 

61 

62 async def events(self) -> AsyncIterator[dict[str, Any]]: 

63 """Yield events from the buffer.""" 

64 while not self._closed: 

65 try: 

66 event = await asyncio.wait_for( 

67 self.queue.get(), 

68 timeout=1.0, 

69 ) 

70 yield event 

71 self.queue.task_done() 

72 except TimeoutError: 

73 continue 

74 

75 

76class SSERetryTracker: 

77 """Tracks Last-Event-ID for resume support.""" 

78 

79 def __init__(self) -> None: 

80 self._last_event_id: str | None = None 

81 self._event_history: dict[str, dict[str, Any]] = {} 

82 self._max_history = 1000 

83 

84 @property 

85 def last_event_id(self) -> str | None: 

86 return self._last_event_id 

87 

88 def set_last_event_id(self, event_id: str) -> None: 

89 """Set the last event ID from request header.""" 

90 self._last_event_id = event_id 

91 

92 def record_event(self, event_id: str, event: str, data: str) -> None: 

93 """Record an event for potential resume.""" 

94 self._last_event_id = event_id 

95 self._event_history[event_id] = { 

96 "event": event, 

97 "data": data, 

98 "id": event_id, 

99 } 

100 

101 # Trim history if too large 

102 if len(self._event_history) > self._max_history: 

103 # Remove oldest 

104 oldest = next(iter(self._event_history)) 

105 del self._event_history[oldest] 

106 

107 def get_events_after(self, event_id: str) -> list[dict[str, Any]]: 

108 """Get events after a given event ID for resumption.""" 

109 result = [] 

110 found = False 

111 

112 for eid, event in self._event_history.items(): 

113 if eid == event_id: 

114 found = True 

115 continue 

116 if found: 

117 result.append(event) 

118 

119 return result 

120 

121 

122class SSEConnectionEvents: 

123 """Callbacks for SSE connection lifecycle events.""" 

124 

125 def __init__( 

126 self, 

127 on_connect: Callable | None = None, 

128 on_disconnect: Callable | None = None, 

129 on_error: Callable | None = None, 

130 ): 

131 self.on_connect = on_connect 

132 self.on_disconnect = on_disconnect 

133 self.on_error = on_error 

134 

135 async def fire_connect(self, *args: Any, **kwargs: Any) -> None: 

136 if self.on_connect: 

137 if asyncio.iscoroutinefunction(self.on_connect): 

138 await self.on_connect(*args, **kwargs) 

139 else: 

140 self.on_connect(*args, **kwargs) 

141 

142 async def fire_disconnect(self, *args: Any, **kwargs: Any) -> None: 

143 if self.on_disconnect: 

144 if asyncio.iscoroutinefunction(self.on_disconnect): 

145 await self.on_disconnect(*args, **kwargs) 

146 else: 

147 self.on_disconnect(*args, **kwargs) 

148 

149 async def fire_error(self, error: Exception, *args: Any, **kwargs: Any) -> None: 

150 if self.on_error: 

151 if asyncio.iscoroutinefunction(self.on_error): 

152 await self.on_error(error, *args, **kwargs) 

153 else: 

154 self.on_error(error, *args, **kwargs) 

155 

156 

157class SSEResponse: 

158 """Enhanced SSE response with backpressure and retry support. 

159 

160 Uses a :class:`~lexigram.web.sse.heartbeat.SSEHeartbeatScheduler` shared 

161 across all active connections with the same ``heartbeat_interval``, so 

162 only **one** background asyncio task is needed regardless of concurrency. 

163 """ 

164 

165 def __init__( 

166 self, 

167 generator: AsyncIterator[Any], 

168 *, 

169 backpressure: bool = True, 

170 max_buffer: int = 100, 

171 retry_timeout: int = 5000, 

172 heartbeat_interval: float = 30.0, 

173 events: SSEConnectionEvents | None = None, 

174 ): 

175 self.generator = generator 

176 self.backpressure = backpressure 

177 self.max_buffer = max_buffer 

178 self.retry_timeout = retry_timeout 

179 self.heartbeat_interval = heartbeat_interval 

180 self.events = events or SSEConnectionEvents() 

181 self._handler = SSEBackpressureHandler(max_buffer) 

182 self._retry_tracker = SSERetryTracker() 

183 

184 async def stream(self) -> AsyncIterator[str]: 

185 """Stream SSE events with backpressure handling.""" 

186 from lexigram.web.sse.heartbeat import get_heartbeat_scheduler 

187 

188 await self.events.fire_connect() 

189 

190 # Register with the shared heartbeat scheduler instead of creating 

191 # a per-connection asyncio task (fixes W8.2 unbounded-task issue). 

192 scheduler = get_heartbeat_scheduler(self.heartbeat_interval) 

193 await scheduler.register(self._handler) 

194 

195 try: 

196 # Send retry timeout 

197 yield f"retry: {self.retry_timeout}\n\n" 

198 

199 async for event in self.generator: 

200 # Format as SSE 

201 if isinstance(event, dict): 

202 data = event.get("data", "") 

203 event_type = event.get("event", "message") 

204 event_id = event.get("id") 

205 else: 

206 data = str(event) 

207 event_type = "message" 

208 event_id = None 

209 

210 # Record for retry 

211 if event_id: 

212 self._retry_tracker.record_event(event_id, event_type, data) 

213 

214 # Send with backpressure 

215 success = True 

216 if self.backpressure: 

217 success = await self._handler.send(event_type, data, event_id) 

218 

219 if success: 

220 output = f"event: {event_type}\n" 

221 for line in data.split("\n"): 

222 output += f"data: {line}\n" 

223 if event_id: 

224 output += f"id: {event_id}\n" 

225 output += "\n" 

226 yield output 

227 

228 except Exception as e: # noqa: BLE001 — event dispatch errors must not crash the SSE stream 

229 await self.events.fire_error(e) 

230 raise 

231 finally: 

232 await scheduler.unregister(self._handler) 

233 await self.events.fire_disconnect() 

234 self._handler.close() 

235 

236 def to_response(self) -> StreamingResponse: 

237 """Convert to a Starlette StreamingResponse.""" 

238 return StreamingResponse( 

239 self.stream(), 

240 media_type="text/event-stream", 

241 )