feat: async events on asgi

This commit is contained in:
Alex committed 2026-06-08 22:02:21 +01:00
1 parent 1f06f31aa3
commit 73ed1cf607
13 files changed
+1270 -1054

No files matched your search

-135
View File
@@ -1,135 +0,0 @@
"""GET /api/messages/<message_id>/events — chat-stream reconnect endpoint.
Authenticates the caller, verifies ``message_id`` belongs to the user,
then hands off to ``build_message_event_stream`` for snapshot+tail.
"""
from __future__ import annotations
import logging
import re
from typing import Iterator, Optional
from flask import Blueprint, Response, jsonify, make_response, request, stream_with_context
from sqlalchemy import text
from application.core.settings import settings
from application.storage.db.session import db_readonly
from application.streaming.event_replay import (
DEFAULT_KEEPALIVE_SECONDS,
DEFAULT_POLL_TIMEOUT_SECONDS,
build_message_event_stream,
)
logger = logging.getLogger(__name__)
messages_bp = Blueprint("message_stream", __name__)
# A message_id is the canonical UUID hex format. Reject anything else
# before the SQL layer so a malformed cookie can't surface as a 500.
_MESSAGE_ID_RE = re.compile(
r"^[0-9a-fA-F]{8}-[0-9a-fA-F]{4}-[0-9a-fA-F]{4}-"
r"[0-9a-fA-F]{4}-[0-9a-fA-F]{12}$"
)
# ``sequence_no`` is a non-negative decimal integer. Anything else is
# corrupt client state — fall through to a fresh-replay cursor and let
# the snapshot reader catch the client up.
_SEQUENCE_NO_RE = re.compile(r"^\d+$")
def _normalise_last_event_id(raw: Optional[str]) -> Optional[int]:
if raw is None:
return None
raw = raw.strip()
if not raw or not _SEQUENCE_NO_RE.match(raw):
return None
return int(raw)
def _user_owns_message(message_id: str, user_id: str) -> bool:
"""Return True iff ``message_id`` belongs to ``user_id``."""
try:
with db_readonly() as conn:
row = conn.execute(
text(
"""
SELECT 1 FROM conversation_messages
WHERE id = CAST(:id AS uuid)
AND user_id = :u
LIMIT 1
"""
),
{"id": message_id, "u": user_id},
).first()
return row is not None
except Exception:
logger.exception(
"Ownership lookup failed for message_id=%s user_id=%s",
message_id,
user_id,
)
return False
@messages_bp.route("/api/messages/<message_id>/events", methods=["GET"])
def stream_message_events(message_id: str) -> Response:
decoded = getattr(request, "decoded_token", None)
user_id = decoded.get("sub") if isinstance(decoded, dict) else None
if not user_id:
return make_response(
jsonify({"success": False, "message": "Authentication required"}),
401,
)
if not _MESSAGE_ID_RE.match(message_id):
return make_response(
jsonify({"success": False, "message": "Invalid message id"}),
400,
)
if not _user_owns_message(message_id, user_id):
# Don't disclose whether the row exists — a malicious caller
# gets the same 404 whether the id is bogus, taken by another
# user, or simply gone.
return make_response(
jsonify({"success": False, "message": "Not found"}),
404,
)
raw_cursor = request.headers.get("Last-Event-ID") or request.args.get(
"last_event_id"
)
last_event_id = _normalise_last_event_id(raw_cursor)
keepalive_seconds = float(
getattr(settings, "SSE_KEEPALIVE_SECONDS", DEFAULT_KEEPALIVE_SECONDS)
)
@stream_with_context
def generate() -> Iterator[str]:
try:
yield from build_message_event_stream(
message_id,
last_event_id=last_event_id,
keepalive_seconds=keepalive_seconds,
poll_timeout_seconds=DEFAULT_POLL_TIMEOUT_SECONDS,
)
except GeneratorExit:
return
except Exception:
logger.exception(
"Reconnect stream crashed for message_id=%s user_id=%s",
message_id,
user_id,
)
response = Response(generate(), mimetype="text/event-stream")
response.headers["Cache-Control"] = "no-store"
response.headers["X-Accel-Buffering"] = "no"
response.headers["Connection"] = "keep-alive"
logger.info(
"message.event.connect message_id=%s user_id=%s last_event_id=%s",
message_id,
user_id,
last_event_id if last_event_id is not None else "-",
)
return response
+243
View File
@@ -0,0 +1,243 @@
"""Native-async (ASGI) SSE reader routes, mounted ahead of the Flask app.
These Starlette routes serve the chat-stream *reconnect* path on the event
loop, so a long-lived, mostly-idle tail costs a coroutine instead of one of
the 32 a2wsgi threadpool slots (see ``application/asgi.py``). They are the
sole reconnect reader — the old Flask blueprint has been removed. The heavy
*producer* (``POST /api/answer/stream`` → agent → LLM) stays on the sync
path untouched.
Auth, message-id validation, ``Last-Event-ID`` parsing and ownership are
done here; the snapshot/tail wire format is shared with the producer's
journal via ``build_message_event_stream_async`` → ``format_sse_event``.
"""
from __future__ import annotations
import logging
import re
from typing import Optional
import anyio
from sqlalchemy import text
from starlette.requests import Request
from starlette.responses import JSONResponse, StreamingResponse
from starlette.routing import Route
from application.auth import handle_auth
from application.core.settings import settings
from application.events.keys import connection_counter_key
from application.storage.db.session import db_readonly
from application.streaming.async_event_replay import (
build_message_event_stream_async,
)
from application.streaming.async_redis import get_async_redis_instance
from application.streaming.event_replay import (
DEFAULT_KEEPALIVE_SECONDS,
DEFAULT_POLL_TIMEOUT_SECONDS,
)
logger = logging.getLogger(__name__)
# Per-user concurrent-connection counter TTL (seconds) — orphaned counts
# from a hard crash self-heal after this window, mirroring the /api/events
# notification stream. The reconnect reader shares the same counter key, so
# the cap bounds a user's *total* live SSE footprint.
_COUNTER_TTL_SECONDS = 3600
# A message_id is the canonical UUID hex format. Reject anything else before
# the SQL layer so a malformed cookie can't surface as a 500.
_MESSAGE_ID_RE = re.compile(
r"^[0-9a-fA-F]{8}-[0-9a-fA-F]{4}-[0-9a-fA-F]{4}-"
r"[0-9a-fA-F]{4}-[0-9a-fA-F]{12}$"
)
# ``sequence_no`` is a non-negative decimal integer. Anything else is corrupt
# client state — fall through to a fresh-replay cursor.
_SEQUENCE_NO_RE = re.compile(r"^\d+$")
def _normalise_last_event_id(raw: Optional[str]) -> Optional[int]:
"""Parse a ``Last-Event-ID`` cursor; ``None`` for missing/invalid."""
if raw is None:
return None
raw = raw.strip()
if not raw or not _SEQUENCE_NO_RE.match(raw):
return None
return int(raw)
def _user_owns_message(message_id: str, user_id: str) -> bool:
"""Return True iff ``message_id`` belongs to ``user_id``."""
try:
with db_readonly() as conn:
row = conn.execute(
text(
"""
SELECT 1 FROM conversation_messages
WHERE id = CAST(:id AS uuid)
AND user_id = :u
LIMIT 1
"""
),
{"id": message_id, "u": user_id},
).first()
return row is not None
except Exception:
logger.exception(
"Ownership lookup failed for message_id=%s user_id=%s",
message_id,
user_id,
)
return False
_SSE_HEADERS = {
"Cache-Control": "no-store",
"X-Accel-Buffering": "no",
"Connection": "keep-alive",
# Marks the response as served by the event-loop reader rather than the
# WSGI-threaded Flask fallback. Purely diagnostic — the frontend reads
# the body via fetch+getReader and ignores response headers.
"X-SSE-Transport": "async",
}
def _json(message: str, status_code: int) -> JSONResponse:
return JSONResponse(
{"success": False, "message": message}, status_code=status_code
)
async def _acquire_stream_slot(user_id: str):
"""Reserve a per-user connection slot; returns ``(redis, key)`` to release.
Mirrors the ``/api/events`` cap (INCR + safety TTL, reject when the
post-increment count exceeds ``SSE_MAX_CONCURRENT_PER_USER``). Returns
``(None, None)`` when the cap is disabled or Redis is unavailable
(fail-open, like the notification stream). Raises ``_CapExceeded`` when
the user is over the cap so the caller can 429.
"""
cap = int(getattr(settings, "SSE_MAX_CONCURRENT_PER_USER", 0))
if cap <= 0:
return None, None
redis = await get_async_redis_instance()
if redis is None:
return None, None
key = connection_counter_key(user_id)
try:
current = int(await redis.incr(key))
except Exception:
logger.debug("async SSE counter INCR failed for user=%s", user_id)
return None, None
# EXPIRE failure must not bypass the cap, so it's best-effort after INCR.
try:
await redis.expire(key, _COUNTER_TTL_SECONDS)
except Exception:
logger.debug("async SSE counter EXPIRE failed for user=%s", user_id)
if current > cap:
await _release_stream_slot(redis, key)
raise _CapExceeded()
return redis, key
async def _release_stream_slot(redis, key) -> None:
if redis is None or key is None:
return
try:
await redis.decr(key)
except Exception:
logger.debug("async SSE counter DECR failed for key=%s", key)
class _CapExceeded(Exception):
"""Raised when a user is over their concurrent-stream cap."""
async def _counted_stream(inner, redis, key):
"""Wrap the reader so the per-user slot is released when it ends.
The slot is reserved before the response starts (so over-cap surfaces as
HTTP 429, not mid-stream); the release runs in ``finally`` on terminal
close, client disconnect, or error. Shielded so a disconnect-cancellation
can't skip the DECR and leak the count.
"""
try:
async for line in inner:
yield line
finally:
with anyio.CancelScope(shield=True):
await _release_stream_slot(redis, key)
async def stream_message_events(request: Request) -> JSONResponse | StreamingResponse:
"""GET /api/messages/{message_id}/events — async reconnect tail.
Mirrors the Flask handler's gates (auth → id format → ownership →
cursor → per-user connection cap) then streams snapshot+tail off the
event loop.
"""
# ``handle_auth`` only reads ``request.headers.get("Authorization")``;
# Starlette's headers are case-insensitive, so the Flask helper works
# verbatim. With AUTH_TYPE unset it returns ``{"sub": "local"}``.
decoded = handle_auth(request)
if isinstance(decoded, dict) and "error" in decoded:
return _json("Authentication error: invalid token", 401)
user_id = decoded.get("sub") if isinstance(decoded, dict) else None
if not user_id:
return _json("Authentication required", 401)
message_id = request.path_params["message_id"]
if not _MESSAGE_ID_RE.match(message_id):
return _json("Invalid message id", 400)
# Ownership check is a sync DB read — push it off the loop.
owns = await anyio.to_thread.run_sync(_user_owns_message, message_id, user_id)
if not owns:
# Same opaque 404 as the Flask route — don't disclose existence.
return _json("Not found", 404)
# Per-user concurrent-connection cap — reserve before the response opens
# so an over-cap caller gets a clean 429 instead of a mid-stream cutoff.
try:
redis, counter_key = await _acquire_stream_slot(user_id)
except _CapExceeded:
logger.warning("sse.reconnect.rejected user_id=%s (over cap)", user_id)
return _json("Too many concurrent SSE connections", 429)
raw_cursor = request.headers.get("Last-Event-ID") or request.query_params.get(
"last_event_id"
)
last_event_id = _normalise_last_event_id(raw_cursor)
keepalive_seconds = float(
getattr(settings, "SSE_KEEPALIVE_SECONDS", DEFAULT_KEEPALIVE_SECONDS)
)
logger.info(
"message.event.connect.async message_id=%s user_id=%s last_event_id=%s",
message_id,
user_id,
last_event_id if last_event_id is not None else "-",
)
stream = build_message_event_stream_async(
message_id,
last_event_id=last_event_id,
user_id=user_id,
keepalive_seconds=keepalive_seconds,
poll_timeout_seconds=DEFAULT_POLL_TIMEOUT_SECONDS,
)
return StreamingResponse(
_counted_stream(stream, redis, counter_key),
media_type="text/event-stream",
headers=_SSE_HEADERS,
)
# Mounted in ``application/asgi.py`` ahead of the Flask catch-all. Keep
# each route's path identical to the Flask blueprint it shadows.
async_sse_routes = [
Route(
"/api/messages/{message_id}/events",
stream_message_events,
methods=["GET"],
),
]
-2
View File
@@ -16,7 +16,6 @@ setup_logging()
from application.api import api # noqa: E402
from application.api.answer import answer # noqa: E402
from application.api.answer.routes.messages import messages_bp # noqa: E402
from application.api.devices import devices_bp # noqa: E402
from application.api.events.routes import events # noqa: E402
from application.api.internal.routes import internal # noqa: E402
@@ -59,7 +58,6 @@ app = Flask(__name__)
app.register_blueprint(user)
app.register_blueprint(answer)
app.register_blueprint(events)
app.register_blueprint(messages_bp)
app.register_blueprint(internal)
app.register_blueprint(connector)
app.register_blueprint(devices_bp)
+7
View File
@@ -8,6 +8,7 @@ from starlette.middleware import Middleware
from starlette.middleware.cors import CORSMiddleware
from starlette.routing import Mount
from application.api.async_sse import async_sse_routes
from application.app import app as flask_app
from application.mcp_server import mcp
@@ -18,6 +19,12 @@ mcp_app = mcp.http_app(path="/")
asgi_app = Starlette(
routes=[
Mount("/mcp", app=mcp_app),
# Native-async SSE readers intercept their exact paths before the
# Flask catch-all, so a mostly-idle reconnect tail rides the event
# loop instead of pinning a WSGI threadpool slot. Order matters:
# Starlette matches routes top-to-bottom, so these must precede the
# Mount("/") that hands everything else to Flask.
*async_sse_routes,
Mount("/", app=WSGIMiddleware(flask_app, workers=_WSGI_THREADPOOL)),
],
middleware=[
@@ -123,6 +123,7 @@ class MessageEventsRepository:
self,
message_id: str,
last_sequence_no: Optional[int] = None,
user_id: Optional[str] = None,
) -> list[dict]:
"""Return events with ``sequence_no > last_sequence_no``.
@@ -133,23 +134,38 @@ class MessageEventsRepository:
data the planner may pick a bitmap+sort. Either way the result
is sorted on ``sequence_no``.
When ``user_id`` is given the scan joins ``conversation_messages``
and filters on ``cm.user_id`` — a non-owner gets an empty result.
This lets the reconnect reader re-assert ownership at the data
layer rather than trusting only the route gate.
Returns a ``list`` (not a generator) so the underlying
``Result`` is fully drained before the caller can issue
another query on the same connection.
"""
cursor = -1 if last_sequence_no is None else int(last_sequence_no)
rows = self._conn.execute(
text(
"""
params = {"message_id": str(message_id), "cursor": cursor}
if user_id is None:
sql = """
SELECT message_id, sequence_no, event_type, payload, created_at
FROM message_events
WHERE message_id = CAST(:message_id AS uuid)
AND sequence_no > :cursor
ORDER BY sequence_no ASC
"""
),
{"message_id": str(message_id), "cursor": cursor},
).fetchall()
else:
params["u"] = user_id
sql = """
SELECT me.message_id, me.sequence_no, me.event_type,
me.payload, me.created_at
FROM message_events me
JOIN conversation_messages cm ON cm.id = me.message_id
WHERE me.message_id = CAST(:message_id AS uuid)
AND cm.user_id = :u
AND me.sequence_no > :cursor
ORDER BY me.sequence_no ASC
"""
rows = self._conn.execute(text(sql), params).fetchall()
return [row_to_dict(row) for row in rows]
def cleanup_older_than(self, ttl_days: int) -> int:
@@ -0,0 +1,119 @@
"""Async Redis pub/sub Topic for the native-async SSE reader.
Event-loop twin of :class:`application.streaming.broadcast_channel.Topic`.
Same contract — ``subscribe`` yields ``None`` on poll timeout (so the
caller can emit keepalives / run the watchdog) and ``bytes`` per delivered
message, fires ``on_subscribe`` once after Redis acks SUBSCRIBE, and tears
the pubsub down cleanly on client disconnect — but awaitable so an idle
stream costs a coroutine instead of a WSGI thread.
Publishing stays on the sync side (the producer writes via
``broadcast_channel.Topic.publish``); this is read-only fan-out.
"""
from __future__ import annotations
import inspect
import logging
from typing import AsyncIterator, Awaitable, Callable, Optional, Union
import anyio
from application.streaming.async_redis import get_async_redis_instance
logger = logging.getLogger(__name__)
OnSubscribe = Callable[[], Union[None, Awaitable[None]]]
class AsyncTopic:
"""An async pub/sub channel identified by a string name."""
def __init__(self, name: str) -> None:
self.name = name
async def subscribe(
self,
on_subscribe: Optional[OnSubscribe] = None,
poll_timeout: float = 1.0,
) -> AsyncIterator[Optional[bytes]]:
"""Subscribe to the topic; yield raw payloads or ``None`` on tick.
``on_subscribe`` runs (and is awaited if it returns a coroutine)
after Redis acks SUBSCRIBE — use it to seed snapshot state that
must be ordered after the subscriber is live but before the first
live message is processed. If Redis is unavailable, returns
immediately without yielding so the caller can fall back to a
direct snapshot read. Cleanly unsubscribes on close / disconnect.
"""
redis = await get_async_redis_instance()
if redis is None:
logger.debug(
"Async Redis unavailable; subscribe to %s yielded nothing",
self.name,
)
return
pubsub = redis.pubsub()
on_subscribe_fired = False
try:
try:
await pubsub.subscribe(self.name)
except Exception:
# Transient subscribe failure is treated like "Redis
# unavailable": yield nothing, let the caller fall back to
# its own snapshot read. The finally block still tears the
# pubsub down cleanly.
logger.exception("async pubsub.subscribe failed for %s", self.name)
return
while True:
try:
msg = await pubsub.get_message(timeout=poll_timeout)
except Exception:
logger.exception(
"async pubsub.get_message failed for %s", self.name
)
return
if msg is None:
yield None
continue
msg_type = msg.get("type")
if msg_type == "subscribe":
if not on_subscribe_fired and on_subscribe is not None:
try:
result = on_subscribe()
if inspect.isawaitable(result):
await result
except Exception:
logger.exception(
"on_subscribe callback failed for %s", self.name
)
on_subscribe_fired = True
continue
if msg_type != "message":
continue
data = msg.get("data")
if data is None:
continue
yield data if isinstance(data, bytes) else str(data).encode("utf-8")
finally:
# Client disconnect cancels this generator at the ``await
# get_message`` above; without shielding, the cancellation could
# re-fire mid-teardown and skip ``aclose()``, leaking the pooled
# connection back to nothing. Shield so unsubscribe + aclose
# always complete and the connection returns to the pool.
with anyio.CancelScope(shield=True):
if on_subscribe_fired:
try:
await pubsub.unsubscribe(self.name)
except Exception:
logger.debug(
"async pubsub unsubscribe error for %s",
self.name,
exc_info=True,
)
try:
await pubsub.aclose()
except Exception:
logger.debug(
"async pubsub close error for %s", self.name, exc_info=True
)
+232
View File
@@ -0,0 +1,232 @@
"""Native-async snapshot+tail iterator for chat-stream reconnect.
The sole reconnect reader: a ``: connected`` prelude, a snapshot flush
inside the SUBSCRIBE-ack callback, a dedup'd live tail, keepalive +
producer-liveness watchdog, and close-on-terminal — all as an async
generator driven off the event loop instead of a WSGI thread.
Wire format, dedup floor, and terminal detection come from
``event_replay`` (``read_snapshot_lines``, ``format_sse_event``,
``_decode_pubsub_message``, ``_payload_is_terminal``,
``_check_producer_liveness``), the same primitives the producer's journal
writes through — so the reader and writer cannot drift on wire shape. The
only sync I/O (snapshot read, watchdog DB probe) is pushed to a worker
thread via ``anyio.to_thread`` so it never blocks the loop.
"""
from __future__ import annotations
import logging
import time
from typing import AsyncIterator, Optional
import anyio
from application.streaming.async_broadcast_channel import AsyncTopic
from application.streaming.event_replay import (
DEFAULT_KEEPALIVE_SECONDS,
DEFAULT_POLL_TIMEOUT_SECONDS,
DEFAULT_PRODUCER_IDLE_SECONDS,
DEFAULT_WATCHDOG_INTERVAL_SECONDS,
_check_producer_liveness,
_decode_pubsub_message,
_payload_is_terminal,
format_sse_event,
read_snapshot_lines,
)
from application.streaming.keys import message_topic_name
logger = logging.getLogger(__name__)
# The snapshot read and watchdog probe are the reader's only DB I/O; they run
# in worker threads and each borrows a connection from the app-wide SQLAlchemy
# pool (pool_size=10 + max_overflow=20 = 30). The reader scales to many
# concurrent event-loop streams, so without a bound a burst of aligned
# watchdog/snapshot ticks could exhaust the pool and starve every other route.
# Cap the reader's concurrent DB-thread usage well below the pool so it can
# never monopolise it; excess ticks queue briefly (watchdog cadence is 5s, so
# a short queue delay is harmless).
_MAX_CONCURRENT_DB_READS = 8
_db_read_limiter: Optional[anyio.CapacityLimiter] = None
def _get_db_read_limiter() -> anyio.CapacityLimiter:
"""Lazily build the shared limiter on the (single) event loop.
Created on first use rather than at import so it binds to the running
loop; creation is synchronous, so the single-worker loop has no race.
"""
global _db_read_limiter
if _db_read_limiter is None:
_db_read_limiter = anyio.CapacityLimiter(_MAX_CONCURRENT_DB_READS)
return _db_read_limiter
async def build_message_event_stream_async(
message_id: str,
last_event_id: Optional[int] = None,
*,
user_id: Optional[str] = None,
keepalive_seconds: float = DEFAULT_KEEPALIVE_SECONDS,
poll_timeout_seconds: float = DEFAULT_POLL_TIMEOUT_SECONDS,
watchdog_interval_seconds: float = DEFAULT_WATCHDOG_INTERVAL_SECONDS,
producer_idle_seconds: float = DEFAULT_PRODUCER_IDLE_SECONDS,
) -> AsyncIterator[str]:
"""Yield SSE-formatted lines for one ``message_id`` reconnect stream.
First frame is ``: connected``; subsequent frames are snapshot rows,
live-tail events, or ``: keepalive`` comments. Runs until the client
disconnects or a terminal event is delivered.
"""
yield ": connected\n\n"
replay_buffer: list[str] = []
max_replayed_seq: Optional[int] = last_event_id
replay_done = False
replay_failed = False
terminal_in_snapshot = False
async def _load_snapshot() -> None:
nonlocal max_replayed_seq, replay_failed, terminal_in_snapshot
try:
lines, max_seq, terminal = await anyio.to_thread.run_sync(
read_snapshot_lines,
message_id,
last_event_id,
user_id,
limiter=_get_db_read_limiter(),
)
except Exception:
logger.exception(
"Snapshot read failed for message_id=%s last_event_id=%s",
message_id,
last_event_id,
)
replay_failed = True
return
replay_buffer.extend(lines)
max_replayed_seq = max_seq
terminal_in_snapshot = terminal
async def _on_subscribe() -> None:
# SUBSCRIBE acked — Postgres reads from this point capture every
# committed row; pub/sub messages published after this are queued
# at the connection level until the loop polls again.
nonlocal replay_done
try:
await _load_snapshot()
finally:
replay_done = True
topic = AsyncTopic(message_topic_name(message_id))
last_keepalive = time.monotonic()
last_watchdog_check = float("-inf")
watchdog_synthetic_seq = -1
try:
async for payload in topic.subscribe(
on_subscribe=_on_subscribe,
poll_timeout=poll_timeout_seconds,
):
# Flush snapshot exactly once after the SUBSCRIBE callback ran.
if replay_done and replay_buffer:
for line in replay_buffer:
yield line
replay_buffer.clear()
if terminal_in_snapshot:
# Original stream already finished; tailing would just
# emit keepalives forever.
return
if replay_failed:
yield format_sse_event(
{
"type": "error",
"error": "Stream replay failed; please refresh to load the latest state.",
"code": "snapshot_failed",
"message_id": message_id,
},
sequence_no=-1,
)
return
now = time.monotonic()
if payload is None:
# Idle tick — gate the watchdog on ``replay_done`` so we
# don't race the snapshot read on the first iteration.
if (
replay_done
and watchdog_interval_seconds >= 0
and now - last_watchdog_check >= watchdog_interval_seconds
):
last_watchdog_check = now
terminal_payload = await anyio.to_thread.run_sync(
_check_producer_liveness,
message_id,
user_id,
producer_idle_seconds,
limiter=_get_db_read_limiter(),
)
if terminal_payload is not None:
yield format_sse_event(
terminal_payload,
sequence_no=watchdog_synthetic_seq,
)
return
if now - last_keepalive >= keepalive_seconds:
yield ": keepalive\n\n"
last_keepalive = now
continue
envelope = _decode_pubsub_message(payload)
if envelope is None:
continue
seq = envelope.get("sequence_no")
inner = envelope.get("payload")
if (
not isinstance(seq, int)
or isinstance(seq, bool)
or not isinstance(inner, dict)
):
continue
if max_replayed_seq is not None and seq <= max_replayed_seq:
# Snapshot already covered this id — drop the duplicate.
continue
yield format_sse_event(inner, seq)
max_replayed_seq = seq
last_keepalive = now
if _payload_is_terminal(inner, envelope.get("event_type")):
return
# Subscribe exited without yielding (Redis unavailable / subscribe
# raised). The snapshot half is still in Postgres — read it
# directly so a Redis-only outage doesn't cost the client their
# backlog. Gate on ``replay_done`` so we don't double-read.
if not replay_done:
await _load_snapshot()
replay_done = True
for line in replay_buffer:
yield line
replay_buffer.clear()
if replay_failed:
yield format_sse_event(
{
"type": "error",
"error": "Stream replay failed; please refresh to load the latest state.",
"code": "snapshot_failed",
"message_id": message_id,
},
sequence_no=-1,
)
return
if terminal_in_snapshot:
return
except Exception:
# GeneratorExit / CancelledError are BaseException subclasses, so a
# client disconnect bypasses this handler and propagates to close
# the inner AsyncTopic generator (tearing its pubsub down in that
# generator's finally). Only genuine bugs land here.
logger.exception(
"Async reconnect stream crashed for message_id=%s", message_id
)
return
+47
View File
@@ -0,0 +1,47 @@
"""Lazy async Redis client for the native-async SSE reader.
Async twin of :func:`application.cache.get_redis_instance`. The
Starlette-mounted reader (``application.api.async_sse``) tails pub/sub on
the event loop, so it needs a ``redis.asyncio`` client rather than the
sync one used by the producer side. The app runs a single ASGI worker /
event loop, so a module-level singleton is sufficient and avoids
reconnecting per request.
"""
from __future__ import annotations
import logging
from typing import Optional
import redis.asyncio as aioredis
from application.core.settings import settings
logger = logging.getLogger(__name__)
_async_redis: Optional[aioredis.Redis] = None
_creation_failed = False
async def get_async_redis_instance() -> Optional[aioredis.Redis]:
"""Return a process-wide async Redis client, or ``None`` if unavailable.
``from_url`` builds the client without opening a socket (connection is
lazy), so a transient broker outage surfaces later on the first command
rather than here. Mirrors the sync client's ``socket_connect_timeout``
and ``health_check_interval`` so a half-open TCP can't wedge the tail
loop past its keepalive cadence.
"""
global _async_redis, _creation_failed
if _async_redis is None and not _creation_failed:
try:
_async_redis = aioredis.Redis.from_url(
settings.CACHE_REDIS_URL,
socket_connect_timeout=2,
health_check_interval=10,
)
except ValueError as e:
logger.error("Invalid Redis URL for async client: %s", e)
_creation_failed = True
_async_redis = None
return _async_redis
+64 -238
View File
@@ -1,8 +1,13 @@
"""Snapshot+tail iterator for chat-stream reconnect-after-disconnect.
"""Shared snapshot/replay primitives for chat-stream reconnect.
Subscribe to ``channel:{message_id}``, snapshot ``message_events``
rows past ``last_event_id`` inside the SUBSCRIBE-ack callback, flush
snapshot, then tail live pub/sub (dedup'd by ``sequence_no``). See
The reconnect reader itself is the native-async generator in
``async_event_replay.build_message_event_stream_async``; this module holds
the pieces both it and the producer's journal depend on: the SSE wire
format (``format_sse_event``), the ``message_events`` snapshot read
(``read_snapshot_lines``), the producer-liveness watchdog probe
(``_check_producer_liveness``), and the pub/sub envelope encode/decode.
Keeping them here lets the async reader and the sync journal agree on the
exact wire shape and dedup/terminal rules. See
``docs/runbooks/sse-notifications.md``.
"""
@@ -11,8 +16,7 @@ from __future__ import annotations
import json
import logging
import re
import time
from typing import Iterator, Optional
from typing import Optional
from sqlalchemy import text as sql_text
@@ -20,8 +24,6 @@ from application.storage.db.repositories.message_events import (
MessageEventsRepository,
)
from application.storage.db.session import db_readonly
from application.streaming.broadcast_channel import Topic
from application.streaming.keys import message_topic_name
logger = logging.getLogger(__name__)
@@ -81,11 +83,16 @@ def format_sse_event(payload: dict, sequence_no: int) -> str:
def _check_producer_liveness(
message_id: str, idle_seconds: float
message_id: str, user_id: Optional[str], idle_seconds: float
) -> Optional[dict]:
"""Inspect ``conversation_messages`` and return a terminal SSE
payload when the producer is no longer alive, else ``None``.
When ``user_id`` is given the lookup is scoped to ``AND user_id = :u``
(defence in depth: this long-lived re-read re-asserts the ownership the
route gated on, so a stream cannot keep tailing a row it no longer
owns). A non-matching row reads as missing → a terminal ``error``.
Three terminal cases collapse into a single DB round-trip:
- ``status='complete'`` — the live finalize ran but its journal
@@ -100,11 +107,17 @@ def _check_producer_liveness(
Synthesise ``error`` so the client doesn't hang on keepalives
until the proxy idle-timeout kicks in.
"""
owner_clause = " AND user_id = :u" if user_id is not None else ""
params = {"id": message_id, "idle_secs": float(idle_seconds)}
if user_id is not None:
params["u"] = user_id
try:
with db_readonly() as conn:
row = conn.execute(
sql_text(
"""
# ``owner_clause`` is a fixed literal (no user input in the
# SQL string); ``user_id`` is bound via ``:u``.
f"""
SELECT
status,
message_metadata->>'error' AS err,
@@ -118,10 +131,10 @@ def _check_producer_liveness(
) < now() - make_interval(secs => :idle_secs)
AS is_stale
FROM conversation_messages
WHERE id = CAST(:id AS uuid)
WHERE id = CAST(:id AS uuid){owner_clause}
"""
),
{"id": message_id, "idle_secs": float(idle_seconds)},
params,
).first()
except Exception:
logger.exception(
@@ -161,236 +174,49 @@ def _check_producer_liveness(
return None
def build_message_event_stream(
message_id: str,
last_event_id: Optional[int] = None,
*,
keepalive_seconds: float = DEFAULT_KEEPALIVE_SECONDS,
poll_timeout_seconds: float = DEFAULT_POLL_TIMEOUT_SECONDS,
watchdog_interval_seconds: float = DEFAULT_WATCHDOG_INTERVAL_SECONDS,
producer_idle_seconds: float = DEFAULT_PRODUCER_IDLE_SECONDS,
) -> Iterator[str]:
"""Yield SSE-formatted lines for one ``message_id`` reconnect stream.
def read_snapshot_lines(
message_id: str, last_event_id: Optional[int], user_id: Optional[str] = None
) -> tuple[list[str], Optional[int], bool]:
"""Read journal rows after ``last_event_id`` as SSE-formatted lines.
First frame is ``: connected``; subsequent frames are snapshot rows,
live-tail events, or ``: keepalive`` comments. Runs until the client
disconnects.
Returns ``(lines, max_sequence_no, terminal)``: ``max_sequence_no`` is
seeded with ``last_event_id`` and advanced past every row read,
``terminal`` is True if any row carried a terminal ``end``/``error``.
Raises on DB error so the caller can drive its replay-failed path.
Used by ``async_event_replay.build_message_event_stream_async`` (the
reconnect reader); it shares ``format_sse_event`` / ``_payload_is_terminal``
with the producer's journal writer so reader and writer never drift on
wire shape or terminal semantics.
"""
yield ": connected\n\n"
# Replay buffer — populated inside ``_on_subscribe`` (or the
# Redis-unavailable fallback below), drained on the first iteration
# of the subscribe loop after the callback runs.
replay_buffer: list[str] = []
# Dedup floor: seeded with the client's cursor so an empty snapshot
# still rejects re-published live events with seq <= last_event_id.
# Advanced by snapshot rows AND by yielded live events, so any
# republish past the snapshot ceiling is also dropped.
max_replayed_seq: Optional[int] = last_event_id
replay_done = False
replay_failed = False
# Set when a snapshot row carries a terminal ``end`` / ``error``
# event. After flushing the buffer the generator returns; if we
# kept tailing we'd loop on keepalives forever for a stream that
# already finished.
terminal_in_snapshot = False
def _read_snapshot_into_buffer() -> None:
nonlocal max_replayed_seq, replay_failed, terminal_in_snapshot
try:
with db_readonly() as conn:
rows = MessageEventsRepository(conn).read_after(
message_id, last_sequence_no=last_event_id
)
for row in rows:
seq = int(row["sequence_no"])
payload = row.get("payload")
if not isinstance(payload, dict):
# ``record_event`` rejects non-dict payloads at the
# write gate, so this can only be a legacy row from
# before that contract or a direct SQL insert. The
# original synthetic fallback (``{"type": event_type}``)
# used to ship a malformed envelope here — drop the
# row instead so a corrupt journal entry doesn't
# poison a reconnect.
logger.warning(
"Skipping non-dict payload from message_events: "
"message_id=%s seq=%s type=%s",
message_id,
seq,
row.get("event_type"),
)
continue
replay_buffer.append(format_sse_event(payload, seq))
if max_replayed_seq is None or seq > max_replayed_seq:
max_replayed_seq = seq
if _payload_is_terminal(payload, row.get("event_type")):
terminal_in_snapshot = True
except Exception:
logger.exception(
"Snapshot read failed for message_id=%s last_event_id=%s",
lines: list[str] = []
max_seq = last_event_id
terminal = False
with db_readonly() as conn:
rows = MessageEventsRepository(conn).read_after(
message_id, last_sequence_no=last_event_id, user_id=user_id
)
for row in rows:
seq = int(row["sequence_no"])
payload = row.get("payload")
if not isinstance(payload, dict):
# ``record_event`` rejects non-dict payloads at the write gate,
# so this is a legacy/direct-SQL row — drop it rather than ship
# a malformed envelope that would poison a reconnect.
logger.warning(
"Skipping non-dict payload from message_events: "
"message_id=%s seq=%s type=%s",
message_id,
last_event_id,
seq,
row.get("event_type"),
)
replay_failed = True
def _on_subscribe() -> None:
# SUBSCRIBE has been acked — Postgres reads from this point
# capture every row that's been committed. Pub/sub messages
# published after this point are queued at the connection level
# until the outer loop calls ``get_message`` again.
nonlocal replay_done
try:
_read_snapshot_into_buffer()
finally:
# Flip even on failure so the outer loop continues to live
# tail and the client doesn't hang waiting for a snapshot
# flush that will never come.
replay_done = True
topic = Topic(message_topic_name(message_id))
last_keepalive = time.monotonic()
# Rate-limit the watchdog's DB hit. ``-inf`` makes the first idle
# tick after replay_done fire immediately so a snapshot-already-
# terminal-in-DB case is surfaced before any keepalive cadence.
# Subsequent checks are gated by ``watchdog_interval_seconds``.
last_watchdog_check = float("-inf")
# Synthetic terminal events emitted by the watchdog use the same
# ``sequence_no=-1`` convention as the snapshot-failure path so the
# frontend's strict ``\d+`` cursor regex rejects them as a
# ``Last-Event-ID`` for any future reconnect. The chosen
# discriminator ensures a manual page refresh after a watchdog-fired
# error doesn't loop on the same synthetic id.
watchdog_synthetic_seq = -1
try:
for payload in topic.subscribe(
on_subscribe=_on_subscribe,
poll_timeout=poll_timeout_seconds,
):
# Flush snapshot exactly once after the SUBSCRIBE callback
# has run and produced a buffer.
if replay_done and replay_buffer:
for line in replay_buffer:
yield line
replay_buffer.clear()
if terminal_in_snapshot:
# The original stream already finished; tailing
# would just emit keepalives forever and pin both a
# client reconnect promise and a server WSGI thread.
return
if replay_failed:
# Snapshot read failed (DB blip / transient timeout). Emit a
# terminal ``error`` event and return — the client only
# reconnects after the original stream has already moved on,
# so without a snapshot there's nothing live left to tail and
# holding the connection open would just emit keepalives
# until the proxy idle-timeout fires. ``code`` preserves the
# snapshot-vs-agent-loop distinction so a future client can
# opt into a refetch instead of a hard failure.
yield format_sse_event(
{
"type": "error",
"error": "Stream replay failed; please refresh to load the latest state.",
"code": "snapshot_failed",
"message_id": message_id,
},
sequence_no=-1,
)
return
now = time.monotonic()
if payload is None:
# Idle tick — check both keepalive and watchdog. The
# watchdog only kicks in once the snapshot half has been
# flushed (``replay_done``) so we don't race the
# snapshot read on the first iteration.
if (
replay_done
and watchdog_interval_seconds >= 0
and now - last_watchdog_check >= watchdog_interval_seconds
):
last_watchdog_check = now
terminal_payload = _check_producer_liveness(
message_id, producer_idle_seconds
)
if terminal_payload is not None:
yield format_sse_event(
terminal_payload,
sequence_no=watchdog_synthetic_seq,
)
return
if now - last_keepalive >= keepalive_seconds:
yield ": keepalive\n\n"
last_keepalive = now
continue
envelope = _decode_pubsub_message(payload)
if envelope is None:
continue
seq = envelope.get("sequence_no")
inner = envelope.get("payload")
if (
not isinstance(seq, int)
or isinstance(seq, bool)
or not isinstance(inner, dict)
):
continue
if max_replayed_seq is not None and seq <= max_replayed_seq:
# Snapshot already covered this id — drop the duplicate.
continue
yield format_sse_event(inner, seq)
# Advance the dedup floor on the live path too, so a stale
# republish of an already-yielded seq (process restart, retry
# tool, etc.) is dropped on a later iteration.
max_replayed_seq = seq
last_keepalive = now
if _payload_is_terminal(inner, envelope.get("event_type")):
# Live tail just delivered the terminal event — close
# out the reconnect stream so the client's drain
# promise resolves and the WSGI thread is freed.
return
# Subscribe exited without ever yielding (Redis unavailable,
# ``pubsub.subscribe`` raised, or the inner loop died between
# SUBSCRIBE-ack and the first poll). The snapshot half is in
# Postgres and is still serviceable — read it directly so a
# Redis-only outage doesn't cost the client their reconnect
# backlog. Gate the read on ``replay_done`` rather than
# ``subscribe_started``: if ``_on_subscribe`` already populated
# the buffer, re-reading would append the same rows twice and
# double the answer chunks on the client (the per-message
# reconnect dispatcher does not dedup by ``id``).
if not replay_done:
_read_snapshot_into_buffer()
replay_done = True
for line in replay_buffer:
yield line
replay_buffer.clear()
if replay_failed:
# Mirror the live-tail branch: emit a terminal ``error`` so
# the frontend's existing end/error handling drives the UI
# to a failed state instead of relying on the proxy timeout.
yield format_sse_event(
{
"type": "error",
"error": "Stream replay failed; please refresh to load the latest state.",
"code": "snapshot_failed",
"message_id": message_id,
},
sequence_no=-1,
)
return
# Same close-on-terminal contract as the live-tail branch.
# Without it a Redis-down + already-completed-stream client
# would also hang on a never-ending generator.
if terminal_in_snapshot:
return
except GeneratorExit:
# Client disconnect — let the underlying ``Topic.subscribe``
# ``finally`` block tear down its pubsub cleanly.
return
continue
lines.append(format_sse_event(payload, seq))
if max_seq is None or seq > max_seq:
max_seq = seq
if _payload_is_terminal(payload, row.get("event_type")):
terminal = True
return lines, max_seq, terminal
def _decode_pubsub_message(raw) -> Optional[dict]:
@@ -1,14 +1,17 @@
"""Integration tests for the end-to-end snapshot+tail handoff.
Exercises the publisher → journal → reconnect endpoint round-trip
without mocking the journal layer, so a regression in any of:
Exercises the publisher → journal → reconnect-reader round-trip without
mocking the journal layer, so a regression in any of:
- complete_stream's _emit closure
- record_event's commit-per-call contract
- build_message_event_stream's snapshot-from-DB path
- the reconnect route's auth + ownership gates
- the shared snapshot read (``event_replay.read_snapshot_lines``)
- the reconnect route's ownership SQL (``async_sse._user_owns_message``)
- message_events repo SQL
would surface here as a failed integration assertion.
The live tail itself (pub/sub) and the full async HTTP route are covered by
``scripts/e2e_async_sse.py`` against a real Redis + uvicorn. These tests pin
the DB-backed halves against a transactional ``pg_conn`` fixture.
"""
from __future__ import annotations
@@ -56,55 +59,47 @@ def _seed_message(conn, user_id: str | None = None):
return user_id, str(msg_id)
def _emitted_ids(lines: list[str]) -> list[int]:
"""Extract the ``id:`` sequence numbers from formatted SSE frames."""
return sorted(
int(line.split("\n", 1)[0].split(": ", 1)[1])
for line in lines
if line.startswith("id: ")
)
@pytest.mark.integration
class TestSnapshotPlusTailRoundTrip:
def test_record_event_then_snapshot_returns_what_was_journaled(
self, pg_conn,
):
"""End-to-end of the journal half: ``record_event`` writes through
a real ``MessageEventsRepository``, ``build_message_event_stream``
reads the snapshot back via the same repo + a stub Topic.
"""End-to-end of the journal half: ``record_event`` writes through a
real ``MessageEventsRepository``, ``read_snapshot_lines`` reads the
snapshot back via the same repo and formats it for the wire — the
exact primitive the async reader replays through.
"""
from application.streaming import event_replay
from application.streaming.event_replay import read_snapshot_lines
from application.streaming.message_journal import record_event
_, message_id = _seed_message(pg_conn)
with _patch_journal_session(pg_conn):
# Three events stamp the journal.
record_event(message_id, 0, "answer", {"type": "answer", "answer": "A"})
record_event(message_id, 1, "answer", {"type": "answer", "answer": "B"})
record_event(message_id, 2, "end", {"type": "end"})
# Reconnect path: subscribe yields nothing (Redis-down
# branch); the post-loop fallback runs the snapshot read
# synchronously and yields the journal contents.
def _empty_subscribe(self, on_subscribe=None, poll_timeout=1.0):
return
yield # pragma: no cover
lines, max_seq, terminal = read_snapshot_lines(message_id, None)
with patch.object(
event_replay.Topic,
"subscribe",
_empty_subscribe,
create=False,
):
gen = event_replay.build_message_event_stream(
message_id,
last_event_id=None,
keepalive_seconds=0.05,
poll_timeout_seconds=0.01,
)
out = list(gen)
# Prelude + 3 snapshot frames.
assert out[0] == ": connected\n\n"
assert "id: 0" in out[1] and '"answer": "A"' in out[1]
assert "id: 1" in out[2] and '"answer": "B"' in out[2]
assert "id: 2" in out[3] and '"type": "end"' in out[3]
assert len(lines) == 3
assert "id: 0" in lines[0] and '"answer": "A"' in lines[0]
assert "id: 1" in lines[1] and '"answer": "B"' in lines[1]
assert "id: 2" in lines[2] and '"type": "end"' in lines[2]
assert max_seq == 2
# A terminal event in the snapshot tells the reader to close.
assert terminal is True
def test_snapshot_resumes_past_last_event_id(self, pg_conn):
from application.streaming import event_replay
from application.streaming.event_replay import read_snapshot_lines
from application.streaming.message_journal import record_event
_, message_id = _seed_message(pg_conn)
@@ -115,119 +110,81 @@ class TestSnapshotPlusTailRoundTrip:
message_id, seq, "answer", {"type": "answer", "answer": str(seq)}
)
def _empty_subscribe(self, on_subscribe=None, poll_timeout=1.0):
return
yield # pragma: no cover
# Client says it has seen up through seq=2; expect 3 + 4.
lines, max_seq, terminal = read_snapshot_lines(message_id, 2)
with patch.object(
event_replay.Topic,
"subscribe",
_empty_subscribe,
create=False,
):
# Client says it has seen up through seq=2; expect 3 + 4.
out = list(
event_replay.build_message_event_stream(
message_id,
last_event_id=2,
keepalive_seconds=0.05,
poll_timeout_seconds=0.01,
)
)
assert _emitted_ids(lines) == [3, 4]
assert max_seq == 4
assert terminal is False
ids_seen = [line for line in out if line.startswith("id: ")]
# Multi-line records: extract the id integers we delivered.
emitted = sorted(
int(line.split(": ", 1)[1].split("\n")[0])
for line in out
if line.startswith("id: ")
)
# Filter for non-negative (the snapshot-failure synthetic uses -1).
emitted = [e for e in emitted if e >= 0]
assert emitted == [3, 4]
assert ids_seen # sanity
def test_reconnect_route_round_trip(self, pg_conn, flask_app):
"""``/api/messages/<id>/events`` returns the journaled events
for an authenticated owner.
def test_ownership_sql_accepts_owner_and_rejects_others(self, pg_conn):
"""The async route's ownership gate runs real SQL against
``conversation_messages`` — the owner passes, everyone else 404s.
"""
from flask import Flask, request
from application.api import async_sse
from application.api.answer.routes.messages import messages_bp
user_id, message_id = _seed_message(pg_conn)
@contextmanager
def _yield():
yield pg_conn
with patch("application.api.async_sse.db_readonly", _yield):
assert async_sse._user_owns_message(message_id, user_id) is True
assert async_sse._user_owns_message(message_id, "different-user") is False
# A well-formed but unknown id is also not owned.
assert async_sse._user_owns_message(str(_uuid.uuid4()), user_id) is False
def test_snapshot_read_is_user_scoped(self, pg_conn):
"""``read_snapshot_lines(..., user_id=)`` re-asserts ownership at the
data layer: the owner gets the journal rows, a non-owner gets none.
"""
from application.streaming.event_replay import read_snapshot_lines
from application.streaming.message_journal import record_event
# Build a fresh Flask app routing to the reconnect blueprint
# plus a tiny auth shim that injects the test user.
user_id, message_id = _seed_message(pg_conn)
app = Flask(__name__)
app.register_blueprint(messages_bp)
app.config["TESTING"] = True
@app.before_request
def _shim_auth():
request.decoded_token = {"sub": user_id}
with _patch_journal_session(pg_conn):
record_event(message_id, 0, "answer", {"type": "answer", "answer": "x"})
record_event(message_id, 1, "end", {"type": "end"})
from application.streaming import event_replay
owner_lines, _, owner_terminal = read_snapshot_lines(
message_id, None, user_id
)
other_lines, _, other_terminal = read_snapshot_lines(
message_id, None, "different-user"
)
# Unscoped (user_id=None) still returns everything.
unscoped_lines, _, _ = read_snapshot_lines(message_id, None)
def _empty_subscribe(self, on_subscribe=None, poll_timeout=1.0):
return
yield # pragma: no cover
assert _emitted_ids(owner_lines) == [0, 1] and owner_terminal is True
assert other_lines == [] and other_terminal is False
assert _emitted_ids(unscoped_lines) == [0, 1]
with patch.object(
event_replay.Topic,
"subscribe",
_empty_subscribe,
create=False,
), patch(
"application.api.answer.routes.messages.db_readonly"
) as ro:
ro.return_value.__enter__.return_value = pg_conn
with app.test_client() as c:
r = c.get(f"/api/messages/{message_id}/events")
assert r.status_code == 200
body = b""
for chunk in r.iter_encoded():
body += chunk
if body.count(b"\n\n") >= 4:
break
r.close()
# Both journaled events present in the response.
text = body.decode("utf-8")
assert ": connected" in text
assert '"answer": "x"' in text
assert '"type": "end"' in text
# The seq lines are correct.
assert "id: 0" in text and "id: 1" in text
def test_reconnect_rejects_non_owner(self, pg_conn, flask_app):
from flask import Flask, request
from application.api.answer.routes.messages import messages_bp
def test_watchdog_is_user_scoped(self, pg_conn):
"""``_check_producer_liveness(..., user_id)`` only sees the caller's
own row; a non-owner reads as missing (terminal), not as the row's
real status.
"""
from application.streaming.event_replay import _check_producer_liveness
# Seed a row and flip it to a terminal status the watchdog reports.
user_id, message_id = _seed_message(pg_conn)
app = Flask(__name__)
app.register_blueprint(messages_bp)
pg_conn.execute(
sql_text(
"UPDATE conversation_messages SET status='complete' WHERE id = :id"
),
{"id": message_id},
)
@app.before_request
def _shim_auth():
request.decoded_token = {"sub": "different-user"}
@contextmanager
def _yield():
yield pg_conn
# Make the ownership check use the test connection.
with patch(
"application.api.answer.routes.messages.db_readonly"
) as ro:
from contextlib import contextmanager as _cm
with patch("application.streaming.event_replay.db_readonly", _yield):
owner = _check_producer_liveness(message_id, user_id, 90.0)
other = _check_producer_liveness(message_id, "different-user", 90.0)
@_cm
def _yield():
yield pg_conn
ro.side_effect = lambda: _yield()
with app.test_client() as c:
r = c.get(f"/api/messages/{message_id}/events")
assert r.status_code == 404
# Owner sees the real terminal state; non-owner sees "missing".
assert owner == {"type": "end"}
assert other is not None and other.get("code") == "message_missing"
+252
View File
@@ -0,0 +1,252 @@
"""Tests for ``application/api/async_sse.py``.
Native-async reconnect endpoint: GET /api/messages/<id>/events. Auth gate,
ownership gate, malformed-id rejection, Last-Event-ID normalisation, and the
SSE response shape (headers + ``: connected`` prelude). The route is a
Starlette endpoint, so it's driven through Starlette's TestClient over a
minimal app built from ``async_sse_routes``.
"""
from __future__ import annotations
from unittest.mock import AsyncMock, patch
import pytest
from starlette.applications import Starlette
from starlette.testclient import TestClient
from application.api.async_sse import (
_MESSAGE_ID_RE,
_normalise_last_event_id,
async_sse_routes,
)
from application.core.settings import settings
VALID_UUID = "67d65e8f-e7fb-4df1-9e6e-99ea6c830206"
_AUTH = "application.api.async_sse.handle_auth"
_OWNS = "application.api.async_sse._user_owns_message"
_STREAM = "application.api.async_sse.build_message_event_stream_async"
_AREDIS = "application.api.async_sse.get_async_redis_instance"
def _client() -> TestClient:
return TestClient(Starlette(routes=async_sse_routes))
@pytest.fixture(autouse=True)
def _no_redis_by_default():
"""Disable the per-user cap (no Redis) so gate tests stay hermetic.
Cap-specific tests override this with their own mock Redis.
"""
with patch(_AREDIS, AsyncMock(return_value=None)):
yield
def _mock_redis(incr_value: int) -> AsyncMock:
redis = AsyncMock()
redis.incr = AsyncMock(return_value=incr_value)
redis.expire = AsyncMock(return_value=True)
redis.decr = AsyncMock(return_value=incr_value - 1)
return redis
def _fake_stream(record: dict | None = None):
"""Async builder stub: records the cursor, yields the prelude, returns."""
async def _gen(message_id, last_event_id=None, **kwargs):
if record is not None:
record["message_id"] = message_id
record["last_event_id"] = last_event_id
yield ": connected\n\n"
return _gen
# ── pure helpers ─────────────────────────────────────────────────────────
@pytest.mark.unit
class TestNormaliseLastEventId:
def test_none_passthrough(self):
assert _normalise_last_event_id(None) is None
def test_empty_string(self):
assert _normalise_last_event_id("") is None
def test_whitespace_only(self):
assert _normalise_last_event_id(" ") is None
def test_valid_int(self):
assert _normalise_last_event_id("42") == 42
def test_stripped_whitespace(self):
assert _normalise_last_event_id(" 7 ") == 7
def test_zero_is_valid(self):
assert _normalise_last_event_id("0") == 0
def test_negative_rejected(self):
# We expose only non-negative cursors; -1 is reserved for the
# snapshot-failure synthetic terminal event.
assert _normalise_last_event_id("-1") is None
def test_non_numeric_rejected(self):
for bad in ("foo", "1.5", "1e3", "abc-123", "null"):
assert _normalise_last_event_id(bad) is None, bad
@pytest.mark.unit
class TestMessageIdRegex:
def test_canonical_uuid_accepted(self):
assert _MESSAGE_ID_RE.match(VALID_UUID)
def test_uppercase_uuid_accepted(self):
assert _MESSAGE_ID_RE.match(VALID_UUID.upper())
def test_no_dashes_rejected(self):
assert not _MESSAGE_ID_RE.match(VALID_UUID.replace("-", ""))
def test_legacy_mongo_id_rejected(self):
assert not _MESSAGE_ID_RE.match("507f1f77bcf86cd799439011")
# ── auth gate ──────────────────────────────────────────────────────────────
@pytest.mark.unit
class TestAuthGate:
def test_401_when_handle_auth_returns_none(self):
with patch(_AUTH, return_value=None):
r = _client().get(f"/api/messages/{VALID_UUID}/events")
assert r.status_code == 401
def test_401_when_decoded_token_missing_sub(self):
with patch(_AUTH, return_value={"email": "x@y"}):
r = _client().get(f"/api/messages/{VALID_UUID}/events")
assert r.status_code == 401
def test_401_when_handle_auth_returns_error(self):
with patch(_AUTH, return_value={"error": "invalid_token"}):
r = _client().get(f"/api/messages/{VALID_UUID}/events")
assert r.status_code == 401
# ── message-id validation ───────────────────────────────────────────────────
@pytest.mark.unit
class TestMessageIdValidation:
def test_400_on_malformed_id(self):
with patch(_AUTH, return_value={"sub": "alice"}):
r = _client().get("/api/messages/not-a-uuid/events")
assert r.status_code == 400
# ── ownership gate ──────────────────────────────────────────────────────────
@pytest.mark.unit
class TestOwnershipGate:
def test_404_when_user_does_not_own_message(self):
with patch(_AUTH, return_value={"sub": "alice"}), patch(
_OWNS, return_value=False
):
r = _client().get(f"/api/messages/{VALID_UUID}/events")
assert r.status_code == 404
def test_200_when_user_owns_message(self):
with patch(_AUTH, return_value={"sub": "alice"}), patch(
_OWNS, return_value=True
), patch(_STREAM, _fake_stream()):
r = _client().get(f"/api/messages/{VALID_UUID}/events")
assert r.status_code == 200
assert r.headers.get("content-type", "").startswith("text/event-stream")
assert r.headers.get("Cache-Control") == "no-store"
assert r.headers.get("X-Accel-Buffering") == "no"
assert r.headers.get("X-SSE-Transport") == "async"
assert ": connected" in r.text
# ── Last-Event-ID parsing ───────────────────────────────────────────────────
@pytest.mark.unit
class TestLastEventIdParsing:
def test_header_passes_through_to_builder(self):
captured: dict = {}
with patch(_AUTH, return_value={"sub": "alice"}), patch(
_OWNS, return_value=True
), patch(_STREAM, _fake_stream(captured)):
_client().get(
f"/api/messages/{VALID_UUID}/events",
headers={"Last-Event-ID": "12"},
)
assert captured["message_id"] == VALID_UUID
assert captured["last_event_id"] == 12
def test_query_param_fallback(self):
captured: dict = {}
with patch(_AUTH, return_value={"sub": "alice"}), patch(
_OWNS, return_value=True
), patch(_STREAM, _fake_stream(captured)):
_client().get(f"/api/messages/{VALID_UUID}/events?last_event_id=5")
assert captured["last_event_id"] == 5
def test_invalid_cursor_normalised_to_none(self):
captured: dict = {}
with patch(_AUTH, return_value={"sub": "alice"}), patch(
_OWNS, return_value=True
), patch(_STREAM, _fake_stream(captured)):
_client().get(
f"/api/messages/{VALID_UUID}/events",
headers={"Last-Event-ID": "definitely-not-a-number"},
)
assert captured["last_event_id"] is None
# ── per-user concurrent-connection cap ──────────────────────────────────────
@pytest.mark.unit
class TestConnectionCap:
def _cap(self) -> int:
return int(settings.SSE_MAX_CONCURRENT_PER_USER) or 8
def test_429_when_over_cap(self):
cap = self._cap()
redis = _mock_redis(incr_value=cap + 1) # post-incr count exceeds cap
with patch(_AUTH, return_value={"sub": "alice"}), patch(
_OWNS, return_value=True
), patch(_STREAM, _fake_stream()), patch(
_AREDIS, AsyncMock(return_value=redis)
):
r = _client().get(f"/api/messages/{VALID_UUID}/events")
assert r.status_code == 429
# The increment is rolled back so a rejected attempt doesn't wedge
# the counter at the cap forever.
redis.decr.assert_awaited_once()
def test_200_and_slot_released_when_under_cap(self):
redis = _mock_redis(incr_value=1)
with patch(_AUTH, return_value={"sub": "alice"}), patch(
_OWNS, return_value=True
), patch(_STREAM, _fake_stream()), patch(
_AREDIS, AsyncMock(return_value=redis)
):
r = _client().get(f"/api/messages/{VALID_UUID}/events")
assert r.status_code == 200
assert ": connected" in r.text
redis.incr.assert_awaited_once()
# Slot released when the stream finishes (terminal/close).
redis.decr.assert_awaited_once()
def test_cap_skipped_when_redis_unavailable(self):
# The autouse fixture already makes get_async_redis_instance -> None;
# the stream is served (fail-open, like /api/events).
with patch(_AUTH, return_value={"sub": "alice"}), patch(
_OWNS, return_value=True
), patch(_STREAM, _fake_stream()):
r = _client().get(f"/api/messages/{VALID_UUID}/events")
assert r.status_code == 200
-220
View File
@@ -1,220 +0,0 @@
"""Tests for ``application/api/answer/routes/messages.py``.
Reconnect endpoint: GET /api/messages/<id>/events. Auth gate, ownership
gate, malformed-id rejection, Last-Event-ID normalisation, and a smoke
test that the SSE response shape matches the user-events endpoint.
"""
from __future__ import annotations
from unittest.mock import patch
from flask import Flask, request
from application.api.answer.routes.messages import (
_MESSAGE_ID_RE,
_normalise_last_event_id,
messages_bp,
)
def _make_app(decoded_token=None):
app = Flask(__name__)
app.register_blueprint(messages_bp)
app.config["TESTING"] = True
@app.before_request
def _shim_auth():
request.decoded_token = decoded_token
return app
VALID_UUID = "67d65e8f-e7fb-4df1-9e6e-99ea6c830206"
class TestNormaliseLastEventId:
def test_none_passthrough(self):
assert _normalise_last_event_id(None) is None
def test_empty_string(self):
assert _normalise_last_event_id("") is None
def test_whitespace_only(self):
assert _normalise_last_event_id(" ") is None
def test_valid_int(self):
assert _normalise_last_event_id("42") == 42
def test_stripped_whitespace(self):
assert _normalise_last_event_id(" 7 ") == 7
def test_zero_is_valid(self):
assert _normalise_last_event_id("0") == 0
def test_negative_rejected(self):
# We expose only non-negative cursors; -1 is reserved for the
# snapshot-failure synthetic terminal event and shouldn't
# round-trip back.
assert _normalise_last_event_id("-1") is None
def test_non_numeric_rejected(self):
for bad in ("foo", "1.5", "1e3", "abc-123", "null"):
assert _normalise_last_event_id(bad) is None, bad
class TestMessageIdRegex:
def test_canonical_uuid_accepted(self):
assert _MESSAGE_ID_RE.match(VALID_UUID)
def test_uppercase_uuid_accepted(self):
assert _MESSAGE_ID_RE.match(VALID_UUID.upper())
def test_no_dashes_rejected(self):
assert not _MESSAGE_ID_RE.match(VALID_UUID.replace("-", ""))
def test_legacy_mongo_id_rejected(self):
# 24-char hex with no dashes — a Mongo objectid-shaped string
# that happened to leak through somewhere.
assert not _MESSAGE_ID_RE.match("507f1f77bcf86cd799439011")
class TestAuthGate:
def test_401_when_no_decoded_token(self):
app = _make_app(decoded_token=None)
with app.test_client() as c:
r = c.get(f"/api/messages/{VALID_UUID}/events")
assert r.status_code == 401
def test_401_when_decoded_token_missing_sub(self):
app = _make_app(decoded_token={"email": "x@y"})
with app.test_client() as c:
r = c.get(f"/api/messages/{VALID_UUID}/events")
assert r.status_code == 401
class TestMessageIdValidation:
def test_400_on_malformed_id(self):
app = _make_app(decoded_token={"sub": "alice"})
with app.test_client() as c:
r = c.get("/api/messages/not-a-uuid/events")
assert r.status_code == 400
class TestOwnershipGate:
def test_404_when_user_does_not_own_message(self):
from application.api.answer.routes import messages as messages_module
app = _make_app(decoded_token={"sub": "alice"})
with patch.object(
messages_module, "_user_owns_message", return_value=False
):
with app.test_client() as c:
r = c.get(f"/api/messages/{VALID_UUID}/events")
assert r.status_code == 404
def test_200_when_user_owns_message(self):
from application.api.answer.routes import messages as messages_module
app = _make_app(decoded_token={"sub": "alice"})
# Have build_message_event_stream yield just the prelude then
# exit so the test can drain the response without blocking on
# a live pubsub subscription.
def _fake_builder(message_id, last_event_id=None, **kwargs):
yield ": connected\n\n"
with patch.object(
messages_module, "_user_owns_message", return_value=True
), patch.object(
messages_module, "build_message_event_stream", _fake_builder
):
with app.test_client() as c:
r = c.get(f"/api/messages/{VALID_UUID}/events")
assert r.status_code == 200
assert r.mimetype == "text/event-stream"
assert r.headers.get("Cache-Control") == "no-store"
assert r.headers.get("X-Accel-Buffering") == "no"
body = b""
for chunk in r.iter_encoded():
body += chunk
if b": connected" in body:
break
r.close()
assert b": connected" in body
class TestLastEventIdParsing:
def test_header_passes_through_to_builder(self):
from application.api.answer.routes import messages as messages_module
captured = {}
def _fake_builder(message_id, last_event_id=None, **kwargs):
captured["message_id"] = message_id
captured["last_event_id"] = last_event_id
yield ": connected\n\n"
app = _make_app(decoded_token={"sub": "alice"})
with patch.object(
messages_module, "_user_owns_message", return_value=True
), patch.object(
messages_module, "build_message_event_stream", _fake_builder
):
with app.test_client() as c:
r = c.get(
f"/api/messages/{VALID_UUID}/events",
headers={"Last-Event-ID": "12"},
)
# Drain a tick.
next(iter(r.iter_encoded()), None)
r.close()
assert captured["message_id"] == VALID_UUID
assert captured["last_event_id"] == 12
def test_query_param_fallback(self):
from application.api.answer.routes import messages as messages_module
captured = {}
def _fake_builder(message_id, last_event_id=None, **kwargs):
captured["last_event_id"] = last_event_id
yield ": connected\n\n"
app = _make_app(decoded_token={"sub": "alice"})
with patch.object(
messages_module, "_user_owns_message", return_value=True
), patch.object(
messages_module, "build_message_event_stream", _fake_builder
):
with app.test_client() as c:
r = c.get(
f"/api/messages/{VALID_UUID}/events?last_event_id=5"
)
next(iter(r.iter_encoded()), None)
r.close()
assert captured["last_event_id"] == 5
def test_invalid_cursor_normalised_to_none(self):
from application.api.answer.routes import messages as messages_module
captured = {}
def _fake_builder(message_id, last_event_id=None, **kwargs):
captured["last_event_id"] = last_event_id
yield ": connected\n\n"
app = _make_app(decoded_token={"sub": "alice"})
with patch.object(
messages_module, "_user_owns_message", return_value=True
), patch.object(
messages_module, "build_message_event_stream", _fake_builder
):
with app.test_client() as c:
r = c.get(
f"/api/messages/{VALID_UUID}/events",
headers={"Last-Event-ID": "definitely-not-a-number"},
)
next(iter(r.iter_encoded()), None)
r.close()
assert captured["last_event_id"] is None
+194 -320
View File
@@ -1,37 +1,48 @@
"""Unit tests for ``application/streaming/event_replay.py``.
"""Unit tests for the chat-stream reconnect snapshot+tail boundary.
The replay generator is the snapshot+tail boundary the chat-stream
reconnect path lives on. The boundary correctness invariants worth
locking down:
``event_replay`` holds the shared leaf primitives (SSE wire format, pub/sub
envelope encode/decode, snapshot read, watchdog probe); the reader itself is
the async generator ``async_event_replay.build_message_event_stream_async``.
Boundary correctness invariants worth locking down:
- Snapshot replay yields rows in ``sequence_no`` order with the SSE
``id:`` header set to that sequence_no.
- Snapshot replay yields rows in ``sequence_no`` order with the SSE ``id:``
header set to that sequence_no.
- Live tail dedupes pub/sub messages whose ``sequence_no`` is already
covered by the snapshot.
- A backlog read failure inside ``on_subscribe`` doesn't wedge the
generator — it surfaces a terminal ``error`` event (``code:
snapshot_failed``) and returns, freeing the WSGI thread instead of
pinning it on keepalives the client can't hear as terminal.
snapshot_failed``) and returns.
- Keepalive comments fire after the configured silence window.
- ``encode_pubsub_message`` round-trips with ``_decode_pubsub_message``.
The snapshot read and watchdog probe run via ``anyio.to_thread`` inside the
async generator but are the same ``event_replay`` functions, so tests patch
them at ``application.streaming.event_replay.*`` exactly as before.
"""
from __future__ import annotations
import inspect
import json
from typing import Iterator
from typing import AsyncIterator
from unittest.mock import MagicMock, patch
import pytest
from application.streaming.async_event_replay import (
build_message_event_stream_async,
)
from application.streaming.event_replay import (
_SSE_LINE_SPLIT_PATTERN,
_decode_pubsub_message,
build_message_event_stream,
encode_pubsub_message,
format_sse_event,
)
_ASYNC_TOPIC = "application.streaming.async_event_replay.AsyncTopic.subscribe"
_READONLY = "application.streaming.event_replay.db_readonly"
_REPO = "application.streaming.event_replay.MessageEventsRepository"
# ── format_sse_event ────────────────────────────────────────────────────
@@ -91,19 +102,20 @@ class TestPubsubEnvelope:
assert envelope == {"sequence_no": 1, "payload": {}}
# ── build_message_event_stream ──────────────────────────────────────────
# ── build_message_event_stream_async ────────────────────────────────────
def _fake_topic_subscribe(messages: list, *, fire_callback: bool = True):
"""Build a ``Topic.subscribe`` mock that fires ``on_subscribe`` then
yields the supplied bytes in order, then yields ``None`` ticks
indefinitely so the generator can keep running until the test
closes it.
def _fake_subscribe(messages: list, *, fire_callback: bool = True):
"""Build an ``AsyncTopic.subscribe`` mock that fires ``on_subscribe``
then yields the supplied bytes in order, then yields ``None`` ticks
indefinitely so the generator keeps running until the test closes it.
"""
def _impl(self, on_subscribe=None, poll_timeout=1.0):
async def _impl(self, on_subscribe=None, poll_timeout=1.0):
if fire_callback and on_subscribe is not None:
on_subscribe()
res = on_subscribe()
if inspect.isawaitable(res):
await res
for m in messages:
yield m
while True:
@@ -112,43 +124,57 @@ def _fake_topic_subscribe(messages: list, *, fire_callback: bool = True):
return _impl
def _drain(gen: Iterator[str], *, max_items: int = 50) -> list[str]:
out = []
def _subscribe_returns_immediately(*, fire_callback: bool):
"""``AsyncTopic.subscribe`` that exits without yielding (Redis-down).
With ``fire_callback`` it runs ``on_subscribe`` first (the
SUBSCRIBE-ack-then-get_message-dies race); without it, nothing runs
(subscribe itself failed), exercising the post-loop fallback read.
"""
async def _impl(self, on_subscribe=None, poll_timeout=1.0):
if fire_callback and on_subscribe is not None:
res = on_subscribe()
if inspect.isawaitable(res):
await res
return
yield # pragma: no cover — make the function an async generator
return _impl
async def _drain(agen: AsyncIterator[str], *, max_items: int = 50) -> list[str]:
out: list[str] = []
try:
for _ in range(max_items):
out.append(next(gen))
except StopIteration:
out.append(await agen.__anext__())
except StopAsyncIteration:
pass
finally:
gen.close()
await agen.aclose()
return out
@pytest.mark.unit
@pytest.mark.asyncio
class TestBuildMessageEventStream:
def test_yields_connected_prelude_first(self):
with patch(
"application.streaming.event_replay.db_readonly"
) as mock_readonly, patch(
"application.streaming.event_replay.MessageEventsRepository"
) as mock_repo_cls, patch(
"application.streaming.event_replay.Topic.subscribe",
_fake_topic_subscribe([]),
create=False,
async def test_yields_connected_prelude_first(self):
with patch(_READONLY) as mock_readonly, patch(_REPO) as mock_repo_cls, patch(
_ASYNC_TOPIC, _fake_subscribe([])
):
mock_readonly.return_value.__enter__.return_value = MagicMock()
mock_repo_cls.return_value.read_after.return_value = []
gen = build_message_event_stream(
gen = build_message_event_stream_async(
"msg-1",
last_event_id=None,
keepalive_seconds=0.05,
poll_timeout_seconds=0.01,
)
first = next(gen)
gen.close()
first = await gen.__anext__()
await gen.aclose()
assert first == ": connected\n\n"
def test_snapshot_replays_in_sequence_order(self):
async def test_snapshot_replays_in_sequence_order(self):
rows = [
{
"sequence_no": 0,
@@ -161,39 +187,27 @@ class TestBuildMessageEventStream:
"payload": {"type": "answer", "answer": "B"},
},
]
with patch(
"application.streaming.event_replay.db_readonly"
) as mock_readonly, patch(
"application.streaming.event_replay.MessageEventsRepository"
) as mock_repo_cls, patch(
"application.streaming.event_replay.Topic.subscribe",
_fake_topic_subscribe([]),
create=False,
with patch(_READONLY) as mock_readonly, patch(_REPO) as mock_repo_cls, patch(
_ASYNC_TOPIC, _fake_subscribe([])
):
mock_readonly.return_value.__enter__.return_value = MagicMock()
mock_repo_cls.return_value.read_after.return_value = rows
gen = build_message_event_stream(
gen = build_message_event_stream_async(
"msg-1",
last_event_id=None,
keepalive_seconds=0.05,
poll_timeout_seconds=0.01,
)
out = _drain(gen, max_items=4)
out = await _drain(gen, max_items=4)
# Expect: prelude, two snapshot frames, then keepalive ticks (None
# not yielded as a string — generator yields keepalive comments
# instead). With poll_timeout=1.0 and short test window we may
# only see the prelude + snapshot before close.
assert out[0] == ": connected\n\n"
assert "id: 0" in out[1]
assert "id: 1" in out[2]
# The repo was queried with the right cursor.
mock_repo_cls.return_value.read_after.assert_called_once_with(
"msg-1", last_sequence_no=None
"msg-1", last_sequence_no=None, user_id=None
)
def test_live_tail_dedupes_against_snapshot(self):
async def test_live_tail_dedupes_against_snapshot(self):
snapshot_rows = [
{
"sequence_no": 5,
@@ -208,25 +222,18 @@ class TestBuildMessageEventStream:
"msg-1", 6, "answer", {"type": "answer", "answer": "fresh"}
).encode("utf-8")
with patch(
"application.streaming.event_replay.db_readonly"
) as mock_readonly, patch(
"application.streaming.event_replay.MessageEventsRepository"
) as mock_repo_cls, patch(
"application.streaming.event_replay.Topic.subscribe",
_fake_topic_subscribe([live_envelope, live_new_envelope]),
create=False,
with patch(_READONLY) as mock_readonly, patch(_REPO) as mock_repo_cls, patch(
_ASYNC_TOPIC, _fake_subscribe([live_envelope, live_new_envelope])
):
mock_readonly.return_value.__enter__.return_value = MagicMock()
mock_repo_cls.return_value.read_after.return_value = snapshot_rows
gen = build_message_event_stream(
gen = build_message_event_stream_async(
"msg-1",
last_event_id=None,
keepalive_seconds=0.05,
poll_timeout_seconds=0.01,
)
out = _drain(gen, max_items=4)
out = await _drain(gen, max_items=4)
assert out[0] == ": connected\n\n"
assert "id: 5" in out[1]
@@ -234,51 +241,36 @@ class TestBuildMessageEventStream:
# The duplicate live event (seq=5) is dropped.
assert "id: 6" in out[2]
assert '"answer": "fresh"' in out[2]
# No third frame past the fresh one (only None ticks).
def test_live_tail_passes_through_when_seq_strictly_greater_than_replay(self):
async def test_live_tail_passes_through_when_seq_strictly_greater_than_replay(self):
"""No snapshot rows; every live event is fresh and yielded."""
live = encode_pubsub_message(
"msg-1", 0, "answer", {"type": "answer", "answer": "x"}
).encode("utf-8")
with patch(
"application.streaming.event_replay.db_readonly"
) as mock_readonly, patch(
"application.streaming.event_replay.MessageEventsRepository"
) as mock_repo_cls, patch(
"application.streaming.event_replay.Topic.subscribe",
_fake_topic_subscribe([live]),
create=False,
with patch(_READONLY) as mock_readonly, patch(_REPO) as mock_repo_cls, patch(
_ASYNC_TOPIC, _fake_subscribe([live])
):
mock_readonly.return_value.__enter__.return_value = MagicMock()
mock_repo_cls.return_value.read_after.return_value = []
gen = build_message_event_stream(
gen = build_message_event_stream_async(
"msg-1",
last_event_id=None,
keepalive_seconds=0.05,
poll_timeout_seconds=0.01,
)
out = _drain(gen, max_items=3)
out = await _drain(gen, max_items=3)
assert out[0] == ": connected\n\n"
assert "id: 0" in out[1]
assert '"answer": "x"' in out[1]
def test_snapshot_read_failure_surfaces_synthetic_event(self):
with patch(
"application.streaming.event_replay.db_readonly"
) as mock_readonly, patch(
"application.streaming.event_replay.MessageEventsRepository"
) as mock_repo_cls, patch(
"application.streaming.event_replay.Topic.subscribe",
_fake_topic_subscribe([]),
create=False,
async def test_snapshot_read_failure_surfaces_synthetic_event(self):
with patch(_READONLY) as mock_readonly, patch(_REPO) as mock_repo_cls, patch(
_ASYNC_TOPIC, _fake_subscribe([])
):
mock_readonly.return_value.__enter__.return_value = MagicMock()
mock_repo_cls.return_value.read_after.side_effect = RuntimeError("boom")
gen = build_message_event_stream(
gen = build_message_event_stream_async(
"msg-1",
last_event_id=None,
keepalive_seconds=0.05,
@@ -286,125 +278,101 @@ class TestBuildMessageEventStream:
)
# Ask for more than the expected output so we'd notice a
# regression where the generator keeps emitting keepalives.
out = _drain(gen, max_items=10)
out = await _drain(gen, max_items=10)
assert out[0] == ": connected\n\n"
# Terminal ``error`` event with the snapshot-specific ``code`` so
# the frontend's existing end/error contract closes the stream.
assert '"type": "error"' in out[1]
assert '"code": "snapshot_failed"' in out[1]
assert "id: -1" in out[1]
# Generator must return after the synthetic — otherwise the
# client hangs on keepalives waiting for a terminal event that
# would never arrive.
# Generator must return after the synthetic.
assert len(out) == 2
def test_malformed_pubsub_message_dropped_silently(self):
async def test_malformed_pubsub_message_dropped_silently(self):
bad = b"not-json"
good = encode_pubsub_message(
"msg-1", 0, "answer", {"type": "answer", "answer": "ok"}
).encode("utf-8")
with patch(
"application.streaming.event_replay.db_readonly"
) as mock_readonly, patch(
"application.streaming.event_replay.MessageEventsRepository"
) as mock_repo_cls, patch(
"application.streaming.event_replay.Topic.subscribe",
_fake_topic_subscribe([bad, good]),
create=False,
with patch(_READONLY) as mock_readonly, patch(_REPO) as mock_repo_cls, patch(
_ASYNC_TOPIC, _fake_subscribe([bad, good])
):
mock_readonly.return_value.__enter__.return_value = MagicMock()
mock_repo_cls.return_value.read_after.return_value = []
gen = build_message_event_stream(
gen = build_message_event_stream_async(
"msg-1",
last_event_id=None,
keepalive_seconds=0.05,
poll_timeout_seconds=0.01,
)
out = _drain(gen, max_items=3)
out = await _drain(gen, max_items=3)
assert out[0] == ": connected\n\n"
# Bad message dropped; good one yielded.
assert "id: 0" in out[1]
assert '"answer": "ok"' in out[1]
def test_pubsub_envelope_with_non_int_sequence_dropped(self):
async def test_pubsub_envelope_with_non_int_sequence_dropped(self):
bad = json.dumps(
{"sequence_no": "not-int", "payload": {"type": "x"}}
).encode("utf-8")
good = encode_pubsub_message(
"msg-1", 0, "answer", {"type": "answer", "answer": "ok"}
).encode("utf-8")
with patch(
"application.streaming.event_replay.db_readonly"
) as mock_readonly, patch(
"application.streaming.event_replay.MessageEventsRepository"
) as mock_repo_cls, patch(
"application.streaming.event_replay.Topic.subscribe",
_fake_topic_subscribe([bad, good]),
create=False,
with patch(_READONLY) as mock_readonly, patch(_REPO) as mock_repo_cls, patch(
_ASYNC_TOPIC, _fake_subscribe([bad, good])
):
mock_readonly.return_value.__enter__.return_value = MagicMock()
mock_repo_cls.return_value.read_after.return_value = []
gen = build_message_event_stream(
gen = build_message_event_stream_async(
"msg-1",
last_event_id=None,
keepalive_seconds=0.05,
poll_timeout_seconds=0.01,
)
out = _drain(gen, max_items=3)
out = await _drain(gen, max_items=3)
assert out[0] == ": connected\n\n"
assert "id: 0" in out[1]
@pytest.mark.unit
@pytest.mark.asyncio
class TestDedupFloorSeededFromCursor:
"""Regressions for the dedup-floor bugs the round-1 review flagged.
With ``max_replayed_seq`` initialised to ``last_event_id``, an
empty snapshot still rejects republished live events the client
has already seen. Advancing on yield protects against republish
past the snapshot ceiling.
With ``max_replayed_seq`` initialised to ``last_event_id``, an empty
snapshot still rejects republished live events the client has already
seen. Advancing on yield protects against republish past the snapshot
ceiling.
"""
def test_empty_snapshot_dedups_against_last_event_id(self):
# No snapshot rows; live event with seq=3 (already seen by
# client at last_event_id=5) must be dropped.
async def test_empty_snapshot_dedups_against_last_event_id(self):
# No snapshot rows; live event with seq=3 (already seen by client
# at last_event_id=5) must be dropped.
live_dup = encode_pubsub_message(
"msg-1", 3, "answer", {"type": "answer", "answer": "stale"}
).encode("utf-8")
live_fresh = encode_pubsub_message(
"msg-1", 6, "answer", {"type": "answer", "answer": "fresh"}
).encode("utf-8")
with patch(
"application.streaming.event_replay.db_readonly"
) as mock_readonly, patch(
"application.streaming.event_replay.MessageEventsRepository"
) as mock_repo_cls, patch(
"application.streaming.event_replay.Topic.subscribe",
_fake_topic_subscribe([live_dup, live_fresh]),
create=False,
with patch(_READONLY) as mock_readonly, patch(_REPO) as mock_repo_cls, patch(
_ASYNC_TOPIC, _fake_subscribe([live_dup, live_fresh])
):
mock_readonly.return_value.__enter__.return_value = MagicMock()
mock_repo_cls.return_value.read_after.return_value = []
gen = build_message_event_stream(
gen = build_message_event_stream_async(
"msg-1",
last_event_id=5,
keepalive_seconds=0.05,
poll_timeout_seconds=0.01,
)
out = _drain(gen, max_items=3)
out = await _drain(gen, max_items=3)
assert out[0] == ": connected\n\n"
# Stale live event dropped; only the fresh one yielded.
assert "id: 6" in out[1]
assert '"answer": "fresh"' in out[1]
def test_yielded_live_event_advances_dedup_floor(self):
async def test_yielded_live_event_advances_dedup_floor(self):
"""A republish of an already-yielded seq must be dropped."""
live_first = encode_pubsub_message(
"msg-1", 0, "answer", {"type": "answer", "answer": "first"}
@@ -415,25 +383,18 @@ class TestDedupFloorSeededFromCursor:
live_third = encode_pubsub_message(
"msg-1", 1, "answer", {"type": "answer", "answer": "third"}
).encode("utf-8")
with patch(
"application.streaming.event_replay.db_readonly"
) as mock_readonly, patch(
"application.streaming.event_replay.MessageEventsRepository"
) as mock_repo_cls, patch(
"application.streaming.event_replay.Topic.subscribe",
_fake_topic_subscribe([live_first, live_dup, live_third]),
create=False,
with patch(_READONLY) as mock_readonly, patch(_REPO) as mock_repo_cls, patch(
_ASYNC_TOPIC, _fake_subscribe([live_first, live_dup, live_third])
):
mock_readonly.return_value.__enter__.return_value = MagicMock()
mock_repo_cls.return_value.read_after.return_value = []
gen = build_message_event_stream(
gen = build_message_event_stream_async(
"msg-1",
last_event_id=None,
keepalive_seconds=0.05,
poll_timeout_seconds=0.01,
)
out = _drain(gen, max_items=4)
out = await _drain(gen, max_items=4)
assert out[0] == ": connected\n\n"
assert "id: 0" in out[1]
@@ -444,13 +405,13 @@ class TestDedupFloorSeededFromCursor:
@pytest.mark.unit
@pytest.mark.asyncio
class TestSnapshotWhenSubscribeUnavailable:
"""Regression for C1: when Redis is down, ``Topic.subscribe`` exits
immediately without firing ``on_subscribe``. The snapshot is in
Postgres and must still be served.
"""When Redis is down, ``AsyncTopic.subscribe`` exits immediately. The
snapshot is in Postgres and must still be served.
"""
def test_snapshot_served_when_subscribe_returns_immediately(self):
async def test_snapshot_served_when_subscribe_returns_immediately(self):
rows = [
{
"sequence_no": 0,
@@ -458,44 +419,30 @@ class TestSnapshotWhenSubscribeUnavailable:
"payload": {"type": "answer", "answer": "from snapshot"},
},
]
# Topic.subscribe yields nothing (Redis-down behaviour).
def _empty_subscribe(self, on_subscribe=None, poll_timeout=1.0):
return
yield # pragma: no cover (make the function a generator)
with patch(
"application.streaming.event_replay.db_readonly"
) as mock_readonly, patch(
"application.streaming.event_replay.MessageEventsRepository"
) as mock_repo_cls, patch(
"application.streaming.event_replay.Topic.subscribe",
_empty_subscribe,
create=False,
with patch(_READONLY) as mock_readonly, patch(_REPO) as mock_repo_cls, patch(
_ASYNC_TOPIC, _subscribe_returns_immediately(fire_callback=False)
):
mock_readonly.return_value.__enter__.return_value = MagicMock()
mock_repo_cls.return_value.read_after.return_value = rows
gen = build_message_event_stream(
gen = build_message_event_stream_async(
"msg-1",
last_event_id=None,
keepalive_seconds=0.05,
poll_timeout_seconds=0.01,
)
out = list(gen)
out = await _drain(gen, max_items=50)
assert out[0] == ": connected\n\n"
# Snapshot row served via the post-subscribe fallback path.
assert "id: 0" in out[1]
assert '"answer": "from snapshot"' in out[1]
def test_callback_fired_then_subscribe_dies_does_not_duplicate(self):
"""Regression: if ``on_subscribe`` ran and populated the buffer,
a subsequent inner-generator failure (e.g. transient Redis
``get_message`` exception between SUBSCRIBE-ack and the first
poll) must not trigger a second snapshot read. Re-reading would
append the same rows twice and double the answer chunks on the
client (the per-message reconnect dispatcher does not dedup
by ``id``).
async def test_callback_fired_then_subscribe_dies_does_not_duplicate(self):
"""If ``on_subscribe`` ran and populated the buffer, a subsequent
inner-generator failure must not trigger a second snapshot read.
Re-reading would append the same rows twice and double the answer
chunks on the client (the reconnect dispatcher does not dedup by
``id``).
"""
rows = [
{
@@ -509,46 +456,26 @@ class TestSnapshotWhenSubscribeUnavailable:
"payload": {"type": "answer", "answer": "second"},
},
]
# Mimic the broadcast_channel race: SUBSCRIBE acks, on_subscribe
# runs, then the inner ``get_message`` raises and the generator
# returns without ever yielding.
def _subscribe_dies_after_callback(
self, on_subscribe=None, poll_timeout=1.0
):
if on_subscribe is not None:
on_subscribe()
return
yield # pragma: no cover (make the function a generator)
repo_mock = MagicMock()
repo_mock.read_after.return_value = rows
with patch(
"application.streaming.event_replay.db_readonly"
) as mock_readonly, patch(
"application.streaming.event_replay.MessageEventsRepository",
return_value=repo_mock,
with patch(_READONLY) as mock_readonly, patch(
_REPO, return_value=repo_mock
), patch(
"application.streaming.event_replay.Topic.subscribe",
_subscribe_dies_after_callback,
create=False,
_ASYNC_TOPIC, _subscribe_returns_immediately(fire_callback=True)
):
mock_readonly.return_value.__enter__.return_value = MagicMock()
gen = build_message_event_stream(
gen = build_message_event_stream_async(
"msg-1",
last_event_id=None,
keepalive_seconds=0.05,
poll_timeout_seconds=0.01,
)
out = list(gen)
out = await _drain(gen, max_items=50)
# The snapshot must have been read exactly once — re-reading
# would have re-appended the rows behind the originals.
# The snapshot must have been read exactly once.
assert repo_mock.read_after.call_count == 1
assert out[0] == ": connected\n\n"
# Each row appears exactly once, in order.
assert "id: 0" in out[1]
assert '"answer": "first"' in out[1]
assert "id: 1" in out[2]
@@ -558,15 +485,14 @@ class TestSnapshotWhenSubscribeUnavailable:
@pytest.mark.unit
@pytest.mark.asyncio
class TestTerminalEventClosesStream:
"""Regression: ``/api/messages/<id>/events`` is a live tail that
keeps emitting keepalives. Without explicit close-on-terminal the
client's drain promise never resolves and the WSGI thread is
pinned waiting for events that won't come for an already-finished
stream.
"""Without explicit close-on-terminal the client's drain promise never
resolves and the connection is pinned waiting for events that won't
come for an already-finished stream.
"""
def test_terminal_in_snapshot_closes_after_flush(self):
async def test_terminal_in_snapshot_closes_after_flush(self):
rows = [
{
"sequence_no": 0,
@@ -579,28 +505,18 @@ class TestTerminalEventClosesStream:
"payload": {"type": "end"},
},
]
with patch(
"application.streaming.event_replay.db_readonly"
) as mock_readonly, patch(
"application.streaming.event_replay.MessageEventsRepository"
) as mock_repo_cls, patch(
"application.streaming.event_replay.Topic.subscribe",
_fake_topic_subscribe([]),
create=False,
with patch(_READONLY) as mock_readonly, patch(_REPO) as mock_repo_cls, patch(
_ASYNC_TOPIC, _fake_subscribe([])
):
mock_readonly.return_value.__enter__.return_value = MagicMock()
mock_repo_cls.return_value.read_after.return_value = rows
gen = build_message_event_stream(
gen = build_message_event_stream_async(
"msg-1",
last_event_id=None,
keepalive_seconds=0.05,
poll_timeout_seconds=0.01,
)
# Drain to exhaustion. Without the close-on-terminal fix
# this would hang; with it, the generator returns after
# flushing the snapshot.
out = list(gen)
out = await _drain(gen, max_items=50)
assert out[0] == ": connected\n\n"
assert "id: 0" in out[1]
@@ -608,10 +524,10 @@ class TestTerminalEventClosesStream:
# No keepalives or further frames after the terminal.
assert all("keepalive" not in line for line in out[3:])
def test_terminal_falls_back_to_event_type_when_payload_lacks_type(self):
# Belt-and-suspenders: a journal write that records ``end`` only
# in the column (e.g. an abort handler that didn't seed
# ``payload.type``) must still terminate the replay.
async def test_terminal_falls_back_to_event_type_when_payload_lacks_type(self):
# A journal write that records ``end`` only in the column (e.g. an
# abort handler that didn't seed ``payload.type``) must still
# terminate the replay.
rows = [
{
"sequence_no": 0,
@@ -624,89 +540,68 @@ class TestTerminalEventClosesStream:
"payload": {},
},
]
with patch(
"application.streaming.event_replay.db_readonly"
) as mock_readonly, patch(
"application.streaming.event_replay.MessageEventsRepository"
) as mock_repo_cls, patch(
"application.streaming.event_replay.Topic.subscribe",
_fake_topic_subscribe([]),
create=False,
with patch(_READONLY) as mock_readonly, patch(_REPO) as mock_repo_cls, patch(
_ASYNC_TOPIC, _fake_subscribe([])
):
mock_readonly.return_value.__enter__.return_value = MagicMock()
mock_repo_cls.return_value.read_after.return_value = rows
gen = build_message_event_stream(
gen = build_message_event_stream_async(
"msg-1",
last_event_id=None,
keepalive_seconds=0.05,
poll_timeout_seconds=0.01,
)
out = list(gen)
out = await _drain(gen, max_items=50)
assert out[0] == ": connected\n\n"
assert "id: 0" in out[1]
assert "id: 1" in out[2]
assert all("keepalive" not in line for line in out[3:])
def test_terminal_in_live_tail_closes(self):
async def test_terminal_in_live_tail_closes(self):
live_answer = encode_pubsub_message(
"msg-1", 0, "answer", {"type": "answer", "answer": "x"}
).encode("utf-8")
live_end = encode_pubsub_message(
"msg-1", 1, "end", {"type": "end"}
).encode("utf-8")
with patch(
"application.streaming.event_replay.db_readonly"
) as mock_readonly, patch(
"application.streaming.event_replay.MessageEventsRepository"
) as mock_repo_cls, patch(
"application.streaming.event_replay.Topic.subscribe",
_fake_topic_subscribe([live_answer, live_end]),
create=False,
with patch(_READONLY) as mock_readonly, patch(_REPO) as mock_repo_cls, patch(
_ASYNC_TOPIC, _fake_subscribe([live_answer, live_end])
):
mock_readonly.return_value.__enter__.return_value = MagicMock()
mock_repo_cls.return_value.read_after.return_value = []
gen = build_message_event_stream(
gen = build_message_event_stream_async(
"msg-1",
last_event_id=None,
keepalive_seconds=0.05,
poll_timeout_seconds=0.01,
)
out = list(gen)
out = await _drain(gen, max_items=50)
assert out[0] == ": connected\n\n"
assert "id: 0" in out[1] and '"answer": "x"' in out[1]
assert "id: 1" in out[2] and '"type": "end"' in out[2]
assert all("keepalive" not in line for line in out[3:])
def test_error_event_also_closes(self):
async def test_error_event_also_closes(self):
"""The agent's catch-all path emits ``error`` with no trailing
``end`` — treating ``error`` as terminal closes that path too.
"""
live_err = encode_pubsub_message(
"msg-1", 0, "error", {"type": "error", "error": "boom"}
).encode("utf-8")
with patch(
"application.streaming.event_replay.db_readonly"
) as mock_readonly, patch(
"application.streaming.event_replay.MessageEventsRepository"
) as mock_repo_cls, patch(
"application.streaming.event_replay.Topic.subscribe",
_fake_topic_subscribe([live_err]),
create=False,
with patch(_READONLY) as mock_readonly, patch(_REPO) as mock_repo_cls, patch(
_ASYNC_TOPIC, _fake_subscribe([live_err])
):
mock_readonly.return_value.__enter__.return_value = MagicMock()
mock_repo_cls.return_value.read_after.return_value = []
gen = build_message_event_stream(
gen = build_message_event_stream_async(
"msg-1",
last_event_id=None,
keepalive_seconds=0.05,
poll_timeout_seconds=0.01,
)
out = list(gen)
out = await _drain(gen, max_items=50)
assert out[0] == ": connected\n\n"
assert "id: 0" in out[1] and '"type": "error"' in out[1]
@@ -719,36 +614,34 @@ def test_sse_line_split_pattern_handles_all_terminators():
@pytest.mark.unit
@pytest.mark.asyncio
class TestWatchdogClosesIdleReconnect:
"""Without the watchdog a reconnect stream with a non-terminal
snapshot and a dead producer (worker crash between chunks and
finalize) would emit keepalives forever — the live tail blocks
on a pub/sub that nobody publishes to and the frontend's drain
promise never resolves. The watchdog periodically inspects
``conversation_messages`` and closes the stream with a terminal
SSE event when the row has gone terminal in the DB or the
producer's heartbeat has gone stale.
"""Without the watchdog a reconnect stream with a non-terminal snapshot
and a dead producer would emit keepalives forever. The watchdog
periodically inspects ``conversation_messages`` and closes the stream
with a terminal SSE event when the row has gone terminal in the DB or
the producer's heartbeat has gone stale.
"""
@staticmethod
def _patch_subscribe_idle_forever():
"""Build a fake ``Topic.subscribe`` that fires ``on_subscribe``
and then yields ``None`` ticks indefinitely (i.e. the producer
is gone). Returns the generator function and a list collecting
the ``on_subscribe`` calls so tests can assert ordering.
def _subscribe_idle_forever():
"""``AsyncTopic.subscribe`` that fires ``on_subscribe`` then yields
``None`` ticks indefinitely (i.e. the producer is gone).
"""
def _impl(self, on_subscribe=None, poll_timeout=1.0):
async def _impl(self, on_subscribe=None, poll_timeout=1.0):
if on_subscribe is not None:
on_subscribe()
res = on_subscribe()
if inspect.isawaitable(res):
await res
while True:
yield None
return _impl
def _mock_liveness_row(self, status, err=None, is_stale=False):
"""Build the ``conn.execute(...).first()`` return value the
watchdog SQL expects — ``(status, err, is_stale)``.
"""Build the ``conn.execute(...).first()`` return value the watchdog
SQL expects — ``(status, err, is_stale)``.
"""
return (status, err, is_stale)
@@ -762,25 +655,17 @@ class TestWatchdogClosesIdleReconnect:
):
"""Wire up the patches the watchdog tests share.
The snapshot read goes through the patched
``MessageEventsRepository`` (returns empty), and the watchdog
liveness check goes through ``conn.execute(...).first()`` on the
same ``db_readonly``-yielded ``MagicMock`` connection.
The snapshot read goes through the patched ``MessageEventsRepository``
(returns empty), and the watchdog liveness check goes through
``conn.execute(...).first()`` on the same ``db_readonly``-yielded
``MagicMock`` connection.
"""
mock_conn = MagicMock()
mock_conn.execute.return_value.first.return_value = liveness_row
readonly_patch = patch(
"application.streaming.event_replay.db_readonly"
)
repo_patch = patch(
"application.streaming.event_replay.MessageEventsRepository"
)
subscribe_patch = patch(
"application.streaming.event_replay.Topic.subscribe",
self._patch_subscribe_idle_forever(),
create=False,
)
readonly_patch = patch(_READONLY)
repo_patch = patch(_REPO)
subscribe_patch = patch(_ASYNC_TOPIC, self._subscribe_idle_forever())
mock_readonly = readonly_patch.start()
mock_repo_cls = repo_patch.start()
@@ -789,7 +674,7 @@ class TestWatchdogClosesIdleReconnect:
mock_readonly.return_value.__enter__.return_value = mock_conn
mock_repo_cls.return_value.read_after.return_value = []
gen = build_message_event_stream(
gen = build_message_event_stream_async(
"msg-1",
last_event_id=None,
keepalive_seconds=keepalive_seconds,
@@ -799,34 +684,32 @@ class TestWatchdogClosesIdleReconnect:
)
return gen, [readonly_patch, repo_patch, subscribe_patch]
def test_watchdog_emits_synthetic_end_when_status_complete(self):
"""A row that flipped to ``complete`` after the snapshot read
(e.g. finalize ran on another worker, journal write lost) must
async def test_watchdog_emits_synthetic_end_when_status_complete(self):
"""A row that flipped to ``complete`` after the snapshot read must
be surfaced as ``end`` so the client closes cleanly.
"""
gen, patches = self._build_gen_with_liveness(
self._mock_liveness_row("complete")
)
try:
out = _drain(gen, max_items=5)
out = await _drain(gen, max_items=5)
finally:
for p in patches:
p.stop()
assert out[0] == ": connected\n\n"
# Watchdog synthetic: id:-1, ``{"type": "end"}``
terminal = [s for s in out if '"type": "end"' in s]
assert len(terminal) == 1
assert "id: -1" in terminal[0]
def test_watchdog_emits_synthetic_error_when_status_failed(self):
async def test_watchdog_emits_synthetic_error_when_status_failed(self):
gen, patches = self._build_gen_with_liveness(
self._mock_liveness_row(
"failed", err="RuntimeError: upstream blew up"
)
)
try:
out = _drain(gen, max_items=5)
out = await _drain(gen, max_items=5)
finally:
for p in patches:
p.stop()
@@ -835,21 +718,15 @@ class TestWatchdogClosesIdleReconnect:
terminal = [s for s in out if '"type": "error"' in s]
assert len(terminal) == 1
assert '"code": "producer_failed"' in terminal[0]
# The stashed error from metadata is surfaced verbatim so the
# UI can show the real reason instead of a generic message.
assert "RuntimeError: upstream blew up" in terminal[0]
def test_watchdog_emits_synthetic_error_when_producer_stale(self):
"""Non-terminal status + heartbeat older than the threshold ⇒
producing worker is presumed dead. Without this, the live tail
would hang on keepalives until the proxy idle-timeout fires.
"""
async def test_watchdog_emits_synthetic_error_when_producer_stale(self):
gen, patches = self._build_gen_with_liveness(
self._mock_liveness_row("streaming", is_stale=True),
producer_idle_seconds=1.0,
)
try:
out = _drain(gen, max_items=5)
out = await _drain(gen, max_items=5)
finally:
for p in patches:
p.stop()
@@ -859,37 +736,34 @@ class TestWatchdogClosesIdleReconnect:
assert len(terminal) == 1
assert '"code": "producer_stale"' in terminal[0]
def test_watchdog_does_not_fire_while_producer_alive(self):
async def test_watchdog_does_not_fire_while_producer_alive(self):
"""A non-terminal row with a fresh heartbeat is healthy; the
watchdog must keep silent (yield keepalives instead) so a
slow-but-alive stream isn't prematurely terminated.
watchdog must keep silent (yield keepalives instead).
"""
gen, patches = self._build_gen_with_liveness(
self._mock_liveness_row("streaming", is_stale=False),
keepalive_seconds=0.01,
)
try:
out = _drain(gen, max_items=5)
out = await _drain(gen, max_items=5)
finally:
for p in patches:
p.stop()
assert out[0] == ": connected\n\n"
# No synthetic terminal — the rest of the output is keepalives.
assert all(
'"type": "end"' not in s and '"type": "error"' not in s
for s in out
)
assert any("keepalive" in s for s in out)
def test_watchdog_handles_missing_row_as_terminal(self):
"""If the message row got deleted out from under us mid-tail,
the watchdog must close the stream rather than tail forever
on a row that no longer exists.
async def test_watchdog_handles_missing_row_as_terminal(self):
"""If the message row got deleted out from under us mid-tail, the
watchdog must close the stream rather than tail forever.
"""
gen, patches = self._build_gen_with_liveness(None)
try:
out = _drain(gen, max_items=5)
out = await _drain(gen, max_items=5)
finally:
for p in patches:
p.stop()