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
« prev ^ index » next coverage.py v7.15.4, created at 2026-08-25 04:37 +0800
1"""SSE improvements with backpressure and retry support.
3Adds backpressure handling, Last-Event-ID resume, and connection events.
4"""
6from __future__ import annotations
8import asyncio
9from typing import TYPE_CHECKING, Any
11from starlette.responses import StreamingResponse
13if TYPE_CHECKING:
14 from collections.abc import AsyncIterator, Callable
17class SSEBackpressureHandler:
18 """Handles backpressure for SSE streams."""
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
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
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
36 message = {"event": event, "data": data}
37 if event_id:
38 message["id"] = event_id
40 try:
41 self.queue.put_nowait(message)
42 return True
43 except asyncio.QueueFull:
44 # Buffer full - apply backpressure
45 return False
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
58 def close(self) -> None:
59 """Close the handler."""
60 self._closed = True
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
76class SSERetryTracker:
77 """Tracks Last-Event-ID for resume support."""
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
84 @property
85 def last_event_id(self) -> str | None:
86 return self._last_event_id
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
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 }
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]
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
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)
119 return result
122class SSEConnectionEvents:
123 """Callbacks for SSE connection lifecycle events."""
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
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)
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)
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)
157class SSEResponse:
158 """Enhanced SSE response with backpressure and retry support.
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 """
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()
184 async def stream(self) -> AsyncIterator[str]:
185 """Stream SSE events with backpressure handling."""
186 from lexigram.web.sse.heartbeat import get_heartbeat_scheduler
188 await self.events.fire_connect()
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)
195 try:
196 # Send retry timeout
197 yield f"retry: {self.retry_timeout}\n\n"
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
210 # Record for retry
211 if event_id:
212 self._retry_tracker.record_event(event_id, event_type, data)
214 # Send with backpressure
215 success = True
216 if self.backpressure:
217 success = await self._handler.send(event_type, data, event_id)
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
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()
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 )