Coverage for src/lexigram/web/sse/handler.py: 47%

73 statements  

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

1"""SSE Handler Base Class. 

2 

3Provides an abstract base class for Server-Sent Events handlers. 

4""" 

5 

6from __future__ import annotations 

7 

8from abc import ABC, abstractmethod 

9import asyncio 

10from collections.abc import AsyncGenerator 

11from typing import Any, ClassVar 

12 

13from starlette.requests import Request 

14 

15from lexigram.web.exceptions import TooManyConnectionsError 

16from lexigram.web.transport.sse import EventSourceResponse, ServerSentEvent 

17 

18 

19class AbstractSSEHandler(ABC): 

20 """Base class for SSE handlers. 

21 

22 Subclass this to create SSE endpoints with automatic event streaming, 

23 heartbeat support, and connection management. 

24 

25 Attributes: 

26 heartbeat_interval: Seconds between heartbeat events (default: 30) 

27 retry: Client retry interval in milliseconds (default: 3000) 

28 event_types: List of supported event types for documentation 

29 

30 Example: 

31 ```python 

32 @sse_endpoint("/events/{channel}") 

33 class ChannelEventsHandler(AbstractSSEHandler): 

34 heartbeat_interval = 15 

35 event_types = ["message", "join", "leave"] 

36 

37 async def stream(self, request): 

38 channel = request.path_params["channel"] 

39 async for event in channel_events(channel): 

40 yield {"event": "message", "data": event} 

41 ``` 

42 """ 

43 

44 # Configuration 

45 heartbeat_interval: int = 30 

46 retry: int = 3000 

47 event_types: ClassVar[list[str]] = [] 

48 max_connections: ClassVar[int] = 0 # 0 = unlimited 

49 

50 # Connection tracking (class-level, shared across instances) 

51 _active_connections: ClassVar[int] = 0 

52 _connection_lock: ClassVar[asyncio.Lock | None] = None 

53 _connections: set[Request] 

54 

55 # Route metadata (set by decorator) 

56 _path: str | None = None 

57 _guards: ClassVar[list[Any]] = [] 

58 

59 def __init__(self) -> None: 

60 self._connections = set() 

61 

62 @classmethod 

63 def _get_connection_lock(cls) -> asyncio.Lock: 

64 """Return the per-class asyncio lock, creating it lazily on first call. 

65 

66 Each concrete subclass gets its own lock so different handler types 

67 track their connections independently. 

68 """ 

69 if cls._connection_lock is None: 

70 cls._connection_lock = asyncio.Lock() 

71 return cls._connection_lock 

72 

73 @abstractmethod 

74 async def stream(self, request: Request) -> AsyncGenerator[dict[str, Any], None]: 

75 """Generate SSE events. 

76 

77 Override this method to yield events to the client. 

78 

79 Args: 

80 request: The incoming HTTP request 

81 

82 Yields: 

83 Dict with 'event' (optional), 'data', 'id' (optional), 'retry' (optional) 

84 """ 

85 yield {} # pragma: no cover 

86 

87 async def on_connect(self, request: Request) -> None: 

88 """Called when a client connects. 

89 

90 Override to perform setup when a client starts streaming. 

91 

92 Args: 

93 request: The incoming HTTP request 

94 """ 

95 

96 async def on_disconnect(self, request: Request) -> None: 

97 """Called when a client disconnects. 

98 

99 Override to perform cleanup when a client stops streaming. 

100 

101 Args: 

102 request: The incoming HTTP request 

103 """ 

104 

105 async def add(self, connection: Request) -> None: 

106 """Register a connection (satisfies :class:`ConnectionManagerProtocol`).""" 

107 async with type(self)._get_connection_lock(): 

108 self._connections.add(connection) 

109 type(self)._active_connections += 1 

110 

111 async def remove(self, connection: Request) -> None: 

112 """Unregister a connection (satisfies :class:`ConnectionManagerProtocol`).""" 

