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
« prev ^ index » next coverage.py v7.16.0, created at 2026-09-02 17:07 +0000
1"""Channels transport helpers."""
3from __future__ import annotations
5import hashlib
6from collections import deque
7from collections.abc import Mapping
8from typing import Any
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
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
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)
37OUTBOX_EVENT_MESSAGE_TYPE = "object.stream.event"
40def _digest(value: str) -> str:
41 return hashlib.sha256(value.encode("utf-8")).hexdigest()
44def object_group_name(ref: ObjectRef) -> str:
45 """Return a stable Channels group name for an object subject."""
47 return f"object_streams.object.{_digest(f'{ref.model}:{ref.pk}')}"
50def model_group_name(model_label: str) -> str:
51 """Return a stable Channels group name for model-level fanout."""
53 return f"object_streams.model.{_digest(model_label)}"
56def subscription_group_names(subscription: SubscriptionRequest) -> tuple[str, ...]:
57 """Return Channels groups needed to wake a subscription."""
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),)
67def outbox_event_group_names(event: StreamEvent) -> tuple[str, ...]:
68 """Return Channels groups that should receive an outbox event."""
70 return (
71 model_group_name(event.subject.model),
72 object_group_name(event.subject),
73 )
76async def broadcast_outbox_event(event_or_id: ObjectStreamEvent | int) -> None:
77 """Fan out an outbox event id through the configured Channels layer."""
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)
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)
94def broadcast_outbox_event_sync(event_or_id: ObjectStreamEvent | int) -> None:
95 """Synchronous wrapper for management commands and other Django call sites."""
97 async_to_sync(broadcast_outbox_event)(event_or_id)
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)
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))
114class ObjectStreamConsumer(AsyncJsonWebsocketConsumer):
115 """Minimal Channels consumer for JSON object stream subscriptions."""
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
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()
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()
153 async def disconnect(self, code: int) -> None:
154 await self._discard_all_groups()
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)
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
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
174 if not self._remember_outbox_id(outbox_id):
175 return
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
183 await self.session.publish(row)
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)
191 async def prepare_subscription(self, subscription: SubscriptionRequest) -> None:
192 await self._add_subscription_groups(subscription)
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 )
203 async def send_event(self, event: StreamEvent) -> None:
204 await self.send_json(event.as_dict())
206 async def send_resync(self, resync: ResyncRequired) -> None:
207 await self.send_json(resync.as_dict())
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)
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)
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)
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()
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)
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
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