diff --git a/bin/chat-chainlit.py b/bin/chat-chainlit.py index 3db4a54..d53e0e3 100644 --- a/bin/chat-chainlit.py +++ b/bin/chat-chainlit.py @@ -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 diff --git a/src/analysis/store.py b/src/analysis/store.py index 1e2a9ac..cc470a8 100644 --- a/src/analysis/store.py +++ b/src/analysis/store.py @@ -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: diff --git a/src/analysis/summarise.py b/src/analysis/summarise.py index 2628446..57ec669 100644 --- a/src/analysis/summarise.py +++ b/src/analysis/summarise.py @@ -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 diff --git a/src/api/analysis_summary.py b/src/api/analysis_summary.py index 85528ff..bc8c0ef 100644 --- a/src/api/analysis_summary.py +++ b/src/api/analysis_summary.py @@ -14,9 +14,9 @@ """ import asyncio +import contextlib import json import time -import uuid from collections.abc import AsyncIterator from typing import Any @@ -24,26 +24,19 @@ 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 @@ -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: @@ -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") @@ -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) @@ -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 @@ -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" @@ -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), ( @@ -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 @@ -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 "" diff --git a/src/handoff/seed.py b/src/handoff/seed.py index 7f3081c..2016366 100644 --- a/src/handoff/seed.py +++ b/src/handoff/seed.py @@ -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 @@ -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: @@ -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): @@ -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)] diff --git a/tests/analysis/test_store.py b/tests/analysis/test_store.py index b7a322d..f975496 100644 --- a/tests/analysis/test_store.py +++ b/tests/analysis/test_store.py @@ -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" diff --git a/tests/analysis/test_summarise.py b/tests/analysis/test_summarise.py index e4a37f4..83710fe 100644 --- a/tests/analysis/test_summarise.py +++ b/tests/analysis/test_summarise.py @@ -2,7 +2,13 @@ from typing import Any -from analysis.summarise import VERDICT_INSTRUCTION, prompt_input +from analysis.summarise import ( + INEXACT_COUNT_INSTRUCTION, + UNKNOWN_COUNT_INSTRUCTION, + VERDICT_INSTRUCTION, + prompt_input, + summary_instruction, +) def _payload(*fdrs: float) -> dict[str, Any]: @@ -297,3 +303,23 @@ def test_no_column_count_is_claimed_when_the_pathways_disagree() -> None: def test_a_result_with_no_expression_values_claims_no_columns() -> None: assert "expression_columns" not in prompt_input(_payload(1e-9)) assert "exp" not in prompt_input(_payload(1e-9))["pathways"][0] + + +def _input(shown: int, significant: int) -> dict[str, Any]: + return { + "verdict": "has_findings" if significant else "nothing_significant", + "significant_count_is_exact": False, + "significant_among_shown": significant, + "pathways_shown": shown, + "analysis_type": "OVERREPRESENTATION", + } + + +def test_all_significant_is_said_only_when_it_is_true() -> None: + # "Every one of them is significant" was added whenever the count was not + # exact -- including when only 2 of 12 were (review, area 1a). + assert INEXACT_COUNT_INSTRUCTION in summary_instruction(_input(12, 12)) + two = summary_instruction(_input(12, 2)) + assert INEXACT_COUNT_INSTRUCTION not in two + assert UNKNOWN_COUNT_INSTRUCTION in two + assert INEXACT_COUNT_INSTRUCTION not in summary_instruction(_input(12, 0)) diff --git a/tests/api/test_analysis_summary.py b/tests/api/test_analysis_summary.py index 2f468c7..881be7c 100644 --- a/tests/api/test_analysis_summary.py +++ b/tests/api/test_analysis_summary.py @@ -98,6 +98,7 @@ def wired(monkeypatch: pytest.MonkeyPatch) -> _Counter: # the model is never called -- which looks like the feature being broken # and is the tests interfering. monkeypatch.setattr("api.analysis_summary._store", SummaryStore()) + monkeypatch.setattr("api.analysis_summary._generating", {}) monkeypatch.setattr("api.analysis_summary.get_llm", lambda *a, **k: counter) monkeypatch.setattr( "api.analysis_summary.resolve_llm_model", lambda _c: ("openai", "m", None) @@ -819,3 +820,96 @@ async def _odd() -> str: private, public = keys start = _events(_post(public, caller_token=_token(private)).text)[0][1] assert start["release"] is None + + +# --- review, area 1a --------------------------------------------------------- + + +async def _collect(private: str, public: str) -> str: + from types import SimpleNamespace + + request = SimpleNamespace( + app=SimpleNamespace(state=SimpleNamespace(caller_token_key=public)) + ) + response = await analysis_summary( + SummaryRequest( + token=SAMPLE_TOKEN, caller_token=_token(private), disclosure="aggregate" + ), + request, # type: ignore[arg-type] + ) + iterator = cast("AsyncGenerator[str, None]", response.body_iterator) + return "".join([chunk async for chunk in iterator]) + + +def test_two_readers_asking_at_once_get_one_summary( + keys: tuple[str, str], monkeypatch: pytest.MonkeyPatch +) -> None: + # Both missed the cache and generated different texts, and the later one + # overwrote what the earlier reader had seen. + class _Different(_Counter): + async def astream(self, _messages: Any) -> AsyncIterator[Any]: + self.calls += 1 + for piece in (f"version {self.calls} ", "of the summary."): + await asyncio.sleep(0.02) + yield type("Chunk", (), {"content": piece})() + + model = _Different() + monkeypatch.setattr("api.analysis_summary.get_llm", lambda *a, **k: model) + private, public = keys + + async def both() -> list[str]: + return list( + await asyncio.gather(_collect(private, public), _collect(private, public)) + ) + + first, second = asyncio.run(both()) + assert model.calls == 1 + assert "version 1" in first + assert "version 1" in second + + +def test_content_blocks_are_read_and_an_empty_stream_is_a_failure( + keys: tuple[str, str], monkeypatch: pytest.MonkeyPatch +) -> None: + class _Blocks(_Counter): + async def astream(self, _messages: Any) -> AsyncIterator[Any]: + self.calls += 1 + yield type( + "Chunk", (), {"content": [{"type": "text", "text": "From blocks."}]} + )() + + monkeypatch.setattr("api.analysis_summary.get_llm", lambda *a, **k: _Blocks()) + private, public = keys + body = asyncio.run(_collect(private, public)) + assert "From blocks." in body + assert '"state": "summarised"' in body + + class _Silent(_Counter): + async def astream(self, _messages: Any) -> AsyncIterator[Any]: + self.calls += 1 + yield type("Chunk", (), {"content": [{"type": "reasoning"}]})() + + monkeypatch.setattr("api.analysis_summary._store", SummaryStore()) + monkeypatch.setattr("api.analysis_summary.get_llm", lambda *a, **k: _Silent()) + body = asyncio.run(_collect(private, public)) + assert '"state": "summarised"' not in body + + +def test_a_fixed_temperature_model_is_sent_its_temperature( + monkeypatch: pytest.MonkeyPatch, +) -> None: + # The summary built its own model with the default 0.0, which o3-class + # models reject; the chat sent 1.0 and worked. + from agent.graph import resolve_temperature + from api import analysis_summary as endpoint + from util.config_yml import Config + + seen: dict[str, Any] = {} + monkeypatch.setattr( + endpoint, "resolve_llm_model", lambda _c: ("openai", "o3", None) + ) + monkeypatch.setattr(Config, "from_yaml", staticmethod(lambda *a, **k: None)) + monkeypatch.setattr(endpoint, "get_llm", lambda *a, **k: seen.update(k)) + endpoint._summary_llm() + assert seen["temperature"] == resolve_temperature("o3") + assert seen["temperature"] != 0.0 diff --git a/tests/handoff/test_handoff_seed.py b/tests/handoff/test_handoff_seed.py index 6aa0920..4115e9e 100644 --- a/tests/handoff/test_handoff_seed.py +++ b/tests/handoff/test_handoff_seed.py @@ -205,3 +205,24 @@ def test_a_search_question_is_shown_as_text_not_markup() -> None: assert re.search(r"(? None: + # The chat model reads the same data for every follow-up; without the + # summary's rules it answers from input that once gave "12 of 1280". + from analysis.summarise import DATA_RULES, prompt_input + + data = prompt_input({"summary": {"type": "OVERREPRESENTATION"}, "pathways": []}) + text = str(seed.seeded_turn(handoff(), data)[1].content) + assert DATA_RULES in text + assert "Instructions that go with it" in text + + +def test_a_temporary_failure_is_not_reported_as_deletion() -> None: + # Every non-ok outcome said "no longer available", so a timeout told the + # model the reader's result had been deleted (review, area 1a). + failed = str(seed.seeded_turn(handoff(), None, outcome="failed")[1].content) + gone = str(seed.seeded_turn(handoff(), None, outcome="gone")[1].content) + assert "temporary" in failed + assert "no longer available" not in failed + assert "no longer available" in gone