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
7 changes: 5 additions & 2 deletions bin/chat-chainlit.py
Original file line number Diff line number Diff line change
Expand Up @@ -243,11 +243,14 @@ async def continue_from_handoff(handoff_id: str) -> None:
# before it has run would otherwise seed a thread called "None".
thread_id: str = current_thread_id()
data = None
outcome = "ok"
try:
if isinstance(handoff, AnalysisHandoff):
data = await seed.analysis_data(handoff)
data, outcome = await seed.analysis_data_and_outcome(handoff)
seeded = await get_graph().seed_history(
profile, thread_id=thread_id, messages=seed.seeded_turn(handoff, data)
profile,
thread_id=thread_id,
messages=seed.seeded_turn(handoff, data, outcome=outcome),
)
except Exception:
# A failed fetch or graph update must not leave the reader in a chat
Expand Down
7 changes: 7 additions & 0 deletions src/analysis/store.py
Original file line number Diff line number Diff line change
Expand Up @@ -73,6 +73,13 @@ def put(
if not text.strip():
return stored
key = (token, release, tier)
existing = self._entries.get(key)
if existing is not None:
# The first stored summary stays. Two generations for one key can
# finish out of order; overwriting served a reader, on reload or
# in Continue in chat, a summary they were never shown (review,
# area 1a). The later reader is served this one instead.
return existing
self._entries[key] = stored
self._entries.move_to_end(key)
while len(self._entries) > self.max_entries:
Expand Down
47 changes: 47 additions & 0 deletions src/analysis/summarise.py
Original file line number Diff line number Diff line change
Expand Up @@ -249,3 +249,50 @@ def prompt_input(payload: dict[str, Any]) -> dict[str, Any]:
"state how many of the total were significant, and never imply that only "
"the pathways listed here passed."
)

#: When the count is not exact for another reason -- the shown pathways are
#: not all significant, or their order could not be confirmed -- the
#: instruction above would tell the model something false (review, 1a).
UNKNOWN_COUNT_INSTRUCTION = (
"Only the highest-ranked pathways are included here, and the number "
"significant overall cannot be determined from them. Never state how many "
"of the total were significant."
)

#: The rules for reading this data. The summary endpoint's system prompt
#: carries them, and so does a handoff's seeded turn: the chat model reads
#: the same data afterwards, and without them answers follow-ups from the
#: very input that once produced "12 significant out of 1280".
DATA_RULES = """
1. Every quantitative claim must come from the data. Never state a statistic it
does not contain.
2. Follow the verdict instruction exactly. It is computed from the data, not
guessed, and it overrides any impression the numbers give you.
3. Whenever you call a pathway significant, say whether that is before or after
multiple-testing correction.
4. Do not name a pathway that is not in the data.
5. Never state how many pathways were significant overall unless the data says
that count is exact. Only the highest-ranked are included.
""".strip()