113 async with type(self)._get_connection_lock(): 

114 self._connections.discard(connection) 

115 type(self)._active_connections -= 1 

116 

117 async def broadcast(self, message: Any, exclude: Request | None = None) -> None: 

118 """Broadcast is not supported for SSE connections. 

119 

120 SSE is a unidirectional protocol — each client has its own event 

121 generator. Override :meth:`stream` and use a shared async queue 

122 or pub/sub channel to push events to multiple clients. 

123 

124 Raises: 

125 NotImplementedError: Always. 

126 """ 

127 raise NotImplementedError( 

128 "SSE broadcast requires a shared event source; " 

129 "use a pub/sub channel or async queue in your stream() implementation" 

130 ) 

131 

132 @property 

133 def count(self) -> int: 

134 """Return the number of active connections.""" 

135 return type(self)._active_connections 

136 

137 async def _create_event_generator( 

138 self, 

139 request: Request, 

140 ) -> AsyncGenerator[ServerSentEvent, None]: 

141 """Internal: Create the SSE event generator with heartbeat support.""" 

142 try: 

143 await self.on_connect(request) 

144 

145 # Initial retry instruction 

146 if self.retry: 

147 yield ServerSentEvent(data="", retry=self.retry) 

148 

149 # Simple approach: yield from stream with periodic heartbeat 

150 last_heartbeat = asyncio.get_running_loop().time() 

151 

152 async for event_data in self.stream(request): 

153 # Check if heartbeat is due 

154 now = asyncio.get_running_loop().time() 

155 if now - last_heartbeat >= self.heartbeat_interval: 

156 yield ServerSentEvent(data="", event="heartbeat") 

157 last_heartbeat = now 

158 

159 # Yield the actual event 

160 yield ServerSentEvent( 

161 data=event_data.get("data", event_data), 

162 event=event_data.get("event"), 

163 event_id=event_data.get("id"), 

164 retry=event_data.get("retry"), 

165 ) 

166 finally: 

167 await self.on_disconnect(request) 

168 

169 def create_event( 

170 self, 

171 event_type: str, 

172 data: Any, 

173 event_id: str | None = None, 

174 ) -> dict[str, Any]: 

175 """Helper to create an event dict. 

176 

177 Args: 

178 event_type: The event type name 

179 data: Event payload 

180 event_id: Optional event ID for client tracking 

181 

182 Returns: 

183 Dict suitable for yielding from stream() 

184 """ 

185 result = {"event": event_type, "data": data} 

186 if event_id: 

187 result["id"] = event_id 

188 return result 

189 

190 async def handle(self, request: Request) -> EventSourceResponse: 

191 """Handle an SSE request. 

192 

193 Checks ``max_connections`` before accepting; raises 

194 ``TooManyConnectionsError`` (503) if the cap is reached. 

195 Decrements the active connection count when the response stream closes. 

196 

197 Args: 

198 request: The incoming HTTP request 

199 

200 Returns: 

201 EventSourceResponse streaming events to the client 

202 """ 

203 async with type(self)._get_connection_lock(): 

204 if ( 

205 self.max_connections > 0 

206 and type(self)._active_connections >= self.max_connections 

207 ): 

208 raise TooManyConnectionsError( 

209 f"SSE connection limit reached ({self.max_connections})" 

210 ) 

211 # Register directly under the lock to keep the limit-check and 

212 # increment atomic. Intentionally bypasses add() to avoid 

213 # re-acquiring _connection_lock (asyncio.Lock is not reentrant). 

214 self._connections.add(request) 

215 type(self)._active_connections += 1 

216 

217 async def _guarded_generator() -> AsyncGenerator[ServerSentEvent, None]: 

218 try: 

219 async for event in self._create_event_generator(request): 

220 yield event 

221 finally: 

222 await self.remove(request) 

223 

224 return EventSourceResponse(_guarded_generator()) 

225 

226 

227__all__ = ["AbstractSSEHandler"]