Coverage for object_streams/transports/channels.py: 87%

156 statements  

« prev     ^ index     » next       coverage.py v7.16.0, created at 2026-09-02 17:07 +0000

1"""Channels transport helpers.""" 

2 

3from __future__ import annotations 

4 

5import hashlib 

6from collections import deque 

7from collections.abc import Mapping 

8from typing import Any 

9 

10from asgiref.sync import async_to_sync 

11from channels.db import database_sync_to_async 

12from channels.generic.websocket import AsyncJsonWebsocketConsumer 

13from channels.layers import get_channel_layer 

14 

15from object_streams.events import ObjectRef 

16from object_streams.events import StreamEvent 

17from object_streams.models import ObjectStreamEvent 

18from object_streams.registry import ObjectStreamRegistry 

19from object_streams.registry import registry as default_registry 

20from object_streams.sessions import AsyncSubscriptionSession 

21from object_streams.subscriptions import ResyncRequired 

22from object_streams.subscriptions import SubscriptionKind 

23from object_streams.subscriptions import SubscriptionRequest 

24 

25 

26__all__ = ( 

27 "ObjectStreamConsumer", 

28 "broadcast_outbox_event", 

29 "broadcast_outbox_event_sync", 

30 "model_group_name", 

31 "object_group_name", 

32 "outbox_event_group_names", 

33 "subscription_group_names", 

34) 

35 

36 

37OUTBOX_EVENT_MESSAGE_TYPE = "object.stream.event" 

38 

39 

40def _digest(value: str) -> str: 

41 return hashlib.sha256(value.encode("utf-8")).hexdigest() 

42 

43 

44def object_group_name(ref: ObjectRef) -> str: 

45 """Return a stable Channels group name for an object subject.""" 

46 

47 return f"object_streams.object.{_digest(f'{ref.model}:{ref.pk}')}" 

48 

49 

50def model_group_name(model_label: str) -> str: 

51 """Return a stable Channels group name for model-level fanout.""" 

52 

53 return f"object_streams.model.{_digest(model_label)}" 

54 

55 

56def subscription_group_names(subscription: SubscriptionRequest) -> tuple[str, ...]: 

57 """Return Channels groups needed to wake a subscription.""" 

58 

59 if subscription.kind == SubscriptionKind.OBJECT: 

60 if subscription.pk is None: 60 ↛ 61line 60 didn't jump to line 61 because the condition on line 60 was never true

61 msg = "Object subscriptions require a primary key." 

62 raise ValueError(msg) 

63 return (object_group_name(ObjectRef(model=subscription.model, pk=subscription.pk)),) 

64 return (model_group_name(subscription.model),) 

65 

66 

67def outbox_event_group_names(event: StreamEvent) -> tuple[str, ...]: 

68 """Return Channels groups that should receive an outbox event.""" 

69 

70 return ( 

71 model_group_name(event.subject.model), 

72 object_group_name(event.subject), 

73 ) 

74 

75 

76async def broadcast_outbox_event(event_or_id: ObjectStreamEvent | int) -> None: 

77 """Fan out an outbox event id through the configured Channels layer.""" 

78 

79 channel_layer = get_channel_layer() 

80 if channel_layer is None: 80 ↛ 81line 80 didn't jump to line 81 because the condition on line 80 was never true

81 msg = "No Channels channel layer is configured." 

82 raise RuntimeError(msg) 

83 

84 row = await database_sync_to_async(_coerce_outbox_event)(event_or_id) 

85 event = await database_sync_to_async(row.to_stream_event)() 

86 message = { 

87 "type": OUTBOX_EVENT_MESSAGE_TYPE, 

88 "id": row.pk, 

89 } 

90 for group_name in outbox_event_group_names(event): 

91 await channel_layer.group_send(group_name, message) 

92 

93 

94def broadcast_outbox_event_sync(event_or_id: ObjectStreamEvent | int) -> None: 

95 """Synchronous wrapper for management commands and other Django call sites.""" 

96 

97 async_to_sync(broadcast_outbox_event)(event_or_id) 

98 

99 

100def _coerce_outbox_event(event_or_id: ObjectStreamEvent | int) -> ObjectStreamEvent: 

101 if isinstance(event_or_id, ObjectStreamEvent): 101 ↛ 102line 101 didn't jump to line 102 because the condition on line 101 was never true

102 return event_or_id 

103 return _get_outbox_event(event_or_id) 

104 

105 

106def _get_outbox_event(event_id: Any) -> ObjectStreamEvent: 

107 return ObjectStreamEvent.objects.select_related( 

108 "subject_content_type", 

109 "source_content_type", 

110 "source_history_content_type", 

111 ).get(pk=int(event_id)) 

112 

113 

114class ObjectStreamConsumer(AsyncJsonWebsocketConsumer): 

115 """Minimal Channels consumer for JSON object stream subscriptions.""" 

116 

117 registry: ObjectStreamRegistry = default_registry 

118 session_class = AsyncSubscriptionSession 

119 event_dedupe_size = 1024 

120 replay_limit = 1000 

121 max_subscriptions: int | None = 100 

122 max_member_pks: int | None = 10000 

123 

124 def __init__( 

125 self, 

126 *args: Any, 

127 registry: ObjectStreamRegistry | None = None, 

128 session_class: type[AsyncSubscriptionSession] | None = None, 

129 **kwargs: Any, 

130 ): 

131 super().__init__(*args, **kwargs) 

132 if registry is not None: 

133 self.registry = registry 

