mirror of
https://github.com/tiennm99/DocsGPT.git
synced 2026-10-11 12:11:45 +00:00
feat: async events on asgi
This commit is contained in:
1 parent
1f06f31aa3
commit
73ed1cf607
13 files changed
+1270
-1054
No files matched your search
@@ -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
|
||||
@@ -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"],
|
||||
),
|
||||
]
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
)
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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"
|
||||
@@ -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
|
||||
@@ -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
@@ -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()
|
||||
|
||||
Reference in new issue
Block a user