Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
115 changes: 78 additions & 37 deletions bin/chat-chainlit.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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
Expand Down Expand Up @@ -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()

Expand All @@ -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:
Expand Down Expand Up @@ -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)
Expand Down
71 changes: 56 additions & 15 deletions src/gsa/chainlit_flow.py
Original file line number Diff line number Diff line change
Expand Up @@ -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__)
Expand Down Expand Up @@ -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.
Expand Down
16 changes: 16 additions & 0 deletions src/gsa/chat.py
Original file line number Diff line number Diff line change
Expand Up @@ -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.

Expand Down
72 changes: 72 additions & 0 deletions src/gsa/pending.py
Original file line number Diff line number Diff line change
@@ -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()
Loading
Loading