Coverage for src / lexigram / admin / realtime / websocket.py: 34%

185 statements  

« prev     ^ index     » next       coverage.py v7.13.5, created at 2026-08-13 22:14 +0800

1"""WebSocket support for lexigram-admin real-time updates. 

2 

3This module provides WebSocket handlers for bidirectional 

4real-time communication in admin. 

5 

6INT-03: WebSocket support for real-time updates. 

7""" 

8 

9from __future__ import annotations 

10 

11import contextlib 

12from dataclasses import dataclass, field 

13from datetime import UTC, datetime 

14from enum import StrEnum 

15from typing import TYPE_CHECKING, Any 

16 

17if TYPE_CHECKING: 

18 from collections.abc import Callable 

19 

20 

21def _utc_now() -> datetime: 

22 """Return current UTC time.""" 

23 return datetime.now(UTC) 

24 

25 

26# ============================================================================ 

27# Protocols for optional web integration 

28# ============================================================================ 

29 

30# Placeholder types for when lexigram-web is not available 

31# These are replaced via container registration when lexigram-web is present 

32 

33 

34class WebSocket: 

35 """Base WebSocket protocol. 

36 

37 This is a local implementation that does not depend on lexigram-web. 

38 When lexigram-web is available, its WebSocket is registered in the container 

39 and resolved through IoC. 

40 """ 

41 

42 async def send(self, data: str) -> None: 

43 """Send data through WebSocket.""" 

44 

45 async def receive(self) -> str: 

46 """Receive data from WebSocket.""" 

47 return "" 

48 

49 

50class WebSocketHandler: 

51 """Base WebSocket handler class. 

52 

53 Subclass this to implement WebSocket handlers. 

54 """ 

55 

56 async def on_connect(self, websocket: WebSocket) -> None: 

57 """Called when a client connects.""" 

58 

59 async def on_disconnect(self, websocket: WebSocket) -> None: 

60 """Called when a client disconnects.""" 

61 

62 async def on_message(self, websocket: WebSocket, message: str) -> None: 

63 """Called when a message is received.""" 

64 

65 

66# ============================================================================ 

67# Message Types 

68# ============================================================================ 

69 

70 

71class WSMessageType(StrEnum): 

72 """WebSocket message types.""" 

73 

74 # Client -> Server 

75 SUBSCRIBE = "subscribe" 

76 UNSUBSCRIBE = "unsubscribe" 

77 ACTION = "action" 

78 PING = "ping" 

79 

80 # Server -> Client 

81 EVENT = "event" 

82 NOTIFICATION = "notification" 

83 ERROR = "error" 

84 PONG = "pong" 

85 ACK = "ack" 

86 

87 

88@dataclass 

89class WSMessage: 

90 """WebSocket message.""" 

91 

92 type: WSMessageType | str 

93 data: dict[str, Any] = field(default_factory=dict) 

94 id: str | None = None 

95 timestamp: datetime = field(default_factory=_utc_now) 

96 

97 def to_dict(self) -> dict[str, Any]: 

98 """Convert to dictionary.""" 

99 return { 

100 "type": str( 

101 self.type.value if isinstance(self.type, WSMessageType) else self.type, 

102 ), 

103 "data": self.data, 

104 "id": self.id, 

105 "timestamp": self.timestamp.isoformat(), 

106 } 

107 

108 @classmethod 

109 def from_dict(cls, data: dict[str, Any]) -> WSMessage: 

110 """Create from dictionary.""" 

111 msg_type = data.get("type", "") 

112 with contextlib.suppress( 

113 ValueError 

114 ): # Keep as string if not a valid WSMessageType 

115 msg_type = WSMessageType(msg_type) 

116 

117 return cls( 

118 type=msg_type, 

119 data=data.get("data", {}), 

120 id=data.get("id"), 

121 ) 

122 

123 

124# ============================================================================ 

125# Connection Manager 

126# ============================================================================ 

127 

128 

129class AdminWebSocketManager: 

130 """Manager for admin WebSocket connections. 

131 

132 Handles connection tracking, subscriptions, and message routing. 

133 

134 Example: 

135 >>> manager = AdminWebSocketManager() 

136 >>> await manager.connect(websocket, user_id=1) 

137 >>> await manager.subscribe(websocket, resources=["users"]) 

138 >>> await manager.broadcast(message, resource="users") 

139 """ 

140 

141 def __init__(self) -> None: 

142 """Initialize the manager.""" 

143 self._connections: dict[str, Any] = {} # connection_id -> websocket 

144 self._user_connections: dict[Any, set[str]] = {} # user_id -> connection_ids 