def summary_instruction(model_input: dict[str, Any]) -> str:
"""Every instruction that goes with this data, in one place."""
instruction = VERDICT_INSTRUCTION[model_input["verdict"]]
if not model_input["significant_count_is_exact"]:
all_shown_significant = (
model_input["significant_among_shown"] == model_input["pathways_shown"]
)
instruction = f"{instruction} " + (
INEXACT_COUNT_INSTRUCTION
if all_shown_significant
else UNKNOWN_COUNT_INSTRUCTION
)
by_type = TYPE_INSTRUCTION.get(str(model_input.get("analysis_type") or "").upper())
if by_type:
instruction = f"{instruction} {by_type}"
instruction = f"{instruction} {STATISTICS_INSTRUCTION}"
instruction = f"{instruction} {UNMATCHED_INSTRUCTION}"
if model_input.get("identifiers_not_found_names"):
instruction = f"{instruction} {NAMED_UNMATCHED_INSTRUCTION}"
return instruction
122 changes: 82 additions & 40 deletions src/api/analysis_summary.py
Original file line number Diff line number Diff line change
Expand Up @@ -14,36 +14,29 @@
"""

import asyncio
import contextlib
import json
import time
import uuid
from collections.abc import AsyncIterator
from typing import Any

from fastapi import APIRouter, Request
from fastapi.responses import StreamingResponse
from pydantic import BaseModel, Field

from agent.graph import resolve_llm_model
from agent.graph import resolve_llm_model, resolve_temperature
from agent.models import get_llm
from analysis.client import current_release, fetch_not_found, fetch_result
from analysis.disclosure import Tier, for_tier
from analysis.store import Stored, SummaryStore
from analysis.summarise import (
INEXACT_COUNT_INSTRUCTION,
NAMED_UNMATCHED_INSTRUCTION,
STATISTICS_INSTRUCTION,
TYPE_INSTRUCTION,
UNMATCHED_INSTRUCTION,
VERDICT_INSTRUCTION,
prompt_input,
)
from analysis.summarise import DATA_RULES, prompt_input, summary_instruction
from util.caller_token import (
TokenRejectedError,
human_presence_detail,
human_presence_reason,
verify,
)
from util.config_yml import Config
from util.logging import logging
from util.rate_limit import identity_of, limiter_from_env

Expand All @@ -60,6 +53,9 @@
#: Process-local, lost on deploy (research D4). Module state so it outlives
#: a request, as the limiter does.
_store = SummaryStore()
#: Summaries being generated right now, so a second request for the same one
#: waits and reuses it rather than generating a different text.
_generating: dict[tuple[str, str, str], "asyncio.Future[None]"] = {}


def stored_summary(token: str, release: str, tier: str) -> Stored | None:
Expand All @@ -76,19 +72,11 @@ def stored_summary(token: str, release: str, tier: str) -> Stored | None:
ran it.

Rules, in order of importance:
1. Every quantitative claim must come from the data below. Never state a
statistic it does not contain.
2. Follow the verdict instruction exactly. It is computed from the data, not
guessed, and it overrides any impression the numbers give you.
3. Whenever you call a pathway significant, say whether that is before or
after multiple-testing correction.
4. Do not name a pathway that is not in the data below.
5. Never state how many pathways were significant overall unless the data
says that count is exact. Only the highest-ranked are included.
{rules}
6. Do not list sources or citations at the end. The interface renders them
from structured events; a list here is a duplicate.
7. Four short paragraphs at most. Plain prose for a working scientist.
""".strip()
""".strip().format(rules=DATA_RULES)


IMPLEMENTED_TIERS = ("aggregate", "identifiers")
Expand Down Expand Up @@ -177,6 +165,8 @@ async def analysis_summary(body: SummaryRequest, request: Request) -> StreamingR

async def stream() -> AsyncIterator[str]:
state = "failed"
mine: asyncio.Future[None] | None = None
key: tuple[str, str, str] = ("", "", "")
try:
async with asyncio.timeout(SUMMARY_TIMEOUT_SECONDS):
fetched = await fetch_result(body.token)
Expand Down Expand Up @@ -206,6 +196,20 @@ async def stream() -> AsyncIterator[str]:
if release
else None
)
if cached is None and release:
# Another request is generating this very summary: wait
# for it and serve its text, so both readers see one
# summary. Generating twice let the later one overwrite
# what the earlier reader had seen (review, area 1a).
key = (body.token, release, body.disclosure)
inflight = _generating.get(key)
if inflight is not None:
with contextlib.suppress(Exception):
await asyncio.shield(inflight)
cached = _store.get(body.token, release, body.disclosure)
if cached is None:
mine = asyncio.get_running_loop().create_future()
_generating[key] = mine

# Which tier the answer was actually built from.
applied: Tier = body.disclosure
Expand All @@ -225,6 +229,11 @@ async def stream() -> AsyncIterator[str]:
# disclose and silently got the other summary has
# been told nothing and given nothing.
applied = "aggregate"
# An aggregate summary may already be stored: reuse
# it rather than generate a second, different one
# (review, area 1a).
if release:
cached = _store.get(body.token, release, "aggregate")
logger.warning(
"identifier tier requested but the not-found "
"lookup returned nothing; serving aggregate"
Expand Down Expand Up @@ -275,20 +284,8 @@ async def stream() -> AsyncIterator[str]:
yield _done("summarised", time.monotonic() - started)
return

provider, model, base_url = resolve_llm_model(None)
llm = get_llm(provider, model, base_url=base_url, request_timeout=90.0)
instruction = VERDICT_INSTRUCTION[model_input["verdict"]]
if not model_input["significant_count_is_exact"]:
instruction = f"{instruction} {INEXACT_COUNT_INSTRUCTION}"
by_type = TYPE_INSTRUCTION.get(
str(model_input.get("analysis_type") or "").upper()
)
if by_type:
instruction = f"{instruction} {by_type}"
instruction = f"{instruction} {STATISTICS_INSTRUCTION}"
instruction = f"{instruction} {UNMATCHED_INSTRUCTION}"
if model_input.get("identifiers_not_found_names"):
instruction = f"{instruction} {NAMED_UNMATCHED_INSTRUCTION}"
llm = _summary_llm()
instruction = summary_instruction(model_input)
messages = [
("system", SYSTEM_PROMPT),
(
Expand All @@ -299,10 +296,13 @@ async def stream() -> AsyncIterator[str]:
]
produced: list[str] = []
async for chunk in llm.astream(messages):
text = getattr(chunk, "content", "")
if isinstance(text, str) and text:
text = _text(chunk)
if text:
produced.append(text)
yield _sse("token", {"text": text})
if not produced:
# Nothing came back: not a summary, whatever the stream said.
raise RuntimeError("the model returned no text")
if release:
_store.put(
body.token, release, applied, "".join(produced), citations
Expand All @@ -319,11 +319,53 @@ async def stream() -> AsyncIterator[str]:
# terminal event to stop waiting.
logger.exception("summarising an analysis failed")
state = "failed"
finally:
if mine is not None:
if _generating.get(key) is mine:
del _generating[key]
if not mine.done():
mine.set_result(None)
yield _done(state, time.monotonic() - started)

return StreamingResponse(stream(), media_type="text/event-stream")


def new_thread_id() -> str:
"""Unused by the stream, kept for parity with the answer endpoint's ids."""
return f"summary-{uuid.uuid4()}"
def _summary_llm() -> Any:
"""The chat's model, chosen the way the chat chooses it.

The first version passed no config and no temperature: a model or
provider set in config.yml applied to the chat but not to summaries, and
a fixed-temperature model (o3, o4-mini) rejected the default 0.0 -- every
summary failing while the chat worked (review, area 1a).
"""
config = Config.from_yaml()
llm_config = config.llm if config else None
provider, model, base_url = resolve_llm_model(llm_config)
temperature = resolve_temperature(
model, configured=llm_config.temperature if llm_config else None
)
return get_llm(
provider,
model,
base_url=base_url,
request_timeout=90.0,
temperature=temperature,
)


