diff --git a/.ai/ARCHITECTURE.md b/.ai/ARCHITECTURE.md index d4bd2ace56..5ad8beef99 100644 --- a/.ai/ARCHITECTURE.md +++ b/.ai/ARCHITECTURE.md @@ -850,7 +850,7 @@ outputs. Messages are encoded with `msgspec` (msgpack), a hard dependency. | `get(key, default=None)` | Read a value | | `set(key, value, ttl=None)` | Write a value; `ttl` = optional lifetime in seconds | | `delete(key)` | Remove a key | -| `publish(topic, message)` | Append a message to a topic | +| `publish(topic, message, ttl=None)` | Append a message to a topic; `ttl` = release the topic after that long idle | | `subscribe(topic, replay_from=None)` | Return a `Subscription` | **Key expiry (TTL).** `set(key, value, ttl=)` gives a key a bounded @@ -925,6 +925,27 @@ deployment is its own owner and pays no socket overhead. Knobs: `namespace` purpose, since each topic retains that many arbitrary payloads; raise it for a wider reconnect window). +**Topic lifetime.** `publish(topic, message, ttl=None)` takes an optional ttl +(seconds, at least `MIN_TOPIC_TTL` = 1s, so a reconnecting reader is not +outrun). A topic with a ttl is released, buffer and sequence both, once nobody +has published to or read from it for that long; a later publish starts it over +at 1, so a consumer returning with an old cursor gets `SharedStorageGap`. +Without a ttl a topic lives as long as the store. The latest publish's ttl +wins. Streaming publishes with `STREAM_TOPIC_TTL` (300s); user topics default +to no ttl. Every backend expires a topic as a whole, never message by message +while it is read (`tests/shared_storage/test_topic_ttl.py` runs the same cases +on all three). Local: the engine records the ttl on the topic and sweeps idle +ones (no publish, head or poll, and no call holding it) at most once a second, +on the next pub/sub call. The sweep also drops empty, unheld topics, such as +the one a returning reader's poll recreates after a release. Redis: the ttl sits in a third key; the publish +script `PEXPIRE`s all three to 1.25 ttl, and a poll script renews them. A poll +on such a topic blocks in `XREAD` for at most a third of that lifetime, not the +usual 5s, so a waiting reader renews it in time. Diskcache: the ttl +sits in a key; on publish or poll, once less than one ttl is left, every key of +the topic (counter, ttl, buffered messages) is renewed to 1.25 ttl, and at +once when a publish changes the ttl. A poll on such a topic waits at most half +a ttl before it renews again. + *Durability* is controlled by `mode` (the key/value store only — pub/sub is always transient): @@ -982,7 +1003,11 @@ election only reaches processes in the same network + filesystem namespace. in-tree or as a **separate package** — implements: - `get(key, default)` / `set(key, value)` / `delete(key)` — JSON-compatible values -- `publish(topic, message)` and `subscribe(topic, replay_from=None) -> Subscription` +- `publish(topic, message, ttl=None)` and `subscribe(topic, replay_from=None) -> Subscription` +- topic expiry: with a `ttl`, release the topic as a whole once nobody has + published to or read from it for that long, never message by message while + it is read; reject a ttl under `MIN_TOPIC_TTL` (run `test_topic_ttl.py` + against a new backend) - optional `start()` / `close()` (idempotent, called once per worker) and returns a `Subscription` (`__iter__` / `__aiter__` / `close`) that raises @@ -1383,6 +1408,15 @@ forged -- otherwise a client could read or inject into another page's topic. Across worker processes every worker must resolve the same signing secret (`secret_key`). +Each run of streams (from the first stream after idle until none is in flight) +also carries `&downlinkId=`, picked fresh by the client, and the connection id +is `:`: every run gets its own topic and the client's +cursor restarts at 0. Without it a page that sat idle past `STREAM_TOPIC_TTL` would +resume a cursor into a topic the store released, get `{reset: true}`, and fail +its new stream. The id only partitions the page's own space, so it is not +signed (just checked against `[A-Za-z0-9_-]{1,64}`). The downlink lifecycle +record (`connection_key`) stays keyed on the `end_id` alone. + The downlink is hosted in a SharedWorker (`dash-stream-worker.js`, served like the WebSocket worker; `config.stream.worker_url`) so **one connection per browser** serves every tab: browsers cap HTTP/1.1 connections per host at diff --git a/dash/_callback.py b/dash/_callback.py index 1b92907df1..e16ea51cf3 100644 --- a/dash/_callback.py +++ b/dash/_callback.py @@ -2,6 +2,7 @@ import hashlib import inspect import logging +import re import warnings from functools import wraps from typing import Callable, Optional, Any, List, Tuple, Union, Dict, TypeVar, cast @@ -485,8 +486,24 @@ def get_stream_connection_id() -> "str | None": worker must resolve the same secret: set a ``secret_key`` on the server, or cross-worker stream requests will not verify. Single-process apps are fine with no configuration. + + The renderer also sends a ``downlinkId`` it picks fresh for each run of + streams, giving every run its own topic (``:``). A run + then never resumes a cursor into a topic the store released while the page + sat idle. It only partitions the page's own space, so it needs no signing. """ - return get_request_end_id(_get_signing_secret()) + end_id = get_request_end_id(_get_signing_secret()) + if end_id is None: + return None + downlink_id = get_app().backend.request_adapter().args.get("downlinkId") + if not downlink_id: + return end_id + if not _DOWNLINK_ID_RE.fullmatch(downlink_id): + return None + return f"{end_id}:{downlink_id}" + + +_DOWNLINK_ID_RE = re.compile(r"[A-Za-z0-9_-]{1,64}") def _get_signing_secret() -> bytes: diff --git a/dash/_shared_storage/_engine.py b/dash/_shared_storage/_engine.py index fd1fc8b4f0..fc6ac6fd17 100644 --- a/dash/_shared_storage/_engine.py +++ b/dash/_shared_storage/_engine.py @@ -12,13 +12,19 @@ ``poll`` blocks a thread; ``apoll`` parks an asyncio task on a future that ``publish`` resolves from whichever thread it runs on, so an ASGI server can hold thousands of subscriptions without an executor thread each. + +A topic published with a ``ttl`` is dropped, buffer and sequence both, once +nobody has published to, polled or subscribed to it for that long, so +per-session topics do not pile up for the life of the process. A later publish +starts it over at sequence 1. """ import asyncio +import contextlib import threading import time from collections import deque -from typing import Any, Deque, Dict, List, NamedTuple, Optional, Tuple +from typing import Any, Deque, Dict, Iterator, List, NamedTuple, Optional, Tuple # Per-topic replay buffer size. Kept small by default because messages are # arbitrary user payloads and every topic retains up to this many -- unbounded @@ -28,6 +34,10 @@ # Deployments that need a wider reconnect window set buffer_size explicitly. DEFAULT_BUFFER = 32 +# How often idle topics are looked for, so a topic outlives its ttl by at most +# this much. +_SWEEP_INTERVAL = 1.0 + class PollResult(NamedTuple): messages: List[Any] @@ -39,14 +49,19 @@ class PollResult(NamedTuple): class _Topic: # pylint: disable=too-few-public-methods - __slots__ = ("seq", "buf", "cond", "waiters") + __slots__ = ("seq", "buf", "cond", "waiters", "users", "touched", "ttl") - def __init__(self, maxlen: int): + def __init__(self, maxlen: int, now: float): self.seq = 0 self.buf: Deque[Tuple[int, Any]] = deque(maxlen=maxlen) self.cond = threading.Condition() # asyncio tasks parked in apoll(), woken by the next publish/close. self.waiters: List[_Waiter] = [] + # Calls currently holding this topic (a blocked poll among them), and + # when the last one let go. Both guarded by the engine's _topics_lock. + self.users = 0 + self.touched = now + self.ttl: Optional[float] = None def _wake(fut: "asyncio.Future[None]") -> None: @@ -55,8 +70,13 @@ def _wake(fut: "asyncio.Future[None]") -> None: class StoreEngine: - def __init__(self, buffer_size: int = DEFAULT_BUFFER, persistence: Any = None): + def __init__( + self, + buffer_size: int = DEFAULT_BUFFER, + persistence: Any = None, + ): self._buffer_size = buffer_size + self._next_sweep = 0.0 # key -> (value, expiry). expiry is a monotonic deadline, or None for # no TTL. Expired entries are dropped lazily on the next read. self._data: Dict[str, Tuple[Any, Optional[float]]] = {} @@ -144,30 +164,57 @@ def snapshot_keys(self, keys: Any) -> Dict[str, Tuple[Any, Optional[float]]]: return out # --- pub/sub ----------------------------------------------------------- - def _topic(self, name: str) -> _Topic: + @contextlib.contextmanager + def _use(self, name: str) -> Iterator[_Topic]: + """Hold a topic for one call. A held topic is never swept, so a + publish cannot land in a topic that was just dropped from the map.""" + now = time.monotonic() with self._topics_lock: + self._sweep(now) topic = self._topics.get(name) if topic is None: - topic = self._topics[name] = _Topic(self._buffer_size) - return topic - - def publish(self, topic: str, message: Any) -> int: - t = self._topic(topic) - with t.cond: - t.seq += 1 - t.buf.append((t.seq, message)) - t.cond.notify_all() - waiters, t.waiters = t.waiters, [] - seq = t.seq + topic = self._topics[name] = _Topic(self._buffer_size, now) + topic.users += 1 + try: + yield topic + finally: + with self._topics_lock: + topic.users -= 1 + topic.touched = time.monotonic() + + def _sweep(self, now: float) -> None: + """Under ``_topics_lock``: drop topics idle past their ttl, and empty + ones a read left behind, which are no different from a missing one.""" + if now < self._next_sweep: + return + self._next_sweep = now + _SWEEP_INTERVAL + idle = [ + name + for name, t in self._topics.items() + if t.users == 0 + and (t.seq == 0 or (t.ttl is not None and t.touched < now - t.ttl)) + ] + for name in idle: + del self._topics[name] + + def publish(self, topic: str, message: Any, ttl: Optional[float] = None) -> int: + with self._use(topic) as t: + with t.cond: + t.ttl = ttl + t.seq += 1 + t.buf.append((t.seq, message)) + t.cond.notify_all() + waiters, t.waiters = t.waiters, [] + seq = t.seq for loop, fut in waiters: loop.call_soon_threadsafe(_wake, fut) return seq def head_seq(self, topic: str) -> int: """Current highest sequence -- where a fresh subscription starts.""" - t = self._topic(topic) - with t.cond: - return t.seq + with self._use(topic) as t: + with t.cond: + return t.seq def _ready(self, t: _Topic, after_seq: int) -> Optional[PollResult]: """Under ``t.cond``: the result available right now, or None to wait.""" @@ -196,22 +243,25 @@ def poll(self, topic: str, after_seq: int, timeout: float) -> PollResult: elapsed (caller re-polls) or the engine closed. ``gap`` is True when the next expected message was already evicted from the buffer. """ - t = self._topic(topic) deadline = time.monotonic() + timeout - with t.cond: - while True: - res = self._ready(t, after_seq) - if res is not None: - return res - remaining = deadline - time.monotonic() - if remaining <= 0: - return PollResult([], after_seq, False) - t.cond.wait(remaining) + with self._use(topic) as t: + with t.cond: + while True: + res = self._ready(t, after_seq) + if res is not None: + return res + remaining = deadline - time.monotonic() + if remaining <= 0: + return PollResult([], after_seq, False) + t.cond.wait(remaining) async def apoll(self, topic: str, after_seq: int, timeout: float) -> PollResult: """:meth:`poll` for asyncio: parks the task on a future instead of blocking a thread; ``publish`` (from any thread) or ``close`` wakes it.""" - t = self._topic(topic) + with self._use(topic) as t: + return await self._apoll(t, after_seq, timeout) + + async def _apoll(self, t: _Topic, after_seq: int, timeout: float) -> PollResult: loop = asyncio.get_running_loop() deadline = time.monotonic() + timeout while True: diff --git a/dash/_shared_storage/_transport.py b/dash/_shared_storage/_transport.py index fcfe5f3877..c8be5a7863 100644 --- a/dash/_shared_storage/_transport.py +++ b/dash/_shared_storage/_transport.py @@ -128,7 +128,8 @@ def _dispatch(self, req): # pylint: disable=too-many-return-statements self._engine.delete(req[1]) return ("ok", None) if op == "publish": - return ("ok", self._engine.publish(req[1], req[2])) + ttl = req[3] if len(req) > 3 else None + return ("ok", self._engine.publish(req[1], req[2], ttl)) if op == "head": return ("ok", self._engine.head_seq(req[1])) if op == "poll": diff --git a/dash/_shared_storage/base.py b/dash/_shared_storage/base.py index 2f0ac37d19..61cd6c50d2 100644 --- a/dash/_shared_storage/base.py +++ b/dash/_shared_storage/base.py @@ -19,6 +19,7 @@ import abc import asyncio +import math from typing import Any, AsyncIterator, Iterator, List, Optional, Tuple @@ -35,6 +36,18 @@ class SharedStorageGap(SharedStorageError): """ +# Shortest topic ttl a publish accepts. A reconnecting reader can easily be +# gone for a second, and a shorter ttl would drop its topic before it is back. +MIN_TOPIC_TTL = 1.0 + + +def check_topic_ttl(ttl: Optional[float]) -> None: + if ttl is not None and not (math.isfinite(ttl) and ttl >= MIN_TOPIC_TTL): + raise SharedStorageError( + f"topic ttl must be None or a finite {MIN_TOPIC_TTL}s or more, got {ttl!r}" + ) + + class Subscription(abc.ABC): """A live, ordered view of a topic. @@ -133,8 +146,16 @@ def delete(self, key: str) -> None: ... @abc.abstractmethod - def publish(self, topic: str, message: Any) -> None: - """Append ``message`` to ``topic``; delivered to every current subscriber.""" + def publish(self, topic: str, message: Any, ttl: Optional[float] = None) -> None: + """Append ``message`` to ``topic``; delivered to every current subscriber. + + ``ttl`` (seconds, at least ``MIN_TOPIC_TTL``) lets the store release the + topic, buffer and sequence both, once nobody has published to or read + from it for that long. A reader that keeps polling keeps it alive. A + later publish starts the topic over, so a consumer returning with an old + cursor gets ``SharedStorageGap``. ``None`` (the default) keeps the topic + for the life of the store. The latest publish's ``ttl`` applies. + """ # --- asyncio variants -------------------------------------------------- # Code running on an event loop (ASGI request handlers, the streaming @@ -155,9 +176,11 @@ async def aset(self, key: str, value: Any, ttl: Optional[float] = None) -> None: async def adelete(self, key: str) -> None: await asyncio.get_running_loop().run_in_executor(None, self.delete, key) - async def apublish(self, topic: str, message: Any) -> None: + async def apublish( + self, topic: str, message: Any, ttl: Optional[float] = None + ) -> None: await asyncio.get_running_loop().run_in_executor( - None, self.publish, topic, message + None, self.publish, topic, message, ttl ) @abc.abstractmethod diff --git a/dash/_shared_storage/diskcache.py b/dash/_shared_storage/diskcache.py index 8d909ab80e..67f5a5dc2b 100644 --- a/dash/_shared_storage/diskcache.py +++ b/dash/_shared_storage/diskcache.py @@ -13,7 +13,9 @@ sequence numbers, each message is stored under its sequence and old sequences are trimmed to a bounded window, and subscribers poll for sequences past their cursor. A consumer that falls farther behind than the window gets a -``SharedStorageGap``. +``SharedStorageGap``. A topic published with a ``ttl`` keeps it in a ttl key, +and every key of the topic expires together once nobody has published to or +polled it for that long, so abandoned topics leave the cache. """ import time from typing import Any, Optional @@ -21,13 +23,18 @@ from ._codec import decode, encode from ._engine import DEFAULT_BUFFER, PollResult from ._polling import PollingSubscription -from .base import BaseSharedStorage, Subscription +from .base import BaseSharedStorage, Subscription, check_topic_ttl # Poll cycle: short so a subscription's close() stays responsive; diskcache has # no server-side blocking wait, so this is a sleep-poll loop. _POLL_TIMEOUT = 1.0 _POLL_INTERVAL = 0.05 +# Every key of a topic gets this many ttls of life when renewed, and is renewed +# once less than one ttl is left: an idle topic goes between 1 and this many +# ttls after its last use, without rewriting its keys on every call. +_TTL_SLACK = 1.25 + def _require_diskcache(): try: @@ -85,6 +92,10 @@ def _seq(topic: str) -> str: def _msg(topic: str, seq: int) -> str: return f"ss:msg:{topic}:{seq}" + @staticmethod + def _ttl(topic: str) -> str: + return f"ss:ttl:{topic}" + # --- key/value --------------------------------------------------------- def get(self, key: str, default: Any = None) -> Any: raw = self._cache.get(self._kv(key)) @@ -97,19 +108,52 @@ def delete(self, key: str) -> None: self._cache.delete(self._kv(key)) # --- pub/sub ----------------------------------------------------------- - def publish(self, topic: str, message: Any) -> None: + def publish(self, topic: str, message: Any, ttl: Optional[float] = None) -> None: + check_topic_ttl(ttl) payload = encode(message) with self._cache.transact(): seq = int(self._cache.incr(self._seq(topic))) - self._cache.set(self._msg(topic, seq), payload) + # Outlives the counter, so a live topic never misses a message. + expire = ttl * _TTL_SLACK if ttl is not None else None + self._cache.set(self._msg(topic, seq), payload, expire=expire) evicted = seq - self._buffer_size if evicted >= 1: self._cache.delete(self._msg(topic, evicted)) + if ttl is not None: + changed = self._cache.get(self._ttl(topic)) != ttl + self._cache.set(self._ttl(topic), ttl, expire=expire) + self._keep_alive(topic, ttl, changed) + elif self._cache.delete(self._ttl(topic)): + self._expire_topic(topic, None) + + def _topic_keys(self, topic: str): + head = self._head(topic) + floor = max(1, head - self._buffer_size + 1) + yield self._seq(topic) + yield self._ttl(topic) + for seq in range(floor, head + 1): + yield self._msg(topic, seq) + + def _expire_topic(self, topic: str, expire: Optional[float]) -> None: + for key in self._topic_keys(topic): + self._cache.touch(key, expire=expire) + + def _keep_alive(self, topic: str, ttl: float, force: bool = False) -> None: + """Under ``transact``: renew the topic once less than ``ttl`` is left.""" + _value, expire_at = self._cache.get(self._seq(topic), expire_time=True) + if force or expire_at is None or expire_at - time.time() < ttl: + self._expire_topic(topic, ttl * _TTL_SLACK) def _head(self, topic: str) -> int: return int(self._cache.get(self._seq(topic), 0)) def _poll(self, topic: str, after_seq: int, timeout: float) -> PollResult: + ttl = self._cache.get(self._ttl(topic)) + if ttl is not None: + with self._cache.transact(): + self._keep_alive(topic, ttl) + # Come back to renew well before the topic could expire. + timeout = min(timeout, ttl / 2) deadline = time.monotonic() + timeout while True: head = self._head(topic) diff --git a/dash/_shared_storage/local.py b/dash/_shared_storage/local.py index 4d838799b0..c81ddfe299 100644 --- a/dash/_shared_storage/local.py +++ b/dash/_shared_storage/local.py @@ -37,7 +37,13 @@ recv_frame, send_frame, ) -from .base import BaseSharedStorage, SharedStorageError, SharedStorageGap, Subscription +from .base import ( + BaseSharedStorage, + SharedStorageError, + SharedStorageGap, + Subscription, + check_topic_ttl, +) _HAS_AF_UNIX = hasattr(socket, "AF_UNIX") _CLIENT_POLL_TIMEOUT = 20.0 # long-poll cycle for remote subscribers @@ -543,7 +549,7 @@ def _local(self, req): if op == "delete": return engine.delete(req[1]) if op == "publish": - return engine.publish(req[1], req[2]) + return engine.publish(req[1], req[2], req[3] if len(req) > 3 else None) if op == "head": return engine.head_seq(req[1]) raise ValueError(f"unknown op {op!r}") @@ -596,8 +602,11 @@ async def aset(self, key: str, value: Any, ttl: Optional[float] = None) -> None: async def adelete(self, key: str) -> None: await self._acall(["delete", key]) - async def apublish(self, topic: str, message: Any) -> None: - await self._acall(["publish", topic, message]) + async def apublish( + self, topic: str, message: Any, ttl: Optional[float] = None + ) -> None: + check_topic_ttl(ttl) + await self._acall(["publish", topic, message, ttl]) def get(self, key: str, default: Any = None) -> Any: return self._call(["get", key, default]) @@ -608,8 +617,9 @@ def set(self, key: str, value: Any, ttl: Optional[float] = None) -> None: def delete(self, key: str) -> None: self._call(["delete", key]) - def publish(self, topic: str, message: Any) -> None: - self._call(["publish", topic, message]) + def publish(self, topic: str, message: Any, ttl: Optional[float] = None) -> None: + check_topic_ttl(ttl) + self._call(["publish", topic, message, ttl]) def _head(self, topic: str) -> int: return self._call(["head", topic]) diff --git a/dash/_shared_storage/redis.py b/dash/_shared_storage/redis.py index 84e5175324..9b90758bf2 100644 --- a/dash/_shared_storage/redis.py +++ b/dash/_shared_storage/redis.py @@ -10,7 +10,9 @@ Sequence numbers are assigned by an atomic server-side script (``INCR`` + ``XADD``) so concurrent publishers stay strictly ordered; the stream is capped to a bounded window (``MAXLEN``), and a subscriber that falls past the trimmed -floor gets a ``SharedStorageGap``. +floor gets a ``SharedStorageGap``. A topic published with a ``ttl`` keeps it +in a third key, and all three expire that long after the last publish or poll, +so abandoned topics leave Redis. """ import os from typing import Any, Optional @@ -18,23 +20,54 @@ from ._codec import decode, encode from ._engine import DEFAULT_BUFFER, PollResult from ._polling import PollingSubscription -from .base import BaseSharedStorage, Subscription +from .base import BaseSharedStorage, Subscription, check_topic_ttl # Redis Stream XREAD blocks server-side, so a longer cycle than the diskcache # sleep-poll is fine; close() latency is bounded by this. _POLL_TIMEOUT = 5.0 _DEFAULT_URL = "redis://localhost:6379" +# A topic's keys live this many ttls past its last publish or poll. A poll on a +# topic with a ttl blocks for at most a third of that lifetime, so a waiting +# reader renews the keys long before they expire. +_TTL_SLACK = 1.25 + # Allocate the next sequence and append atomically, so concurrent publishers # never produce out-of-order stream IDs. Exact MAXLEN (not '~') keeps the replay # window at exactly buffer_size, so the gap boundary is deterministic. -# KEYS: seq counter, stream key. ARGV: encoded payload, maxlen. +# KEYS: seq counter, stream, ttl. ARGV: encoded payload, maxlen, lifetime in ms +# (0 for none). _PUBLISH_LUA = """ local seq = redis.call('INCR', KEYS[1]) redis.call('XADD', KEYS[2], 'MAXLEN', ARGV[2], seq .. '-0', 'm', ARGV[1]) +if tonumber(ARGV[3]) > 0 then + redis.call('SET', KEYS[3], ARGV[3], 'PX', ARGV[3]) + redis.call('PEXPIRE', KEYS[1], ARGV[3]) + redis.call('PEXPIRE', KEYS[2], ARGV[3]) +else + redis.call('DEL', KEYS[3]) + redis.call('PERSIST', KEYS[1]) + redis.call('PERSIST', KEYS[2]) +end return seq """ +# A poll renews the topic's lifetime, if it has one, and reads the head, the +# oldest buffered entry and the lifetime in ms. KEYS: seq counter, stream, ttl. +_POLL_HEAD_LUA = """ +local ttl = redis.call('GET', KEYS[3]) +if ttl then + redis.call('PEXPIRE', KEYS[1], ttl) + redis.call('PEXPIRE', KEYS[2], ttl) + redis.call('PEXPIRE', KEYS[3], ttl) +end +return { + redis.call('GET', KEYS[1]), + redis.call('XRANGE', KEYS[2], '-', '+', 'COUNT', 1), + ttl, +} +""" + def _require_redis(): try: @@ -87,10 +120,12 @@ def __init__( self._prefix = key_prefix self._buffer_size = buffer_size self._publish_script: Any = None + self._poll_head_script: Any = None def start(self) -> None: if self._publish_script is None: self._publish_script = self._redis.register_script(_PUBLISH_LUA) + self._poll_head_script = self._redis.register_script(_POLL_HEAD_LUA) def close(self) -> None: if self._owns_client: @@ -109,6 +144,12 @@ def _seq(self, topic: str) -> str: def _stream(self, topic: str) -> str: return f"{self._prefix}:stream:{topic}" + def _ttl(self, topic: str) -> str: + return f"{self._prefix}:ttl:{topic}" + + def _topic_keys(self, topic: str): + return [self._seq(topic), self._stream(topic), self._ttl(topic)] + # --- key/value --------------------------------------------------------- def get(self, key: str, default: Any = None) -> Any: raw = self._redis.get(self._kv(key)) @@ -124,11 +165,16 @@ def delete(self, key: str) -> None: self._redis.delete(self._kv(key)) # --- pub/sub ----------------------------------------------------------- - def publish(self, topic: str, message: Any) -> None: + def publish(self, topic: str, message: Any, ttl: Optional[float] = None) -> None: + check_topic_ttl(ttl) self.start() self._publish_script( - keys=[self._seq(topic), self._stream(topic)], - args=[encode(message), self._buffer_size], + keys=self._topic_keys(topic), + args=[ + encode(message), + self._buffer_size, + round(ttl * _TTL_SLACK * 1000) if ttl is not None else 0, + ], ) def _head(self, topic: str) -> int: @@ -136,21 +182,26 @@ def _head(self, topic: str) -> int: return int(raw) if raw is not None else 0 def _poll(self, topic: str, after_seq: int, timeout: float) -> PollResult: + self.start() stream = self._stream(topic) + raw_head, first, lifetime_ms = self._poll_head_script( + keys=self._topic_keys(topic) + ) + head = int(raw_head) if raw_head is not None else 0 # Cursor past the head: it was minted before the stream was reset (the - # key was flushed, or evicted under a maxmemory policy). Gap so the - # consumer resets rather than blocking on XREAD until the sequence climbs - # back past the cursor. - if after_seq > self._head(topic): + # key was flushed, expired, or evicted under a maxmemory policy). Gap so + # the consumer resets rather than blocking on XREAD until the sequence + # climbs back past the cursor. + if after_seq > head: return PollResult([], after_seq, True) # Gap: the next wanted sequence sits below the trimmed floor. Checked # before XREAD, which would otherwise silently resume at the floor. - first = self._redis.xrange(stream, count=1) if first and after_seq + 1 < _seq_of(first[0][0]): return PollResult([], after_seq, True) - entries = self._redis.xread( - {stream: f"{after_seq}-0"}, block=max(1, int(timeout * 1000)) - ) + block_ms = int(timeout * 1000) + if lifetime_ms is not None: + block_ms = min(block_ms, int(lifetime_ms) // 3) + entries = self._redis.xread({stream: f"{after_seq}-0"}, block=max(1, block_ms)) if not entries: return PollResult([], after_seq, False) items = entries[0][1] diff --git a/dash/_stream_hub.py b/dash/_stream_hub.py index 3f6a1878fb..944cb4da71 100644 --- a/dash/_stream_hub.py +++ b/dash/_stream_hub.py @@ -17,7 +17,10 @@ backend, never taken from the client), so a page can only ever read or write its own topic. The renderer hosts the downlink in a SharedWorker so every tab of the browser shares one connection: the worker pins the ``end_id`` of the first tab -that streams and sends it with every request for that connection. +that streams and sends it with every request for that connection. It also picks +a fresh downlink id for each run of streams, so the connection id is +``:`` and each run gets its own topic, which the store +releases once it sits idle. The lifecycle record stays keyed on the page. Downlink line shape (one JSON object per NDJSON line):: @@ -91,6 +94,10 @@ POLL_GRACE = 30.0 # How often a pump consults the connection record while a callback runs. DOWNLINK_CHECK_INTERVAL = 2.0 +# How long the store keeps a stream topic nobody publishes to or reads. Long +# enough to outlast the grace windows above many times over, short enough that +# a busy app does not hold every finished run's frames for hours. +STREAM_TOPIC_TTL = 300.0 # The uplink's fast acknowledgement -- the streaming callback's POST returns this # immediately; its outputs arrive on the downlink, not this response. @@ -125,7 +132,8 @@ def stream_topic(connection_id: str) -> str: def connection_key(connection_id: str) -> str: - return f"{_CONN_PREFIX}{connection_id}" + page_id = connection_id.split(":", 1)[0] + return f"{_CONN_PREFIX}{page_id}" def cancel_key(connection_id: str, request_id: str) -> str: @@ -151,7 +159,11 @@ def publish_frame( frame: Any, ) -> None: """Publish one streaming frame onto a connection's downlink topic.""" - storage.publish(stream_topic(connection_id), _envelope(request_id, frame)) + storage.publish( + stream_topic(connection_id), + _envelope(request_id, frame), + ttl=STREAM_TOPIC_TTL, + ) async def apublish_frame( @@ -161,7 +173,11 @@ async def apublish_frame( frame: Any, ) -> None: """:func:`publish_frame` for the pumps: never blocks their event loop.""" - await storage.apublish(stream_topic(connection_id), _envelope(request_id, frame)) + await storage.apublish( + stream_topic(connection_id), + _envelope(request_id, frame), + ttl=STREAM_TOPIC_TTL, + ) # --- downlink lifecycle record --------------------------------------------- diff --git a/dash/dash-renderer/src/utils/streamClient.ts b/dash/dash-renderer/src/utils/streamClient.ts index 5cc6892e1c..793496482a 100644 --- a/dash/dash-renderer/src/utils/streamClient.ts +++ b/dash/dash-renderer/src/utils/streamClient.ts @@ -120,11 +120,15 @@ export class StreamClient implements StreamTransport { // kept while any stream is in flight, so every tab behind a shared worker // publishes to and reads from the same topic. private endId = ''; + // Picked fresh with the endId for each run of streams, so each run reads a + // topic of its own: the previous run's topic may have been released by the + // server while the page sat idle, and its cursor means nothing here. + private downlinkId = ''; private pending = new Map(); private counter = 0; // Last sequence applied; the downlink resumes from here on reconnect. Starts - // at 0 so the first connect replays anything published before it subscribed - // (the uplink POST and the downlink open race). + // each run at 0 so the first connect replays anything published before it + // subscribed (the uplink POST and the downlink open race). private cursor = 0; private downlinkOpen = false; private abort: AbortController | null = null; @@ -194,7 +198,10 @@ export class StreamClient implements StreamTransport { return url; } const delim = url.includes('?') ? '&' : '?'; - return `${url}${delim}endId=${encodeURIComponent(this.endId)}`; + return ( + `${url}${delim}endId=${encodeURIComponent(this.endId)}` + + `&downlinkId=${encodeURIComponent(this.downlinkId)}` + ); } run( @@ -223,6 +230,8 @@ export class StreamClient implements StreamTransport { // streams are in flight the key stays put, so a stream from // another tab (shared worker) lands on the same topic. this.endId = endId || ''; + this.downlinkId = genId(); + this.cursor = 0; } const requestId = `${this.localId}-${++this.counter}`; const settled = new Promise((resolve, reject) => { @@ -447,7 +456,7 @@ export class StreamClient implements StreamTransport { if (!res.ok || !res.body) { throw new Error(`downlink responded ${res.status}`); } - received = await this.consume(res.body); + received = await this.consume(res.body, gen); if (received === 0 && !polling) { // Accepted then closed without a single envelope (a server // mid-shutdown, a proxy dropping idle connections): back @@ -517,7 +526,10 @@ export class StreamClient implements StreamTransport { } /** Relay envelopes until the connection ends; returns how many arrived. */ - private async consume(body: ReadableStream): Promise { + private async consume( + body: ReadableStream, + gen: number + ): Promise { const reader = body.getReader(); const decoder = new TextDecoder(); let buffer = ''; @@ -535,6 +547,10 @@ export class StreamClient implements StreamTransport { if (!line.trim()) { continue; // keepalive blank line } + if (this.loopGen !== gen) { + // Retired while reading: a newer run owns the cursor. + return received; + } received++; this.dispatchEnvelope(JSON.parse(line)); } diff --git a/dash/dash-renderer/tests/streamClient.test.js b/dash/dash-renderer/tests/streamClient.test.js index 44bd7f143b..1b1b04f45d 100644 --- a/dash/dash-renderer/tests/streamClient.test.js +++ b/dash/dash-renderer/tests/streamClient.test.js @@ -60,6 +60,12 @@ function makeFetch() { return {fetchImpl, uplinks, downlinks, mode}; } +// The URL every request of one run carries: the signed endId plus the run's +// downlink id. +const RUN_URL = /^\/cb\?endId=e1&downlinkId=[\w-]+$/; +const downlinkIdOf = url => + new URL(url, 'http://x').searchParams.get('downlinkId'); + const tick = (ms = 5) => new Promise(r => setTimeout(r, ms)); async function waitFor(pred, timeout = 1000) { const end = Date.now() + timeout; @@ -215,6 +221,35 @@ describe('StreamClient', () => { expect(mock.downlinks[1].from).to.equal(0); }); + it('gives each run of streams its own downlink id, starting from 0', async () => { + // The server may release an idle run's topic, so a later run must not + // resume that run's cursor: it reads a fresh topic from the start. + const first = client.run('/cb', {}, 'e1', {output: 'a'}, () => {}); + client.run('/cb', {}, 'e1', {output: 'b'}, () => {}); + await waitFor( + () => mock.uplinks.length === 2 && mock.downlinks.length === 1 + ); + const runA = downlinkIdOf(mock.uplinks[0].url); + expect(mock.uplinks[0].url).to.match(RUN_URL); + expect(downlinkIdOf(mock.uplinks[1].url)).to.equal(runA); + expect(downlinkIdOf(mock.downlinks[0].url)).to.equal(runA); + + const [ridA, ridB] = mock.uplinks.map( + u => u.streamConnection.requestId + ); + mock.downlinks[0].dl.push({rid: ridA, frame: {done: true}, seq: 7}); + mock.downlinks[0].dl.push({rid: ridB, frame: {done: true}, seq: 8}); + await first; + await waitFor(() => client.activeCount === 0); + + client.run('/cb', {}, 'e1', {output: 'c'}, () => {}); + await waitFor(() => mock.downlinks.length === 2); + const runB = downlinkIdOf(mock.uplinks[2].url); + expect(runB).to.not.equal(runA); + expect(downlinkIdOf(mock.downlinks[1].url)).to.equal(runB); + expect(mock.downlinks[1].from).to.equal(0); + }); + it('fails the callback loudly when the uplink is rejected (unverified connection)', async () => { // The server refuses an unverified multiplexed connection with a 403; no // frames will arrive on the downlink, so the request must reject rather @@ -449,7 +484,12 @@ describe('StreamClient cancellation', () => { let err; await settled.catch(e => (err = e)); expect(err.message).to.contain('cancelled'); - expect(mock.cancels).to.deep.equal([{url: '/cb?endId=e1', requestId}]); + expect(mock.cancels.length).to.equal(1); + expect(mock.cancels[0].requestId).to.equal(requestId); + expect(mock.cancels[0].url).to.match(RUN_URL); + expect(downlinkIdOf(mock.cancels[0].url)).to.equal( + downlinkIdOf(mock.uplinks[0].url) + ); // A late frame for the cancelled request is dropped, not delivered. client.dispatchEnvelope({ rid: requestId, @@ -466,10 +506,8 @@ describe('StreamClient cancellation', () => { client.run('/cb', {}, 'e1', {output: 'a'}, () => {}); client.run('/cb', {}, 'e2', {output: 'b'}, () => {}); await waitFor(() => mock.uplinks.length === 2); - expect(mock.uplinks.map(u => u.url)).to.deep.equal([ - '/cb?endId=e1', - '/cb?endId=e1' - ]); + expect(mock.uplinks[0].url).to.match(RUN_URL); + expect(mock.uplinks[1].url).to.equal(mock.uplinks[0].url); expect(client.connectionEndId).to.equal('e1'); }); }); @@ -497,7 +535,7 @@ describe('StreamClient in poll mode', () => { const withFrames = mock.polls.filter(p => p.n > 0); expect(withFrames.length).to.equal(2); expect(mock.polls[mock.polls.length - 1].from).to.equal(1); - expect(mock.polls[0].url).to.equal('/cb?endId=e1'); + expect(mock.polls[0].url).to.match(RUN_URL); expect(client.activeCount).to.equal(0); }); diff --git a/dash/dash-renderer/tests/streamWorkerHost.test.js b/dash/dash-renderer/tests/streamWorkerHost.test.js index 199d5ee1c9..af3eb365d7 100644 --- a/dash/dash-renderer/tests/streamWorkerHost.test.js +++ b/dash/dash-renderer/tests/streamWorkerHost.test.js @@ -5,6 +5,8 @@ import {SharedStreamClient, StreamClient} from '../src/utils/streamClient'; import {attachStreamWorkerHost} from '../src/utils/streamWorkerHost'; import {makeFetch, waitFor} from './helpers/streamMocks'; +const RUN_URL = /^\/cb\?endId=e1&downlinkId=[\w-]+$/; + // The worker side: one StreamClient (one downlink) behind a fake worker scope. // Each "tab" is a MessageChannel: port1 connects to the host, port2 is the // page's SharedStreamClient. @@ -39,7 +41,7 @@ describe('SharedWorker stream transport', () => { await waitFor(() => worker.mock.downlinks.length === 1); // The uplink carried the tab's signed endId and the payload. const conn = worker.mock.uplinks[0].streamConnection; - expect(worker.mock.uplinks[0].url).to.equal('/cb?endId=e1'); + expect(worker.mock.uplinks[0].url).to.match(RUN_URL); expect(worker.client.connectionEndId).to.equal('e1'); expect(worker.mock.uplinks[0].output).to.equal('a.b'); @@ -80,7 +82,7 @@ describe('SharedWorker stream transport', () => { await waitFor(() => worker.mock.uplinks.length === 2); expect(worker.mock.downlinks.length).to.equal(1); // Tab B's stream rides tab A's connection: one endId keys the topic. - expect(worker.mock.uplinks[1].url).to.equal('/cb?endId=e1'); + expect(worker.mock.uplinks[1].url).to.equal(worker.mock.uplinks[0].url); const ridA = worker.mock.uplinks[0].streamConnection.requestId; const ridB = worker.mock.uplinks[1].streamConnection.requestId; const dl = worker.mock.downlinks[0].dl; @@ -108,7 +110,7 @@ describe('SharedWorker stream transport', () => { tabA.release(); // tab A closed await waitFor(() => worker.mock.cancels.length === 1); expect(worker.mock.cancels[0]).to.deep.equal({ - url: '/cb?endId=e1', + url: worker.mock.uplinks[0].url, requestId: ridA }); // Tab B still has a stream in flight: the shared downlink stays open. diff --git a/tests/shared_storage/test_diskcache_backend.py b/tests/shared_storage/test_diskcache_backend.py index 9b833c5a6b..4c02c570e7 100644 --- a/tests/shared_storage/test_diskcache_backend.py +++ b/tests/shared_storage/test_diskcache_backend.py @@ -10,7 +10,7 @@ import pytest -from dash._shared_storage import DiskcacheSharedStorage, SharedStorageGap +from dash._shared_storage import DiskcacheSharedStorage, SharedStorageGap, base CTX = mp.get_context("spawn") @@ -94,6 +94,26 @@ def test_no_gap_at_buffer_edge(tmp_path): store.close() +def test_idle_topic_leaves_the_cache(tmp_path, monkeypatch): + monkeypatch.setattr(base, "MIN_TOPIC_TTL", 0.1) + store = DiskcacheSharedStorage(directory=str(tmp_path / "c")) + for i in range(3): + store.publish("t", f"m{i}", ttl=0.3) + time.sleep(0.5) + keys = [store._seq("t"), store._ttl("t")] + [store._msg("t", n) for n in (1, 2, 3)] + assert all(store._cache.get(k) is None for k in keys) + store.close() + + +def test_no_ttl_sets_no_expiry(tmp_path): + store = DiskcacheSharedStorage(directory=str(tmp_path / "c")) + store.publish("t", "m") + for key in (store._seq("t"), store._msg("t", 1)): + _value, expire = store._cache.get(key, expire_time=True) + assert expire is None + store.close() + + # --- cross-process (shared cache directory) ------------------------------- diff --git a/tests/shared_storage/test_engine.py b/tests/shared_storage/test_engine.py index 26a7e45545..208e28f5cf 100644 --- a/tests/shared_storage/test_engine.py +++ b/tests/shared_storage/test_engine.py @@ -173,7 +173,7 @@ async def scenario(): res = await e.apoll("t", 1, timeout=0.05) assert res.messages == [] and res.last_seq == 1 # A waiter that timed out was removed from the topic. - assert e._topic("t").waiters == [] + assert e._topics["t"].waiters == [] asyncio.run(scenario()) @@ -186,9 +186,90 @@ def test_apoll_wakes_on_close_and_serves_many_waiters(): async def scenario(): waits = [asyncio.ensure_future(e.apoll(f"t{i}", 0, 5.0)) for i in range(200)] await asyncio.sleep(0.05) - assert sum(len(e._topic(f"t{i}").waiters) for i in range(200)) == 200 + assert sum(len(e._topics[f"t{i}"].waiters) for i in range(200)) == 200 threading.Timer(0.05, e.close).start() results = await asyncio.wait_for(asyncio.gather(*waits), 5.0) assert all(r.messages == [] for r in results) asyncio.run(scenario()) + + +def _clocked_engine(monkeypatch): + from dash._shared_storage import _engine + + clock = {"t": 1000.0} + monkeypatch.setattr(_engine.time, "monotonic", lambda: clock["t"]) + return _engine.StoreEngine(), clock + + +def test_idle_topic_is_released(monkeypatch): + e, clock = _clocked_engine(monkeypatch) + for i in range(5): + e.publish("gone", i, ttl=10) + e.publish("kept", "x") + clock["t"] += 12 # past the ttl and a sweep interval + e.publish("other", "x") # any pub/sub call sweeps + assert sorted(e._topics) == ["kept", "other"] + # A later publish starts the topic over. + assert e.publish("gone", "again") == 1 + + +def test_polling_keeps_a_topic_alive(monkeypatch): + e, clock = _clocked_engine(monkeypatch) + e.publish("t", "m", ttl=10) + for _ in range(5): + clock["t"] += 6 # each gap is under the ttl, the total is well past it + e.poll("t", 1, timeout=0) + e.publish("other", "x") + assert e.head_seq("t") == 1 + + +def test_topic_held_by_a_blocked_poll_is_not_released(monkeypatch): + e, clock = _clocked_engine(monkeypatch) + e.publish("t", "m1", ttl=10) + t = e._topics["t"] + got = [] + poller = threading.Thread(target=lambda: got.append(e.poll("t", 1, timeout=100))) + poller.start() + until_parked = time.time() + 2 + while t.users == 0 and time.time() < until_parked: + time.sleep(0.01) + clock["t"] += 60 + e.publish("other", "x") + assert e._topics["t"] is t + e.publish("t", "m2") # still reaches the parked poll + poller.join(timeout=2) + assert got[0].messages == ["m2"] + + +def test_latest_publish_sets_the_ttl(monkeypatch): + e, clock = _clocked_engine(monkeypatch) + e.publish("t", "m", ttl=10) + e.publish("t", "m") + clock["t"] += 10_000 + e.publish("other", "x") + assert e.head_seq("t") == 2 + + +def test_released_topic_gaps_a_stale_cursor(monkeypatch): + # A consumer that comes back after its topic was released gets the gap + # signal, not a silent restart. + e, clock = _clocked_engine(monkeypatch) + for i in range(3): + e.publish("t", i, ttl=10) + clock["t"] += 12 + e.publish("t", "fresh") + assert e.poll("t", 3, timeout=0).gap is True + + +def test_empty_topic_left_by_a_read_is_released(monkeypatch): + # A browser back after its stream topic was released polls it again; the + # empty topic that poll creates has no ttl, but must not stay forever. + e, clock = _clocked_engine(monkeypatch) + e.publish("t", "m", ttl=10) + clock["t"] += 12 + assert e.poll("t", 1, timeout=0).gap is True + e.head_seq("never-published") + clock["t"] += 2 + e.publish("other", "x") + assert list(e._topics) == ["other"] diff --git a/tests/shared_storage/test_redis_backend.py b/tests/shared_storage/test_redis_backend.py index f86a650c02..626b1b76c9 100644 --- a/tests/shared_storage/test_redis_backend.py +++ b/tests/shared_storage/test_redis_backend.py @@ -17,6 +17,7 @@ from dash._shared_storage import ( # noqa: E402 RedisSharedStorage, SharedStorageGap, + base, ) REDIS_URL = os.environ.get("REDIS_URL", "redis://localhost:6379") @@ -122,6 +123,27 @@ def test_no_gap_at_buffer_edge(): store.close() +def test_idle_topic_leaves_redis(monkeypatch): + monkeypatch.setattr(base, "MIN_TOPIC_TTL", 0.1) + prefix = f"dash:sstest:{uuid.uuid4().hex[:12]}" + store = RedisSharedStorage(url=REDIS_URL, key_prefix=prefix) + store.start() + for i in range(3): + store.publish("t", f"m{i}", ttl=0.3) + keys = [store._seq("t"), store._stream("t"), store._ttl("t")] + assert all(0 < store._redis.pttl(k) <= 375 for k in keys) + time.sleep(0.5) + assert store._redis.exists(*keys) == 0 + store.close() + + +def test_no_ttl_sets_no_expiry(store): + store.publish("t", "m") + assert store._redis.pttl(store._seq("t")) == -1 + assert store._redis.pttl(store._stream("t")) == -1 + store._redis.delete(store._seq("t"), store._stream("t")) + + def test_two_instances_share_state(store): """A second client (separate connection pool) sees the first's writes and published messages -- the multi-worker / multi-pod case.""" diff --git a/tests/shared_storage/test_topic_ttl.py b/tests/shared_storage/test_topic_ttl.py new file mode 100644 index 0000000000..fba268ab39 --- /dev/null +++ b/tests/shared_storage/test_topic_ttl.py @@ -0,0 +1,144 @@ +"""Topic ttl contract, run against every backend so they expire topics the +same way: a topic goes as a whole once nobody publishes to or reads it, never +message by message while it is in use.""" +import os +import threading +import time +import uuid + +import pytest + +from dash._shared_storage import ( + DiskcacheSharedStorage, + LocalSharedStorage, + RedisSharedStorage, + SharedStorageError, + _engine, + base, +) + +REDIS_URL = os.environ.get("REDIS_URL", "redis://localhost:6379") +TTL = 0.5 + + +def _redis_available(): + try: + import redis # pylint: disable=import-outside-toplevel + + client = redis.Redis.from_url(REDIS_URL) + client.ping() + client.close() + return True + except Exception: # pylint: disable=broad-except + return False + + +def _make(kind, tmp_path): + tag = uuid.uuid4().hex[:12] + if kind == "local": + return LocalSharedStorage(namespace=f"ttl-{tag}") + if kind == "diskcache": + return DiskcacheSharedStorage(directory=str(tmp_path / "cache")) + if not _redis_available(): + pytest.skip("no Redis reachable at REDIS_URL") + return RedisSharedStorage(url=REDIS_URL, key_prefix=f"dash:sstest:{tag}") + + +BACKENDS = ["local", "diskcache", "redis"] + + +@pytest.fixture(params=BACKENDS) +def store(request, tmp_path, monkeypatch): + monkeypatch.setattr(base, "MIN_TOPIC_TTL", 0.1) + monkeypatch.setattr(_engine, "_SWEEP_INTERVAL", 0.05) + s = _make(request.param, tmp_path) + s.start() + try: + yield s + finally: + s.close() + + +def _idle(store): + time.sleep(TTL * 1.25 + 0.2) + # Local sweeps on the next pub/sub call. + store.publish("other", "x") + + +def test_idle_topic_is_released(store): + for i in range(3): + store.publish("t", f"m{i}", ttl=TTL) + _idle(store) + assert store.subscribe("t", replay_from=0).poll(0.0) == [] + store.publish("t", "again", ttl=TTL) + assert store.subscribe("t", replay_from=0).poll(0.0) == [(1, "again")] + + +def test_a_read_topic_keeps_every_buffered_message(store): + store.publish("t", "m1", ttl=TTL) + reader = store.subscribe("t", replay_from=1) + deadline = time.monotonic() + TTL * 3 + while time.monotonic() < deadline: + assert reader.poll(0.0) == [] + time.sleep(TTL / 5) + # Older than the ttl, but the topic was in use all along. + assert store.subscribe("t", replay_from=0).poll(0.0) == [(1, "m1")] + + +@pytest.mark.parametrize("kind", BACKENDS) +def test_a_waiting_reader_keeps_a_minimum_ttl_topic(kind, tmp_path): + # Real floor: a reader parked in one long poll must not outlast the topic. + s = _make(kind, tmp_path) + s.start() + s.publish("t", "m1", ttl=base.MIN_TOPIC_TTL) + reader = s.subscribe("t", replay_from=1) + errors = [] + + def read(): + try: + for _ in reader: + pass + except Exception as err: # pylint: disable=broad-except + errors.append(err) + + thread = threading.Thread(target=read, daemon=True) + thread.start() + time.sleep(base.MIN_TOPIC_TTL * 3) + try: + assert s.subscribe("t", replay_from=0).poll(0.0) == [(1, "m1")] + assert errors == [] + finally: + reader.close() + thread.join(timeout=10) + s.close() + + +def test_without_ttl_a_topic_is_kept(store): + store.publish("t", "m1") + _idle(store) + assert store.subscribe("t", replay_from=0).poll(0.0) == [(1, "m1")] + + +def test_a_publish_without_ttl_keeps_the_topic(store): + store.publish("t", "m1", ttl=TTL) + store.publish("t", "m2") + _idle(store) + assert store.subscribe("t", replay_from=0).poll(0.0) == [(1, "m1"), (2, "m2")] + + +def test_a_shorter_ttl_applies_at_once(store): + store.publish("t", "m1", ttl=60) + store.publish("t", "m2", ttl=TTL) + _idle(store) + assert store.subscribe("t", replay_from=0).poll(0.0) == [] + + +@pytest.mark.parametrize("kind", BACKENDS) +@pytest.mark.parametrize("ttl", [0, -1, 0.5, float("nan"), float("inf")]) +def test_a_ttl_under_the_floor_is_rejected(kind, ttl, tmp_path): + s = _make(kind, tmp_path) + try: + with pytest.raises(SharedStorageError): + s.publish("t", "m", ttl=ttl) + finally: + s.close() diff --git a/tests/streaming/test_stream_callbacks_integration.py b/tests/streaming/test_stream_callbacks_integration.py index 86f45bbbf7..1643c1e40c 100644 --- a/tests/streaming/test_stream_callbacks_integration.py +++ b/tests/streaming/test_stream_callbacks_integration.py @@ -330,3 +330,63 @@ async def stream_cb(_): before = dash_duo.find_element("#out").text until(lambda: dash_duo.find_element("#out").text != before, timeout=5) assert dash_duo.get_logs() == [] + + +def test_stst012_idle_stream_topics_are_released(dash_duo, monkeypatch): + """Each run of streams gets its own topic, and the store drops it once it + sits idle past ``STREAM_TOPIC_TTL``: page loads do not leave their frames + behind for the life of the process. A page that streams again after its + topic was dropped still gets every frame of the new run.""" + from dash import _stream_hub + from dash._shared_storage import LocalSharedStorage + + ttl = 1.0 + monkeypatch.setattr(_stream_hub, "STREAM_TOPIC_TTL", ttl) + app = Dash(__name__, shared_storage=LocalSharedStorage()) + app.layout = html.Div( + [html.Button("go", id="btn"), html.Div(id="out", children="idle")] + ) + + @app.callback( + Output("out", "children"), + Input("btn", "n_clicks"), + prevent_initial_call=True, + ) + async def stream_cb(n): + for i in range(5): + await asyncio.sleep(0.02) + yield f"{n}-{i}" + yield f"done-{n}" + + def stream_topics(): + engine = app.shared_storage._coord.engine + return {name for name in engine._topics if name.startswith("_dash_stream:")} + + def go_idle(): + # Past the ttl and one sweep interval, so the next publish sweeps. + time.sleep(ttl + 1.5) + + dash_duo.start_server(app) + dash_duo.find_element("#btn").click() + dash_duo.wait_for_text_to_equal("#out", "done-1") + first_page = stream_topics() + assert len(first_page) == 1 + + go_idle() + dash_duo.driver.refresh() + dash_duo.wait_for_text_to_equal("#out", "idle") + dash_duo.find_element("#btn").click() + dash_duo.wait_for_text_to_equal("#out", "done-1") + second_page = stream_topics() + assert len(second_page) == 1 + assert not second_page & first_page + + # Same page, new run after its topic was dropped: a fresh topic, read from + # its start, rather than a stale cursor into the released one. + go_idle() + dash_duo.find_element("#btn").click() + dash_duo.wait_for_text_to_equal("#out", "done-2") + third_run = stream_topics() + assert len(third_run) == 1 + assert not third_run & second_page + assert dash_duo.get_logs() == [] diff --git a/tests/streaming/test_stream_callbacks_unit.py b/tests/streaming/test_stream_callbacks_unit.py index 60d55abe48..2078acdec9 100644 --- a/tests/streaming/test_stream_callbacks_unit.py +++ b/tests/streaming/test_stream_callbacks_unit.py @@ -479,11 +479,11 @@ def test_stcb024_cancelled_pump_publishes_terminal_error(): published = [] class FakeStorage: - def publish(self, topic, message): + def publish(self, topic, message, ttl=None): published.append((topic, message)) # The pump talks to the store through its loop-native methods. - async def apublish(self, topic, message): + async def apublish(self, topic, message, ttl=None): self.publish(topic, message) async def aget(self, key, default=None): diff --git a/tests/streaming/test_stream_transport.py b/tests/streaming/test_stream_transport.py index 458cb93fda..7fe15ee2c1 100644 --- a/tests/streaming/test_stream_transport.py +++ b/tests/streaming/test_stream_transport.py @@ -116,6 +116,39 @@ def test_flask_uplink_without_valid_end_id_is_rejected(): storage.close() +def test_flask_downlink_id_gives_each_run_its_own_topic(): + # The renderer picks a downlink id per run of streams: the run's frames go + # to a topic of its own, while the lifecycle record stays on the page. + from dash import _stream_hub as hub + + app, storage = _streaming_app() + run = f"{CONNECTION_ID}:run1" + out = [] + th = _start_drain(storage, run, out) + client = app.server.test_client() + url = f"{_uplink_url(app)}&downlinkId=run1" + + assert client.post(url, json=_uplink_body("r1")).status_code == 200 + th.join(timeout=5) + _assert_delivered(out) + + resp = client.post(url, json={"streamDownlink": {"from": 0}}) + lines = [line for line in resp.get_data(as_text=True).split("\n") if line] + assert [json.loads(line)["rid"] for line in lines] == ["r1", "r1", "r1"] + assert storage.get(hub.connection_key(run))["mode"] == "poll" + assert hub.connection_key(run) == hub.connection_key(CONNECTION_ID) + storage.close() + + +def test_flask_rejects_a_malformed_downlink_id(): + app, storage = _streaming_app() + resp = app.server.test_client().post( + f"{_uplink_url(app)}&downlinkId=a:b", json={"streamDownlink": {"from": 0}} + ) + assert resp.status_code == 403 + storage.close() + + def test_flask_downlink_rejects_missing_end_id(): # A downlink with no valid signed endId cannot name a topic at all: the # server refuses it (403) rather than serving an attacker-named connection.