Coverage for src/lexigram/web/websocket/rooms.py: 28%

83 statements  

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

1"""WebSocket room management for group messaging. 

2 

3Provides RoomManager for join/leave/broadcast semantics. 

4""" 

5 

6from __future__ import annotations 

7 

8from dataclasses import dataclass 

9from typing import TYPE_CHECKING, Any, cast 

10 

11from lexigram.logging import get_logger 

12from lexigram.result import Err, Ok, Result 

13from lexigram.web.websocket.connection_id import ConnectionIDManager 

14 

15if TYPE_CHECKING: 

16 from starlette.websockets import WebSocket 

17 

18logger = get_logger(__name__) 

19 

20 

21@dataclass(frozen=True) 

22class RoomError: 

23 """Error type for room operations.""" 

24 

25 reason: str 

26 

27 

28class RoomManager: 

29 """Manages WebSocket rooms with join/leave/broadcast semantics.""" 

30 

31 def __init__( 

32 self, 

33 connection_id_manager: ConnectionIDManager | None = None, 

34 ) -> None: 

35 self._rooms: dict[str, set[WebSocket]] = {} 

36 # Track which rooms a user is in (by connection UUID) 

37 self._user_rooms: dict[str, set[str]] = {} 

38 self._connection_id_manager = ( 

39 connection_id_manager 

40 if connection_id_manager is not None 

41 else ConnectionIDManager() 

42 ) 

43 

44 async def join(self, room: str, websocket: WebSocket) -> None: 

45 """Add a WebSocket to a room.""" 

46 if room not in self._rooms: 

47 self._rooms[room] = set() 

48 self._rooms[room].add(websocket) 

49 

50 # Track user's rooms 

51 ws_uuid = self._connection_id_manager.get_or_create(websocket) 

52 if ws_uuid not in self._user_rooms: 

53 self._user_rooms[ws_uuid] = set() 

54 self._user_rooms[ws_uuid].add(room) 

55 

56 async def leave(self, room: str, websocket: WebSocket) -> None: 

57 """Remove a WebSocket from a room.""" 

58 if room in self._rooms: 

59 self._rooms[room].discard(websocket) 

60 if not self._rooms[room]: 

61 del self._rooms[room] 

62 

63 # Update user's room tracking 

64 ws_uuid = self._connection_id_manager.get_or_create(websocket) 

65 if ws_uuid in self._user_rooms: 

66 self._user_rooms[ws_uuid].discard(room) 

67 

68 async def leave_all(self, websocket: WebSocket) -> None: 

69 """Remove a WebSocket from all rooms.""" 

70 ws_uuid = self._connection_id_manager.get_or_create(websocket) 

71 if ws_uuid in self._user_rooms: 

72 rooms = list(self._user_rooms[ws_uuid]) 

73 for room in rooms: 

74 await self.leave(room, websocket) 

75 

76 async def broadcast( 

77 self, 

78 room: str, 

79 message: Any, 

80 exclude: WebSocket | None = None, 

81 ) -> None: 

82 """Broadcast a message to all WebSockets in a room.""" 

83 if room not in self._rooms: 

84 return 

85 

86 # Serialize message if needed 

87 if not isinstance(message, str): 

88 from lexigram import serialization as json 

89 

90 message = json.dumps(message) 

91 

92 # Send to all except excluded 

93 for ws in self._rooms[room]: 

94 if ws is exclude: 

95 continue 

96 try: 

97 await ws.send_text(message) 

98 except (OSError, RuntimeError): 

99 # Remove broken connections 

100 await self.leave(room, ws) 

101 

102 async def send_to( 

103 self, 

104 room: str, 

105 user_id: str, 

106 message: Any, 

107 ) -> Result[None, RoomError]: 

108 """Send a message to a specific user in a room (by user_id). 

109 

110 Note: This requires the WebSocket to have user_id stored in its state. 

111 """ 

112 if room not in self._rooms: 

113 return Err(RoomError(reason="room not found")) 

114 

115 # Serialize message if needed 

116 if not isinstance(message, str): 

117 from lexigram import serialization as json 

118 

119 message = json.dumps(message) 

120 

121 for ws in self._rooms[room]: 

122 if getattr(ws.state, "user_id", None) == user_id: 

123 try: 

124 await ws.send_text(message) 

125 return Ok(None) 

126 except (OSError, RuntimeError) as exc: 

127 logger.debug("websocket_send_failed", error=str(exc)) 

128 return Err(RoomError(reason=f"websocket disconnected: {exc}")) 

129 

130 return Err(RoomError(reason="user not in room")) 

131 

132 def members(self, room: str) -> set[WebSocket]: 

133 """Get all WebSockets in a room.""" 

134 return self._rooms.get(room, set()).copy() 

135 

136 def rooms_for(self, websocket: WebSocket) -> set[str]: 

137 """Get all rooms a WebSocket is in.""" 

138 ws_uuid = self._connection_id_manager.get_or_create(websocket) 

139 return self._user_rooms.get(ws_uuid, set()).copy() 

140 

141 def room_count(self, room: str) -> int: 

142 """Get the number of members in a room.""" 

143 return len(self._rooms.get(room, set())) 

144 

145 

146# Global room manager instance 

147_room_manager: RoomManager | None = None 

148 

149 

150def get_room_manager(context: Any | None = None) -> RoomManager: 

151 """Get the global RoomManager instance.""" 

152 from lexigram.di.resolution.context import get_resolver 

153 

154 resolver = get_resolver(context) 

155 if resolver: 

156 return cast("RoomManager", cast("Any", resolver).resolve_sync(RoomManager)) 

157 

158 global _room_manager 

159 if _room_manager is None: 

160 _room_manager = RoomManager() 

161 return _room_manager 

162 

163 

164def set_room_manager(manager: RoomManager) -> None: 

165 """Set a custom RoomManager instance.""" 

166 global _room_manager 

167 _room_manager = manager