def _text(chunk: Any) -> str:
"""A streamed chunk's text, whether a string or content blocks.

Responses-API models stream content blocks; reading only strings dropped
all of it and still reported `summarised` (review, area 1a).
"""
content = getattr(chunk, "content", "")
if isinstance(content, str):
return content
if isinstance(content, list):
return "".join(
block.get("text", "")
for block in content
if isinstance(block, dict) and block.get("type") in ("text", "output_text")
)
return ""
39 changes: 33 additions & 6 deletions src/handoff/seed.py
Original file line number Diff line number Diff line change
Expand Up @@ -26,7 +26,7 @@

from analysis.client import Fetched, fetch_not_found, fetch_result
from analysis.disclosure import for_tier
from analysis.summarise import prompt_input
from analysis.summarise import DATA_RULES, prompt_input, summary_instruction
from handoff.store import DEFAULT_TTL_SECONDS, AnalysisHandoff, Handoff, SearchHandoff
from util.markdown import escape

Expand All @@ -51,16 +51,31 @@ async def analysis_data(
results on a new release -- in which case the chat still continues the
summary, and says it cannot see the underlying numbers.
"""
data, _ = await analysis_data_and_outcome(
handoff, fetch=fetch, fetch_unmatched=fetch_unmatched
)
return data


async def analysis_data_and_outcome(
handoff: AnalysisHandoff,
*,
fetch: FetchResult = fetch_result,
fetch_unmatched: FetchUnmatched = fetch_not_found,
) -> tuple[dict[str, Any] | None, str]:
"""The data, and the Analysis Service's outcome: "ok", "gone",
"not_found" or "failed". A failure is not a deletion, and the model is
told which (review, area 1a: a timeout said the result was gone)."""
fetched = await fetch(handoff.token)
if fetched.outcome != "ok" or fetched.result is None:
return None
return None, fetched.outcome
data = prompt_input(for_tier(fetched.result, handoff.tier))
if handoff.tier == "identifiers":
# Only at the tier the reader chose on the website, and only then.
unmatched = await fetch_unmatched(handoff.token)
if unmatched:
data["identifiers_not_found_names"] = unmatched
return data
return data, "ok"


def _sources(citations: tuple[tuple[str, str], ...]) -> str:
Expand All @@ -71,7 +86,7 @@ def _sources(citations: tuple[tuple[str, str], ...]) -> str:


def seeded_turn(
handoff: Handoff, data: dict[str, Any] | None = None
handoff: Handoff, data: dict[str, Any] | None = None, *, outcome: str = "gone"
) -> list[BaseMessage]:
"""The conversation turn the chat starts from."""
if isinstance(handoff, SearchHandoff):
Expand All @@ -81,17 +96,29 @@ def seeded_turn(
HumanMessage(handoff.question),
AIMessage(handoff.summary + _sources(handoff.citations)),
]
if data is None:
if data is None and outcome == "failed":
appendix = (
"\n\n(The analysis data behind this summary could not be fetched from "
"Reactome just now -- a temporary failure. The result itself may "
"still exist; only the summary above is known.)"
)
elif data is None:
appendix = (
"\n\n(The analysis result behind this summary is no longer available "
"from Reactome, so only the summary above is known.)"
)
else:
# The same rules and instructions the summary was written under: the
# chat model reads this data for every follow-up, and without them
# answers from the input that once produced "12 of 1280" (review, 1a).
appendix = (
"\n\nThe analysis data this summary was built from "
f"(disclosure tier: {handoff.tier}):\n"
f"```json\n{json.dumps(data, indent=1, default=str)}\n```"
f"```json\n{json.dumps(data, indent=1, default=str)}\n```\n\n"
f"Rules for answering from this data:\n{DATA_RULES}"
)
if "verdict" in data:
appendix += f"\nInstructions that go with it: {summary_instruction(data)}"
return [HumanMessage(HUMAN_TURN), AIMessage(handoff.summary + appendix)]


Expand Down
12 changes: 12 additions & 0 deletions tests/analysis/test_store.py
Original file line number Diff line number Diff line change
Expand Up @@ -77,3 +77,15 @@ def test_reading_a_summary_keeps_it_from_being_evicted() -> None:
store.put("c", "97", "aggregate", "summary c", CITES)
assert store.get("a", "97", "aggregate") is not None
assert store.get("b", "97", "aggregate") is None


def test_a_stored_summary_is_never_overwritten() -> None:
# Two generations for one key can finish out of order; overwriting served
# a reader a summary they were never shown (review, area 1a).
store = SummaryStore()
store.put("t", "97", "aggregate", "first", ())
kept = store.put("t", "97", "aggregate", "second", ())
assert kept.text == "first"
found = store.get("t", "97", "aggregate")
assert found is not None
assert found.text == "first"
Loading
Loading