145 self._resource_subscriptions: dict[ 

146 str, 

147 set[str], 

148 ] = {} # resource -> connection_ids 

149 self._connection_resources: dict[ 

150 str, 

151 set[str], 

152 ] = {} # connection_id -> resources 

153 self._connection_users: dict[str, Any] = {} # connection_id -> user_id 

154 

155 def _generate_connection_id(self) -> str: 

156 """Generate unique connection ID.""" 

157 import uuid 

158 

159 return str(uuid.uuid4())[:12] 

160 

161 async def connect( 

162 self, 

163 websocket: Any, 

164 user_id: Any | None = None, 

165 ) -> str: 

166 """Register a new connection. 

167 

168 Args: 

169 websocket: WebSocket connection 

170 user_id: Optional user ID 

171 

172 Returns: 

173 Connection ID 

174 """ 

175 connection_id = self._generate_connection_id() 

176 

177 self._connections[connection_id] = websocket 

178 self._connection_resources[connection_id] = set() 

179 

180 if user_id: 

181 self._connection_users[connection_id] = user_id 

182 if user_id not in self._user_connections: 

183 self._user_connections[user_id] = set() 

184 self._user_connections[user_id].add(connection_id) 

185 

186 return connection_id 

187 

188 async def disconnect(self, connection_id: str) -> None: 

189 """Unregister a connection. 

190 

191 Args: 

192 connection_id: Connection to remove 

193 """ 

194 # Remove from subscriptions 

195 if connection_id in self._connection_resources: 

196 for resource in self._connection_resources[connection_id]: 

197 if resource in self._resource_subscriptions: 

198 self._resource_subscriptions[resource].discard(connection_id) 

199 del self._connection_resources[connection_id] 

200 

201 # Remove from user connections 

202 if connection_id in self._connection_users: 

203 user_id = self._connection_users[connection_id] 

204 if user_id in self._user_connections: 

205 self._user_connections[user_id].discard(connection_id) 

206 del self._connection_users[connection_id] 

207 

208 # Remove connection 

209 self._connections.pop(connection_id, None) 

210 

211 async def subscribe( 

212 self, 

213 connection_id: str, 

214 resources: list[str], 

215 ) -> None: 

216 """Subscribe connection to resources. 

217 

218 Args: 

219 connection_id: Connection ID 

220 resources: List of resource types to subscribe to 

221 """ 

222 for resource in resources: 

223 if resource not in self._resource_subscriptions: 

224 self._resource_subscriptions[resource] = set() 

225 self._resource_subscriptions[resource].add(connection_id) 

226 

227 if connection_id in self._connection_resources: 

228 self._connection_resources[connection_id].add(resource) 

229 

230 async def unsubscribe( 

231 self, 

232 connection_id: str, 

233 resources: list[str], 

234 ) -> None: 

235 """Unsubscribe connection from resources. 

236 

237 Args: 

238 connection_id: Connection ID 

239 resources: List of resources to unsubscribe from 

240 """ 

241 for resource in resources: 

242 if resource in self._resource_subscriptions: 

243 self._resource_subscriptions[resource].discard(connection_id) 

244 

245 if connection_id in self._connection_resources: 

246 self._connection_resources[connection_id].discard(resource) 

247 

248 async def send( 

249 self, 

250 connection_id: str, 

251 message: WSMessage | dict[str, Any], 

252 ) -> bool: 

253 """Send message to specific connection. 

254 

255 Args: 

256 connection_id: Target connection 

257 message: Message to send 

258 

259 Returns: 

260 True if sent successfully 

261 """ 

262 websocket = self._connections.get(connection_id) 

263 if not websocket: 

264 return False 

265 

266 try: 

267 data = message.to_dict() if isinstance(message, WSMessage) else message 

268 await websocket.send_json(data) 

269 return True 

270 except (RuntimeError, ValueError, OSError): 

271 return False 

272 

273 async def broadcast( 

274 self, 

275 message: WSMessage | dict[str, Any], 

276 resource: str | None = None, 

277 user_ids: list[Any] | None = None, 

278 exclude_connections: list[str] | None = None, 

279 ) -> int: 

280 """Broadcast message to connections. 

281 

282 Args: 

283 message: Message to broadcast 

284 resource: Optional resource filter 

285 user_ids: Optional user IDs to target 

286 exclude_connections: Connection IDs to exclude 

287 

288 Returns: 

289 Number of connections that received the message 

290 """ 

291 exclude_connections = exclude_connections or [] 

292 target_connections: set[str] = set() 

293 

294 if user_ids: 

295 for user_id in user_ids: 

