diff --git a/src/lembas/api/chats.py b/src/lembas/api/chats.py index 59c127b..db22ee1 100644 --- a/src/lembas/api/chats.py +++ b/src/lembas/api/chats.py @@ -5,6 +5,7 @@ from __future__ import annotations import asyncio import json import logging +import time from collections.abc import AsyncIterator from fastapi import APIRouter, Depends, Form, HTTPException, Request, status @@ -28,6 +29,10 @@ log = logging.getLogger(__name__) router = APIRouter(prefix="/api/chats", tags=["chats"]) +# Seconds of silence before a comment frame is sent to hold the connection open. +# Well under nginx's 60s default; see services/sse.py:KEEPALIVE. +KEEPALIVE_AFTER = 15.0 + def _owned_chat(db: DBSession, chat_id: str, user_id: str) -> Chat: chat = db.get(Chat, chat_id) # 404 rather than 403 for someone else's chat: whether a given id exists is @@ -240,6 +245,7 @@ async def _follow(chat_id: str, message_id: str) -> AsyncIterator[str]: generation = generation_service.ensure(chat_id, message_id) generation.followers += 1 seen = -1 + last_frame = time.monotonic() try: while True: @@ -251,9 +257,18 @@ async def _follow(chat_id: str, message_id: str) -> AsyncIterator[str]: yield sse.event("tools", _tool_activity(generation.tool_events)) if generation.content: yield sse.event("render", render_markdown(generation.text)) + last_frame = time.monotonic() if generation.done: break + + # A reasoning model can think for a minute or more without emitting + # anything, and an idle connection is what a proxy closes. The + # comment frame keeps it open and is ignored by the browser. + if time.monotonic() - last_frame > KEEPALIVE_AFTER: + yield sse.KEEPALIVE + last_frame = time.monotonic() + # Polling rather than per-follower wakeups: the producer already # works in RENDER_INTERVAL steps, so a short sleep is simpler and # cannot drop a notification. @@ -261,7 +276,7 @@ async def _follow(chat_id: str, message_id: str) -> AsyncIterator[str]: finally: generation.followers = max(0, generation.followers - 1) - # The producer writes the message before marking itself done, so by here + # The producer commits the message before marking itself done, so by here # the row is authoritative and the final bubble can be rendered from it. with session_scope() as db: message = db.get(Message, message_id) @@ -592,7 +607,9 @@ async def regenerate( message.complete = False message.model_id = chat.model_id db.commit() - generation_service.ensure(chat.id, message.id) + # restart, not ensure: this is the one caller that reuses a Message row, and + # the finished generation for it is still registered. + generation_service.restart(chat.id, message.id) return templates.TemplateResponse( request, diff --git a/src/lembas/services/generation.py b/src/lembas/services/generation.py index a356bc4..aaf751f 100644 --- a/src/lembas/services/generation.py +++ b/src/lembas/services/generation.py @@ -120,18 +120,48 @@ def ensure(chat_id: str, message_id: str) -> Generation: Idempotent, because more than one thing can ask for it: the route that created the message, and any page load that finds the message unfinished. + + `_prune` runs first, not after the lookup. Below it, a stale entry could + never expire: the early return is the only path a repeated id takes, so the + sweep was unreachable for exactly the message that needed it. """ + _prune() existing = _RUNNING.get(message_id) if existing is not None: return existing - _prune() generation = Generation(chat_id=chat_id, message_id=message_id) _RUNNING[message_id] = generation _TASKS[message_id] = asyncio.create_task(_run(generation)) return generation +def restart(chat_id: str, message_id: str) -> Generation: + """Produce this reply again, discarding any finished attempt at it. + + `ensure` is idempotent on purpose, and that is load-bearing: a page load + finding an unfinished reply must attach to it rather than start a second + one, and `_follow` calls it too. Regeneration is the one caller that means + the opposite. + + It is also the one caller that reuses an existing Message row -- blanked and + marked incomplete -- rather than creating a new one. The finished Generation + for that id is still in the registry, because finished ones linger + KEEP_FINISHED so a follower arriving at the last moment still gets the final + frames. `ensure` handed that one straight back: no request was made, + `_follow` replayed the previous answer, and the `done` frame re-rendered a + streaming shell because the row said incomplete. That was the reconnect loop, + and the Send button stuck on Stop. + """ + previous = _RUNNING.pop(message_id, None) + task = _TASKS.pop(message_id, None) + if previous is not None and not previous.done: + previous.cancel = True + if task is not None: + task.cancel() + return ensure(chat_id, message_id) + + async def shutdown() -> None: """Stop every running generation, keeping what each has produced.""" for task in list(_TASKS.values()): @@ -295,10 +325,14 @@ async def _run(generation: Generation) -> None: ) title = title or chat_service.fallback_title(question) + # Written *before* `done`, because `_follow` breaks out of its loop the + # moment it sees that flag and immediately re-renders the bubble from + # the row. The other order left a window in which the finished frame + # showed the previous turn's stored values. + _persist(generation, title, time.monotonic() - started) generation.done = True generation.finished_at = datetime.now(UTC) generation.touch() - _persist(generation, title, time.monotonic() - started) def _question_from(payload: dict) -> str: @@ -319,7 +353,21 @@ def _question_from(payload: dict) -> str: def _persist(generation: Generation, title: str, elapsed: float) -> None: - """Write the finished reply, name the chat, and set the unread flag.""" + """Write the finished reply, name the chat, and set the unread flag. + + A generation another one has replaced may not write. A regeneration cancels + its predecessor, whose `finally:` then runs this on the same row -- and it + would overwrite the fresh reply with the abandoned one. + + The test is "someone else owns this row now", not "this one is registered": + an unregistered generation still writes, because that is a direct call + rather than a superseded one. + """ + owner = _RUNNING.get(generation.message_id) + if owner is not None and owner is not generation: + log.debug("skipping persist for superseded generation %s", generation.message_id) + return + try: with session_scope() as db: message = db.get(Message, generation.message_id) @@ -364,5 +412,6 @@ __all__ = [ "ensure", "get", "request_stop", + "restart", "shutdown", ] diff --git a/src/lembas/web/templates/chat/_message.html b/src/lembas/web/templates/chat/_message.html index 0b7b5a0..cf6e51c 100644 --- a/src/lembas/web/templates/chat/_message.html +++ b/src/lembas/web/templates/chat/_message.html @@ -98,7 +98,10 @@ Thinking… {{ icon("chevron-down", "icon--sm reasoning__chevron") }} -
+ {# innerHTML, not beforeend: the frame carries the whole block of + thinking each time, exactly as `render` and `tools` do. Appending it + repeated everything already shown, so the panel grew quadratically. #} + {# Tool activity as it happens. Empty until the model asks for something, diff --git a/tests/test_regenerate.py b/tests/test_regenerate.py new file mode 100644 index 0000000..313ada1 --- /dev/null +++ b/tests/test_regenerate.py @@ -0,0 +1,262 @@ +"""Regenerating a reply, and the registry rules that make it work. + +Regeneration is the one caller that reuses a Message row rather than creating a +new one, which is why it is the one caller `ensure` was wrong for. +""" + +from __future__ import annotations + +import asyncio +from datetime import UTC, datetime, timedelta + +import pytest +from fastapi.testclient import TestClient +from sqlalchemy import select + +from lembas.db.models import Chat, Connection, Message, Model +from lembas.services import generation as generation_service +from lembas.services.crypto import encrypt + + +@pytest.fixture(autouse=True) +def empty_registry(): + """The registry is module state. A test that leaves an entry behind changes + what the next one sees.""" + yield + generation_service._RUNNING.clear() + generation_service._TASKS.clear() + + +@pytest.fixture +def no_upstream(monkeypatch): + """Replace the producer, so a test can watch the registry without a server. + + Records the generations it was asked to run. + """ + started: list = [] + + async def _fake_run(generation): + started.append(generation) + + monkeypatch.setattr(generation_service, "_run", _fake_run) + return started + + +def _connection(db) -> Connection: + connection = Connection( + name="Test", base_url="http://127.0.0.1:1", api_key_encrypted=encrypt("") + ) + db.add(connection) + db.commit() + db.add(Model(connection_id=connection.id, model_id="test-model")) + db.commit() + return connection + + +def _reply(db, chat_id: str, *, complete: bool = True) -> Message: + db.add(Message(chat_id=chat_id, role="user", content="what is lembas?")) + reply = Message(chat_id=chat_id, role="assistant", content="Waybread.", complete=complete) + db.add(reply) + db.commit() + return reply + + +def _finished(chat_id: str, message_id: str, *, text: str = "old") -> generation_service.Generation: + """A generation in the state a just-completed reply leaves behind.""" + generation = generation_service.Generation(chat_id=chat_id, message_id=message_id) + generation.content.append(text) + generation.done = True + generation.finished_at = datetime.now(UTC) + generation_service._RUNNING[message_id] = generation + return generation + + +# --- The bug ----------------------------------------------------------------- +def test_regenerate_starts_a_fresh_generation( + client: TestClient, db, registered, make_chat, no_upstream +): + """The bug: `ensure` handed back the finished generation still sitting in + the registry, so no request was ever made and the browser reconnected to a + stream that had nothing left to say.""" + _connection(db) + chat_id = make_chat() + reply = _reply(db, chat_id) + stale = _finished(chat_id, reply.id) + + response = client.post(f"/api/chats/{chat_id}/messages/{reply.id}/regenerate") + assert response.status_code == 200 + + fresh = generation_service.get(reply.id) + assert fresh is not stale + assert not fresh.done + assert fresh.text == "" + assert no_upstream == [fresh] + + +def test_regenerate_blanks_the_row_and_returns_a_streaming_shell( + client: TestClient, db, registered, make_chat, no_upstream +): + _connection(db) + chat_id = make_chat() + reply = _reply(db, chat_id) + + body = client.post(f"/api/chats/{chat_id}/messages/{reply.id}/regenerate").text + assert "sse-connect=" in body + assert 'sse-swap="render"' in body + + db.expire_all() + row = db.get(Message, reply.id) + assert row.complete is False + assert row.content == "" + + +def test_regenerating_twice_in_a_row_works( + client: TestClient, db, registered, make_chat, no_upstream +): + """The reported symptom was that it worked once, sometimes.""" + _connection(db) + chat_id = make_chat() + reply = _reply(db, chat_id) + + client.post(f"/api/chats/{chat_id}/messages/{reply.id}/regenerate") + first = generation_service.get(reply.id) + first.done = True + first.finished_at = datetime.now(UTC) + + client.post(f"/api/chats/{chat_id}/messages/{reply.id}/regenerate") + assert generation_service.get(reply.id) is not first + assert len(no_upstream) == 2 + + +def test_regenerating_someone_elses_reply_is_not_found( + client: TestClient, db, registered, make_chat, no_upstream +): + from lembas.db.models import User + from lembas.security.passwords import hash_password + + _connection(db) + chat_id = make_chat() + reply = _reply(db, chat_id) + + someone_else = User(name="Sam", email="s@shire.test", password_hash=hash_password("x")) + db.add(someone_else) + db.commit() + db.scalar(select(Chat).where(Chat.id == chat_id)).user_id = someone_else.id + db.commit() + + response = client.post(f"/api/chats/{chat_id}/messages/{reply.id}/regenerate") + assert response.status_code == 404 + + +# --- Registry rules ---------------------------------------------------------- +async def test_ensure_still_attaches_to_an_unfinished_reply(db, no_upstream): + """Idempotence is load-bearing: a page load that finds an unfinished reply + must attach to it, not start a second one.""" + first = generation_service.ensure("chat", "message") + assert generation_service.ensure("chat", "message") is first + await asyncio.sleep(0) # let the scheduled task actually start + assert len(no_upstream) == 1 + + +async def test_prune_expires_a_finished_generation_before_the_lookup(db, no_upstream): + """`_prune` used to sit below the early return, where it could never reach + the one entry that needed it.""" + stale = _finished("chat", "message") + stale.finished_at = datetime.now(UTC) - generation_service.KEEP_FINISHED - timedelta(minutes=1) + + assert generation_service.ensure("chat", "message") is not stale + + +def test_a_superseded_generation_does_not_write_the_row(db, registered, make_chat): + """A cancelled predecessor's `finally:` runs _persist on the same message, + and it must not overwrite the reply that replaced it.""" + chat_id = make_chat() + reply = _reply(db, chat_id) + + abandoned = generation_service.Generation(chat_id=chat_id, message_id=reply.id) + abandoned.content.append("the abandoned attempt") + current = generation_service.Generation(chat_id=chat_id, message_id=reply.id) + current.content.append("the reply that replaced it") + generation_service._RUNNING[reply.id] = current + + generation_service._persist(abandoned, "", 0.0) + + db.expire_all() + assert db.get(Message, reply.id).content == "Waybread." + + +async def test_restart_cancels_a_generation_still_running(db, no_upstream): + live = generation_service.ensure("chat", "message") + generation_service.restart("chat", "message") + assert live.cancel is True + assert generation_service.get("message") is not live + + +# --- Ordering ---------------------------------------------------------------- +async def test_the_reply_is_persisted_before_it_is_marked_done( + db, registered, make_chat, monkeypatch +): + """`_follow` breaks out the instant it sees `done` and re-renders the bubble + from the row, so the row has to be right first.""" + chat_id = make_chat() + reply = _reply(db, chat_id, complete=False) + seen: list[bool] = [] + + def _record(generation, title, elapsed): + seen.append(generation.done) + + monkeypatch.setattr(generation_service, "_persist", _record) + + generation = generation_service.Generation(chat_id=chat_id, message_id=reply.id) + generation_service._RUNNING[reply.id] = generation + # No connection row, so resolve_endpoint fails and _run goes straight to + # its finally: which is the part under test. + await generation_service._run(generation) + + assert seen == [False] + assert generation.done is True + + +# --- The streaming shell ------------------------------------------------------ +def test_live_reasoning_is_replaced_not_appended( + client: TestClient, db, registered, make_chat, no_upstream +): + """The frame carries the whole block each time. Appending it repeated + everything already shown, so the panel grew quadratically.""" + _connection(db) + chat_id = make_chat() + reply = _reply(db, chat_id) + + body = client.post(f"/api/chats/{chat_id}/messages/{reply.id}/regenerate").text + line = next(line for line in body.splitlines() if 'sse-swap="reasoning"' in line) + assert 'hx-swap="innerHTML"' in line + assert "beforeend" not in line + + +async def test_a_silent_generation_gets_a_keepalive(db, registered, make_chat, monkeypatch): + """A model thinking for a minute emits nothing, and an idle connection is + what a proxy closes.""" + from lembas.api import chats as chats_api + + chat_id = make_chat() + reply = _reply(db, chat_id, complete=False) + + generation = generation_service.Generation(chat_id=chat_id, message_id=reply.id) + generation_service._RUNNING[reply.id] = generation + monkeypatch.setattr(chats_api, "KEEPALIVE_AFTER", 0.0) + + frames: list[str] = [] + stream = chats_api._follow(chat_id, reply.id).__aiter__() + + async def _finish(): + await asyncio.sleep(0.05) + generation.done = True + + task = asyncio.create_task(_finish()) + async for frame in stream: + frames.append(frame) + if frame.startswith("event: close"): + break + await task + + assert any(frame == ": keepalive\n\n" for frame in frames)