134 if session_class is not None: 134 ↛ 135line 134 didn't jump to line 135 because the condition on line 134 was never true

135 self.session_class = session_class 

136 self._subscription_groups: dict[str, tuple[str, ...]] = {} 

137 self._group_ref_counts: dict[str, int] = {} 

138 self._seen_outbox_ids: set[int] = set() 

139 self._seen_outbox_order: deque[int] = deque() 

140 

141 async def connect(self) -> None: 

142 self.session = self.session_class( 

143 user=self.scope.get("user"), 

144 request=self.scope, 

145 transport=self, 

146 registry=self.registry, 

147 replay_limit=self.replay_limit, 

148 max_subscriptions=self.max_subscriptions, 

149 max_member_pks=self.max_member_pks, 

150 ) 

151 await self.accept() 

152 

153 async def disconnect(self, code: int) -> None: 

154 await self._discard_all_groups() 

155 

156 async def receive_json(self, content: Any, **kwargs: Any) -> None: 

157 if not isinstance(content, Mapping): 157 ↛ 158line 157 didn't jump to line 158 because the condition on line 157 was never true

158 await self.send_error("invalid_request", "Messages must be JSON objects.") 

159 return 

160 await self.session.handle_message(content) 

161 

162 async def object_stream_event(self, message: Mapping[str, Any]) -> None: 

163 event_id = message.get("id") or message.get("outbox_id") 

164 if event_id is None: 164 ↛ 165line 164 didn't jump to line 165 because the condition on line 164 was never true

165 await self.send_error("invalid_event", "Object stream events require an outbox id.") 

166 return 

167 

168 try: 

169 outbox_id = int(event_id) 

170 except (TypeError, ValueError): 

171 await self.send_error("event_not_found", "Object stream event does not exist.") 

172 return 

173 

174 if not self._remember_outbox_id(outbox_id): 

175 return 

176 

177 try: 

178 row = await database_sync_to_async(_get_outbox_event)(outbox_id) 

179 except ObjectStreamEvent.DoesNotExist: 

180 await self.send_error("event_not_found", "Object stream event does not exist.") 

181 return 

182 

183 await self.session.publish(row) 

184 

185 async def send_subscribed(self, subscription: SubscriptionRequest) -> None: 

186 payload = subscription.as_dict() 

187 payload.pop("op", None) 

188 payload["type"] = "subscribed" 

189 await self.send_json(payload) 

190 

191 async def prepare_subscription(self, subscription: SubscriptionRequest) -> None: 

192 await self._add_subscription_groups(subscription) 

193 

194 async def send_unsubscribed(self, subscription_id: str) -> None: 

195 await self._discard_subscription_groups(subscription_id) 

196 await self.send_json( 

197 { 

198 "type": "unsubscribed", 

199 "subscription_id": subscription_id, 

200 } 

201 ) 

202 

203 async def send_event(self, event: StreamEvent) -> None: 

204 await self.send_json(event.as_dict()) 

205 

206 async def send_resync(self, resync: ResyncRequired) -> None: 

207 await self.send_json(resync.as_dict()) 

208 

209 async def send_error(self, code: str, message: str, *, details: Any = None) -> None: 

210 payload = { 

211 "type": "error", 

212 "code": code, 

213 "message": message, 

214 } 

215 if details is not None: 215 ↛ 216line 215 didn't jump to line 216 because the condition on line 215 was never true

216 payload["details"] = details 

217 await self.send_json(payload) 

218 

219 async def _add_subscription_groups(self, subscription: SubscriptionRequest) -> None: 

220 if subscription.subscription_id is None: 220 ↛ 221line 220 didn't jump to line 221 because the condition on line 220 was never true

221 return 

222 await self._discard_subscription_groups(subscription.subscription_id) 

223 groups = subscription_group_names(subscription) 

224 self._subscription_groups[subscription.subscription_id] = groups 

225 for group_name in groups: 

226 await self._add_group(group_name) 

227 

228 async def _discard_subscription_groups(self, subscription_id: str) -> None: 

229 groups = self._subscription_groups.pop(subscription_id, ()) 

230 for group_name in groups: 

231 await self._discard_group(group_name) 

232 

233 async def _discard_all_groups(self) -> None: 

234 for group_name in tuple(self._group_ref_counts): 

235 self._group_ref_counts[group_name] = 1 

236 await self._discard_group(group_name) 

237 self._subscription_groups.clear() 

238 

239 async def _add_group(self, group_name: str) -> None: 

240 count = self._group_ref_counts.get(group_name, 0) 

241 self._group_ref_counts[group_name] = count + 1 

242 if count == 0 and self.channel_layer is not None: 

243 await self.channel_layer.group_add(group_name, self.channel_name) 

244 

245 async def _discard_group(self, group_name: str) -> None: 

246 count = self._group_ref_counts.get(group_name, 0) 

247 if count <= 1: 

248 self._group_ref_counts.pop(group_name, None) 

249 if self.channel_layer is not None: 

250 await self.channel_layer.group_discard(group_name, self.channel_name) 

251 return 

252 self._group_ref_counts[group_name] = count - 1 

253 

254 def _remember_outbox_id(self, outbox_id: int) -> bool: 

255 if outbox_id in self._seen_outbox_ids: 

256 return False 

257 self._seen_outbox_ids.add(outbox_id) 

258 self._seen_outbox_order.append(outbox_id) 

259 while len(self._seen_outbox_order) > self.event_dedupe_size: 

260 expired_id = self._seen_outbox_order.popleft() 

261 self._seen_outbox_ids.discard(expired_id) 

262 return True