296 if user_id in self._user_connections: 

297 target_connections.update(self._user_connections[user_id]) 

298 elif resource: 

299 if resource in self._resource_subscriptions: 

300 target_connections.update(self._resource_subscriptions[resource]) 

301 else: 

302 target_connections.update(self._connections.keys()) 

303 

304 # Exclude specified connections 

305 target_connections -= set(exclude_connections) 

306 

307 # Send to all targets 

308 sent = 0 

309 for conn_id in target_connections: 

310 if await self.send(conn_id, message): 

311 sent += 1 

312 

313 return sent 

314 

315 async def send_to_user( 

316 self, 

317 user_id: Any, 

318 message: WSMessage | dict[str, Any], 

319 ) -> int: 

320 """Send message to all connections of a user. 

321 

322 Args: 

323 user_id: Target user 

324 message: Message to send 

325 

326 Returns: 

327 Number of connections that received the message 

328 """ 

329 return await self.broadcast(message, user_ids=[user_id]) 

330 

331 @property 

332 def connection_count(self) -> int: 

333 """Get total number of connections.""" 

334 return len(self._connections) 

335 

336 def get_user_connection_count(self, user_id: Any) -> int: 

337 """Get number of connections for a user.""" 

338 return len(self._user_connections.get(user_id, set())) 

339 

340 

341HAS_WEBSOCKET = True # local placeholder is always available 

342 

343# ============================================================================ 

344# Admin WebSocket Handler 

345# ============================================================================ 

346 

347 

348class AdminWebSocketHandler(WebSocketHandler if HAS_WEBSOCKET else object): # type: ignore[misc] 

349 """WebSocket handler for admin real-time updates. 

350 

351 Handles: 

352 - Subscription management for resources 

353 - Real-time event broadcasting 

354 - Client actions (HTMX-like operations) 

355 

356 Usage with lexigram.web: 

357 >>> @websocket_handler("/admin/ws") 

358 ... class AdminWSEndpoint(AdminWebSocketHandler): 

359 ... pass 

360 """ 

361 

362 ping_interval: int = 30 

363 ping_timeout: int = 10 

364 

365 def __init__(self) -> None: 

366 if HAS_WEBSOCKET: 

367 super().__init__() 

368 self._manager = AdminWebSocketManager() 

369 self._connection_ids: dict[Any, str] = {} # websocket -> connection_id 

370 self._action_handlers: dict[str, Callable[..., Any]] = {} 

371 

372 def register_action( 

373 self, 

374 action_name: str, 

375 handler: Callable[..., Any], 

376 ) -> None: 

377 """Register an action handler. 

378 

379 Args: 

380 action_name: Name of the action 

381 handler: Async function to handle the action 

382 """ 

383 self._action_handlers[action_name] = handler 

384 

385 async def on_connect(self, websocket: Any) -> None: 

386 """Handle WebSocket connection.""" 

387 await websocket.accept() 

388 

389 # Get user from websocket state/scope 

390 user = getattr(websocket, "user", None) 

391 if not user and hasattr(websocket, "scope"): 

392 user = websocket.scope.get("user") 

393 

394 user_id = getattr(user, "id", None) if user else None 

395 

396 connection_id = await self._manager.connect(websocket, user_id) 

397 self._connection_ids[websocket] = connection_id 

398 

399 # Send welcome message 

400 await websocket.send_json( 

401 WSMessage( 

402 type=WSMessageType.ACK, 

403 data={ 

404 "connection_id": connection_id, 

405 "message": "Connected to admin WebSocket", 

406 }, 

407 ).to_dict(), 

408 ) 

409 

410 async def on_message(self, websocket: Any, message: dict[str, Any]) -> None: 

411 """Handle incoming WebSocket message.""" 

412 connection_id = self._connection_ids.get(websocket) 

413 if not connection_id: 

414 return 

415 

416 try: 

417 msg = WSMessage.from_dict(message) 

418 

419 from lexigram.admin.realtime.ws_handler_registry import ( 

420 get_ws_message_type_registry, 

421 ) 

422 

423 registry = get_ws_message_type_registry() 

424 await registry.handle_message( 

425 msg.type, 

426 websocket, 

427 msg, 

428 connection_id, 

429 self._manager, 

430 ) 

431 

432 except (RuntimeError, ValueError, OSError) as e: 

433 await websocket.send_json( 

434 WSMessage( 

435 type=WSMessageType.ERROR, 

436 data={"message": str(e)}, 

437 ).to_dict(), 

438 ) 

439 

440 async def _handle_action(self, websocket: Any, msg: WSMessage) -> None: 

