Skip to content
Open
Show file tree
Hide file tree
Changes from 5 commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
7 changes: 6 additions & 1 deletion litestar/channels/plugin.py
Original file line number Diff line number Diff line change
Expand Up @@ -103,6 +103,9 @@ def __init__(
self._subscriber_class = subscriber_class

self._channels: dict[str, set[Subscriber]] = {channel: set() for channel in channels or []}
# Declared channels are kept when empty (they back route handlers); arbitrary channels
# are dropped on unsubscribe to avoid unbounded growth.
self._declared_channels: set[str] = set(channels or [])
Comment thread
lesnik512 marked this conversation as resolved.
Outdated

def encode_data(self, data: LitestarEncodableType) -> bytes:
"""Encode data before storing it in the backend"""
Expand Down Expand Up @@ -226,7 +229,7 @@ async def unsubscribe(self, subscriber: Subscriber, channels: str | Iterable[str
channels_to_unsubscribe: set[str] = set()

for channel in channels:
channel_subscribers = self._channels[channel]
channel_subscribers = self._channels.get(channel, set())

try:
channel_subscribers.remove(subscriber)
Expand All @@ -235,6 +238,8 @@ async def unsubscribe(self, subscriber: Subscriber, channels: str | Iterable[str

if not channel_subscribers:
channels_to_unsubscribe.add(channel)
if channel not in self._declared_channels:
del self._channels[channel]

if all(subscriber not in queues for queues in self._channels.values()):
await subscriber.put(None) # this will stop any running task or generator by breaking the inner loop
Expand Down
42 changes: 42 additions & 0 deletions tests/unit/test_channels/test_plugin.py
Original file line number Diff line number Diff line change
Expand Up @@ -326,6 +326,48 @@ async def test_unsubscribe_last_subscriber_unsubscribes_backend(
assert not plugin._channels.get("foo")


async def test_unsubscribe_removes_arbitrary_channel(memory_backend: MemoryChannelsBackend) -> None:
plugin = ChannelsPlugin(backend=memory_backend, arbitrary_channels_allowed=True)
subscriber = await plugin.subscribe(channels="foo")
assert "foo" in plugin._channels

await plugin.unsubscribe(subscriber, channels="foo")

assert "foo" not in plugin._channels


async def test_unsubscribe_keeps_declared_channel(memory_backend: MemoryChannelsBackend) -> None:
plugin = ChannelsPlugin(backend=memory_backend, channels=["foo"])
subscriber = await plugin.subscribe(channels="foo")

await plugin.unsubscribe(subscriber, channels="foo")

assert "foo" in plugin._channels
assert plugin._channels["foo"] == set()


async def test_unsubscribe_arbitrary_channels_does_not_leak(memory_backend: MemoryChannelsBackend) -> None:
plugin = ChannelsPlugin(backend=memory_backend, arbitrary_channels_allowed=True)

for i in range(1000):
channel = f"channel_{i}"
subscriber = await plugin.subscribe(channels=channel)
await plugin.unsubscribe(subscriber, channels=channel)

assert plugin._channels == {}


async def test_unsubscribe_twice_after_arbitrary_channel_removed(memory_backend: MemoryChannelsBackend) -> None:
# Second unsubscribe must not raise KeyError once the channel entry has been removed.
plugin = ChannelsPlugin(backend=memory_backend, arbitrary_channels_allowed=True)
subscriber = await plugin.subscribe(channels="foo")

await plugin.unsubscribe(subscriber, channels="foo")
assert "foo" not in plugin._channels

await plugin.unsubscribe(subscriber, channels="foo")


async def _populate_channels_backend(*, message_count: int, channel: str, backend: ChannelsBackend) -> list[bytes]:
messages = [f"{channel} - message {i}".encode() for i in range(message_count)]

Expand Down