diff --git a/bin/chat-chainlit.py b/bin/chat-chainlit.py index 059a75c..e70740f 100644 --- a/bin/chat-chainlit.py +++ b/bin/chat-chainlit.py @@ -8,13 +8,16 @@ import time from collections.abc import Awaitable, Callable from pathlib import Path +from typing import Any import chainlit as cl from chainlit.chat_context import chat_contexts +from chainlit.config import config as chainlit_config from chainlit.context import context from chainlit.data.base import BaseDataLayer from chainlit.data.sql_alchemy import SQLAlchemyDataLayer from chainlit.oauth_providers import providers +from chainlit.session import WebsocketSession from chainlit.types import ThreadDict from dotenv import load_dotenv from langchain_community.callbacks import OpenAICallbackHandler @@ -32,13 +35,17 @@ submit_identifiers, ) from analysis.proposals import Proposal, proposals +from gsa import chat as gsa_chat from gsa.chainlit_flow import ( Attachment, matrix_attachment, + receive_matrix, result_file_kwargs, - run_analysis, + run_with_labels, ) from gsa.chat import HOW_TO_RUN_GSA, HOW_TO_RUN_GSA_WITH_A_MATRIX, asks_to_run_gsa +from gsa.pending import pending_matrices +from gsa.upload import Matrix from handoff import seed from handoff.store import AnalysisHandoff, handoffs from handoff.window import acknowledgement, claimed_id @@ -169,40 +176,49 @@ async def resume(thread: ThreadDict) -> None: @cl.on_chat_end async def end() -> None: await static_messages(config, TriggerEvent.on_chat_end) - # Nothing of a finished session should outlive it in memory (review, - # area 2): Chainlit's own per-session message list, our offers, and -- for - # a guest, who can never resume -- the conversation's checkpoints. + # Chainlit calls this on EVERY disconnect, a network blip included, and + # keeps the session for `session_timeout` so a reconnect can resume it. + # Cleaning up here deleted a guest's conversation, offers and waiting + # matrix on every blip (found 2026-10-04 driving a network drop). So the + # cleanup waits out the timeout, and runs only if the session is gone. sid = session_id() + thread_id = current_thread_id() + guest = cl.user_session.get("user") is None + task = asyncio.get_running_loop().create_task( + _clean_up_when_gone(sid, thread_id, guest) + ) + _cleanups.add(task) + task.add_done_callback(_cleanups.discard) + + +#: Cleanups waiting out the session timeout, so they are not collected first. +_cleanups: set[asyncio.Task[None]] = set() +#: Past Chainlit's own clear, so the session is gone by the time we look. +CLEANUP_GRACE_SECONDS = 60 + + +async def _clean_up_when_gone(sid: str, thread_id: str, guest: bool) -> None: + """Nothing of a finished session should outlive it in memory (review, + area 2): Chainlit's per-session message list, our offers and waiting + matrix, and -- for a guest, who can never resume -- the checkpoints.""" + await asyncio.sleep(chainlit_config.project.session_timeout + CLEANUP_GRACE_SECONDS) + if WebsocketSession.get_by_id(sid) is not None: + return # Reconnected: still in use. Its next disconnect tries again. + pending_matrices.drop(sid) chat_contexts.pop(sid, None) proposals.drop_session(sid) _seeding.pop(sid, None) - if cl.user_session.get("user") is None: + if guest: with contextlib.suppress(Exception): - await get_graph().forget_thread(current_thread_id()) - + await get_graph().forget_thread(thread_id) -async def run_gsa_analysis(attachment: Attachment) -> None: - """Drive `gsa.chainlit_flow` with this session's chat operations. - The flow takes these four as arguments so it can be tested without a - browser; this is the only place that knows they are Chainlit. - """ - # Created on the first update, not up front. - # - # It used to be sent before anything else, so it was the first message - # in the thread -- and Chainlit scrolls the latest user message to the - # top, which put the progress line above the fold. It updated faithfully - # for minutes where nobody could see it; the user saw "Started." and then - # nothing. Created lazily, it lands below "Started.", where they are - # looking. +def gsa_chat_ops() -> tuple[Any, Any, Any, Any]: + """send, update_progress, send_file, and finish (which clears progress).""" + # Progress is created on the first update, not up front: sent first, it + # sat above the fold, where Chainlit scrolls the latest user message. progress: cl.Message | None = None - async def ask_for_grouping() -> str | None: - answer = await cl.AskUserMessage( - content="Which group is each sample in?", timeout=600 - ).send() - return (answer or {}).get("output") if answer else None - async def send(text: str) -> None: await cl.Message(content=text).send() @@ -220,23 +236,37 @@ async def send_file(path: Path) -> None: content="", elements=[cl.File(**result_file_kwargs(path))] ).send() - # The progress line is removed rather than marked "Done". - # - # A `finally` that sets "Done." runs on the failure paths too, so a user - # whose analysis died would have been told it finished, one line above - # the message explaining that it had not. `run_analysis` says what - # happened on every path; this only has to stop the spinner. + async def finish() -> None: + # Removed rather than marked "Done": on a failure path that would + # contradict the message explaining it. + if progress is not None: + await progress.remove() + + return send, update_progress, send_file, finish + + +async def run_gsa_analysis(attachment: Attachment) -> None: + """Check an uploaded matrix and keep it, waiting for its labels.""" + send, _, _, _ = gsa_chat_ops() + matrix = await receive_matrix(attachment, send=send) + if matrix is not None: + pending_matrices.put(session_id(), matrix) + + +async def run_with_pending_labels(matrix: Matrix, reply: str) -> None: + send, update_progress, send_file, finish = gsa_chat_ops() try: - await run_analysis( - attachment, - ask_for_grouping=ask_for_grouping, + used = await run_with_labels( + matrix, + reply, send=send, update_progress=update_progress, send_file=send_file, ) finally: - if progress is not None: - await progress.remove() + await finish() + if not used: + pending_matrices.put(session_id(), matrix) async def continue_from_handoff(handoff_id: str) -> None: @@ -671,6 +701,17 @@ async def handle_message(message: cl.Message) -> None: await run_gsa_analysis(attachment) return + # A matrix waiting for labels: a reply that reads as labels runs it. + # Anything else is answered as usual, and the matrix keeps waiting. + waiting = pending_matrices.peek(session_id()) + if waiting is not None and gsa_chat.looks_like_labels( + message.content or "", len(waiting.samples) + ): + matrix = pending_matrices.take(session_id()) + if matrix is not None: + await run_with_pending_labels(matrix, message.content or "") + return + # "yes" to the offer just made is the same as clicking Run. if latest and gene_list.confirms(message.content or ""): proposal = await take_proposal(latest) diff --git a/src/gsa/chainlit_flow.py b/src/gsa/chainlit_flow.py index fab7247..92cdca8 100644 --- a/src/gsa/chainlit_flow.py +++ b/src/gsa/chainlit_flow.py @@ -20,7 +20,7 @@ from gsa import chat from gsa.client import GsaClient from gsa.job import AnalysisFailedError, await_result, submit_uploaded_matrix -from gsa.upload import UploadRejectedError, discard, validate +from gsa.upload import Matrix, UploadRejectedError, discard, validate from util.logging import logging logger = logging.getLogger(__name__) @@ -141,45 +141,86 @@ async def run_analysis( discard(path) -async def _run( - attachment: Attachment, - path: Path, - ask_for_grouping: Any, - send: Any, - update_progress: Any, - send_file: Any, - client: GsaClient | None, -) -> None: +async def receive_matrix(attachment: Attachment, *, send: Any) -> Matrix | None: + """Step one: check the upload and describe it. The file is kept, waiting + for labels, unless it was refused -- then it is deleted here.""" + path = Path(attachment.path) try: # Off the event loop: reading and checking a 20 MB file stalled every # session for seconds (review, area 2). matrix = await asyncio.to_thread(validate, path) except UploadRejectedError as refusal: + discard(path) await send(str(refusal)) - return + return None except OSError: # Unreadable, vanished, or not a file. The user did nothing wrong # that they can act on, so do not describe it as their mistake. logger.exception("could not read an uploaded file") + discard(path) await send("I could not read that file — please try attaching it again.") - return - + return None await send(chat.describe_matrix(matrix)) + return matrix + + +async def run_with_labels( + matrix: Matrix, + reply: str, + *, + send: Any, + update_progress: Any, + send_file: Any, + client: GsaClient | None = None, +) -> bool: + """Step two. False, with the file kept, if the reply is not a grouping + the reader can fix by replying again; True once the matrix is used.""" + try: + grouping = chat.parse_grouping(reply, len(matrix.samples)) + except chat.ReplyUnusableError as unusable: + await send(f"{unusable} Reply with the labels again when you are ready.") + return False + try: + await _run_grouped(matrix, grouping, send, update_progress, send_file, client) + finally: + discard(matrix.path) + return True + +async def _run( + attachment: Attachment, + path: Path, + ask_for_grouping: Any, + send: Any, + update_progress: Any, + send_file: Any, + client: GsaClient | None, +) -> None: + matrix = await receive_matrix(attachment, send=send) + if matrix is None: + return reply = await ask_for_grouping() if not reply: - discard(path) await send( "No labels arrived, so I have not run anything. The file is deleted." ) return - try: grouping = chat.parse_grouping(reply, len(matrix.samples)) except chat.ReplyUnusableError as unusable: await send(f"{unusable} Send the file again when you are ready.") return + await _run_grouped(matrix, grouping, send, update_progress, send_file, client) + +async def _run_grouped( + matrix: Matrix, + grouping: chat.Grouping, + send: Any, + update_progress: Any, + send_file: Any, + client: GsaClient | None, +) -> None: gsa = client or GsaClient() try: # `submit_uploaded_matrix` deletes the file itself, on every path. diff --git a/src/gsa/chat.py b/src/gsa/chat.py index 893ae89..a2b46e5 100644 --- a/src/gsa/chat.py +++ b/src/gsa/chat.py @@ -122,6 +122,22 @@ def parse_grouping(reply: str, sample_count: int) -> Grouping: return Grouping(labels=canonical, group1=group1, group2=group2) +def looks_like_labels(reply: str, sample_count: int) -> bool: + """Whether a message, sent while a matrix waits, is an attempt at labels. + + A grouping that parses certainly is. One that does not but has the shape + of a list -- separators, no question mark, about the right length -- is + a mistake to explain, not a question for the model. + """ + try: + parse_grouping(reply, sample_count) + return True + except ReplyUnusableError: + pass + shaped = any(sep in reply for sep in (",", ";", "\t")) + return shaped and "?" not in reply and len(reply) < 40 * sample_count + 200 + + def describe_progress(status: AnalysisStatus) -> str: """One line, safe to send repeatedly as an edit. diff --git a/src/gsa/pending.py b/src/gsa/pending.py new file mode 100644 index 0000000..9d6ab35 --- /dev/null +++ b/src/gsa/pending.py @@ -0,0 +1,72 @@ +"""Uploaded matrices waiting for their sample labels. + +The labels used to be asked for with Chainlit's `AskUserMessage`, which is +tied to the socket open when it was asked. A reconnect -- a laptop sleeping, +a network change, a backgrounded phone tab -- left it waiting on a socket +that was gone: the reader's "control, control, treated, treated" was answered +by the model as an ordinary question, and ten minutes later the chat said no +labels had arrived (review, area 2). Now the matrix waits here, and the +reader's next message that reads as labels runs it. + +In process memory, like the gene-list offers: plain data, bounded, gone on +restart -- after which the reader attaches the file again. +""" + +import time +from collections import OrderedDict +from dataclasses import dataclass, field + +from gsa.upload import Matrix, discard + +#: How long an uploaded matrix waits for labels. +WAIT_SECONDS = 30 * 60 +MAX_SESSIONS = 500 + + +@dataclass +class PendingMatrices: + wait_seconds: float = WAIT_SECONDS + max_sessions: int = MAX_SESSIONS + _waiting: OrderedDict[str, tuple[Matrix, float]] = field( + default_factory=OrderedDict + ) + + def put(self, session_id: str, matrix: Matrix, now: float | None = None) -> None: + """Keep a matrix for this session; one per session, newest wins. + + Putting the same matrix back -- after a reply that was not usable + labels -- must not delete its file, which dropping the old entry would. + """ + previous = self._waiting.pop(session_id, None) + if previous is not None and previous[0].path != matrix.path: + discard(previous[0].path) + self._waiting[session_id] = (matrix, time.time() if now is None else now) + while len(self._waiting) > self.max_sessions: + _, (old, _) = self._waiting.popitem(last=False) + discard(old.path) + + def peek(self, session_id: str, now: float | None = None) -> Matrix | None: + """The matrix waiting for this session, if it has not expired.""" + found = self._waiting.get(session_id) + if found is None: + return None + matrix, since = found + if (time.time() if now is None else now) - since > self.wait_seconds: + self.drop(session_id) + return None + return matrix + + def take(self, session_id: str, now: float | None = None) -> Matrix | None: + matrix = self.peek(session_id, now) + if matrix is not None: + self._waiting.pop(session_id, None) + return matrix + + def drop(self, session_id: str) -> None: + """Forget it and delete its file.""" + found = self._waiting.pop(session_id, None) + if found is not None: + discard(found[0].path) + + +pending_matrices = PendingMatrices() diff --git a/tests/gsa/test_gsa_pending.py b/tests/gsa/test_gsa_pending.py new file mode 100644 index 0000000..933eaed --- /dev/null +++ b/tests/gsa/test_gsa_pending.py @@ -0,0 +1,73 @@ +"""A matrix waiting for labels, across messages rather than a blocking ask.""" + +from pathlib import Path + +import pytest + +from gsa import chat +from gsa.pending import PendingMatrices +from gsa.upload import Matrix + + +def matrix(tmp_path: Path, name: str = "m.tsv") -> Matrix: + path = tmp_path / name + path.write_text("\tA\tB\nTP53\t1\t2\n") + return Matrix(path=path, size_bytes=10, samples=["A", "B"], gene_count=1) + + +def test_a_matrix_waits_and_is_taken_once(tmp_path: Path) -> None: + store = PendingMatrices() + m = matrix(tmp_path) + store.put("s", m) + assert store.peek("s") is m + assert store.take("s") is m + assert store.take("s") is None + + +def test_putting_the_same_matrix_back_keeps_its_file(tmp_path: Path) -> None: + # After a reply that was not usable labels, the matrix goes back; that + # must not delete the file it is waiting with. + store = PendingMatrices() + m = matrix(tmp_path) + store.put("s", m) + store.put("s", m) + assert m.path.exists() + + +def test_a_new_upload_replaces_and_deletes_the_old(tmp_path: Path) -> None: + store = PendingMatrices() + old, new = matrix(tmp_path, "old.tsv"), matrix(tmp_path, "new.tsv") + store.put("s", old) + store.put("s", new) + assert not old.path.exists() + assert store.peek("s") is new + + +def test_an_expired_matrix_is_gone_with_its_file(tmp_path: Path) -> None: + store = PendingMatrices(wait_seconds=60) + m = matrix(tmp_path) + store.put("s", m, now=0) + assert store.peek("s", now=61) is None + assert not m.path.exists() + + +def test_dropping_a_session_deletes_its_file(tmp_path: Path) -> None: + store = PendingMatrices() + m = matrix(tmp_path) + store.put("s", m) + store.drop("s") + assert not m.path.exists() + + +@pytest.mark.parametrize( + ("reply", "expected"), + [ + ("control, control, treated, treated", True), # usable + ("control, treated", True), # wrong count: a mistake to explain + ("control; control; treated", True), + ("What does this matrix show?", False), # a question for the model + ("tell me about apoptosis", False), + ], +) +def test_what_counts_as_an_attempt_at_labels(reply: str, expected: bool) -> None: + assert chat.looks_like_labels(reply, 4) is expected