441 """Handle an action request.""" 

442 action_name = msg.data.get("action") 

443 if not action_name: 

444 await websocket.send_json( 

445 WSMessage( 

446 type=WSMessageType.ERROR, 

447 data={"message": "Action name required"}, 

448 id=msg.id, 

449 ).to_dict(), 

450 ) 

451 return 

452 

453 handler = self._action_handlers.get(action_name) 

454 if not handler: 

455 await websocket.send_json( 

456 WSMessage( 

457 type=WSMessageType.ERROR, 

458 data={"message": f"Unknown action: {action_name}"}, 

459 id=msg.id, 

460 ).to_dict(), 

461 ) 

462 return 

463 

464 try: 

465 result = await handler(msg.data) 

466 await websocket.send_json( 

467 WSMessage( 

468 type=WSMessageType.ACK, 

469 data={"result": result}, 

470 id=msg.id, 

471 ).to_dict(), 

472 ) 

473 except (RuntimeError, ValueError, OSError) as e: 

474 await websocket.send_json( 

475 WSMessage( 

476 type=WSMessageType.ERROR, 

477 data={"message": str(e)}, 

478 id=msg.id, 

479 ).to_dict(), 

480 ) 

481 

482 async def on_disconnect(self, websocket: Any) -> None: 

483 """Handle WebSocket disconnection.""" 

484 connection_id = self._connection_ids.pop(websocket, None) 

485 if connection_id: 

486 await self._manager.disconnect(connection_id) 

487 

488 

489# ============================================================================ 

490# Resource Change Notifier 

491# ============================================================================ 

492 

493 

494class ResourceChangeNotifier: 

495 """Notifies WebSocket clients of resource changes. 

496 

497 Integrate with your data layer to automatically notify 

498 clients when resources are created, updated, or deleted. 

499 

500 Example: 

501 >>> notifier = ResourceChangeNotifier() 

502 >>> 

503 >>> # After creating a user 

504 >>> await notifier.notify_created("users", user.id, user.to_dict()) 

505 """ 

506 

507 def __init__(self, manager: AdminWebSocketManager | None = None): 

508 self._manager = manager or AdminWebSocketManager() 

509 

510 async def notify_created( 

511 self, 

512 resource: str, 

513 resource_id: Any, 

514 data: dict[str, Any] | None = None, 

515 ) -> int: 

516 """Notify clients of resource creation.""" 

517 message = WSMessage( 

518 type=WSMessageType.EVENT, 

519 data={ 

520 "event": "resource.created", 

521 "resource": resource, 

522 "resource_id": resource_id, 

523 "data": data, 

524 }, 

525 ) 

526 return await self._manager.broadcast(message, resource=resource) 

527 

528 async def notify_updated( 

529 self, 

530 resource: str, 

531 resource_id: Any, 

532 changes: dict[str, Any] | None = None, 

533 ) -> int: 

534 """Notify clients of resource update.""" 

535 message = WSMessage( 

536 type=WSMessageType.EVENT, 

537 data={ 

538 "event": "resource.updated", 

539 "resource": resource, 

540 "resource_id": resource_id, 

541 "changes": changes, 

542 }, 

543 ) 

544 return await self._manager.broadcast(message, resource=resource) 

545 

546 async def notify_deleted( 

547 self, 

548 resource: str, 

549 resource_id: Any, 

550 ) -> int: 

551 """Notify clients of resource deletion.""" 

552 message = WSMessage( 

553 type=WSMessageType.EVENT, 

554 data={ 

555 "event": "resource.deleted", 

556 "resource": resource, 

557 "resource_id": resource_id, 

558 }, 

559 ) 

560 return await self._manager.broadcast(message, resource=resource) 

561 

562 async def notify_bulk_progress( 

563 self, 

564 resource: str, 

565 operation_id: str, 

566 progress: dict[str, Any], 

567 ) -> int: 

568 """Notify clients of bulk operation progress.""" 

569 message = WSMessage( 

570 type=WSMessageType.EVENT, 

571 data={ 

572 "event": "bulk.progress", 

573 "resource": resource, 

574 "operation_id": operation_id, 

575 "progress": progress, 

576 }, 

577 ) 

578 return await self._manager.broadcast(message, resource=resource) 

579 

580 

581__all__ = [ 

582 # Flags 

583 "HAS_WEBSOCKET", 

584 # Handler 

585 "AdminWebSocketHandler", 

586 # Manager 

587 "AdminWebSocketManager", 

588 # Notifier 

589 "ResourceChangeNotifier", 

590 "WSMessage", 

591 # Message types 

592 "WSMessageType", 

593]