From 175790f657b3411ed28c3c37dc14c782254c3792 Mon Sep 17 00:00:00 2001 From: Adam Wright Date: Sun, 4 Oct 2026 03:05:41 +0000 Subject: [PATCH] Bound per-conversation and per-query cost in the agent (review, area 3) - History sent to the model is the recent part: up to 40 messages and 60K characters, plus a handoff's seeded first turn. Nothing was trimmed; threads overflowed the context window (8K-char messages at turn 47, an 8K-emoji message by turn 6) and then failed every turn. - BM25 scores a deduplicated query of at most 64 tokens. rank_bm25 scans every document per query token, repeats included: ~16s of CPU for a 2,000-character question, slowing every session. Documents are scored as before. - OpenAI embeddings have a 30s timeout; with none, a stalled endpoint pinned shared worker threads until every to_thread queued. - The graph compiles once, under a lock: concurrent first requests each opened a Postgres pool. thread_holds_analysis and forget_thread initialize it instead of answering 'no' or doing nothing -- after a restart the web-search guard on analysis threads failed open. Answer sweep 16/16 and the analysis handoff browser test pass on this branch. Co-Authored-By: Claude Opus 5.5 --- src/agent/graph.py | 33 ++++++-- src/agent/history.py | 46 ++++++++++++ src/agent/models.py | 11 ++- src/agent/profiles/base.py | 3 +- src/agent/profiles/plantreactome.py | 3 +- src/agent/profiles/react_to_me.py | 7 +- src/retrievers/csv_chroma.py | 20 ++++- tests/agent/test_history_and_resources.py | 91 +++++++++++++++++++++++ 8 files changed, 200 insertions(+), 14 deletions(-) create mode 100644 src/agent/history.py create mode 100644 tests/agent/test_history_and_resources.py diff --git a/src/agent/graph.py b/src/agent/graph.py index 5d24c9c4..2203afd4 100644 --- a/src/agent/graph.py +++ b/src/agent/graph.py @@ -370,6 +370,20 @@ def __del__(self) -> None: "application shutdown hook." ) + async def _ensure_graph(self) -> dict[str, CompiledStateGraph]: + """Compile once. Concurrent first requests each compiled the graph and + opened their own Postgres pool, and close_pool closed only the last + (review, area 3).""" + if self.graph is not None: + return self.graph + lock = getattr(self, "_init_lock", None) + if lock is None: + lock = self._init_lock = asyncio.Lock() + async with lock: + if self.graph is None: + self.graph = await self.initialize() + return self.graph + async def initialize(self) -> dict[str, CompiledStateGraph]: checkpointer: BaseCheckpointSaver[str] = await self.create_checkpointer() self.checkpointer = checkpointer @@ -409,7 +423,12 @@ async def thread_holds_analysis(self, profile: str, thread_id: str) -> bool: Read from the thread's own history, so it holds for as long as the history does -- across reconnects and restarts. """ - if self.graph is None or profile not in self.graph: + # Built here rather than answering "no": the app's startup never + # initializes the compiled graph, so after a restart this said "no + # analysis" until something else did -- and web search ran on a + # thread seeded with a reader's analysis (review, area 3). + self.graph = await self._ensure_graph() + if profile not in self.graph: return False state = await self.graph[profile].aget_state( RunnableConfig(configurable={"thread_id": thread_id}) @@ -428,6 +447,9 @@ async def forget_thread(self, thread_id: str) -> None: in the MemorySaver for the life of the process, which also serves the chat (review, area 1a: 600 answers grew RSS by 22 MiB). """ + # Initialized first: before anything else had built the graph this was + # a silent no-op, and on Postgres the thread stayed (review, area 3). + self.graph = await self._ensure_graph() if self.checkpointer is not None: await self.checkpointer.adelete_thread(thread_id) @@ -453,8 +475,7 @@ async def astream_answer( tags. What does separate them is order -- the expander runs *inside* retrieval, the answer after it. So a retriever completing is the boundary. """ - if self.graph is None: - self.graph = await self.initialize() + self.graph = await self._ensure_graph() if profile not in self.graph: yield AnswerEvent(kind="done", state="failed") return @@ -577,8 +598,7 @@ async def seed_history( Returns False, and changes nothing, for an unknown profile. """ - if self.graph is None: - self.graph = await self.initialize() + self.graph = await self._ensure_graph() if profile not in self.graph: return False await self.graph[profile].aupdate_state( @@ -597,8 +617,7 @@ async def ainvoke( thread_id: str, enable_postprocess: bool = True, ) -> OutputState: - if self.graph is None: - self.graph = await self.initialize() + self.graph = await self._ensure_graph() if profile not in self.graph: return OutputState() # ainvoke is typed dict[str, Any] | Any; the graph's output schema is diff --git a/src/agent/history.py b/src/agent/history.py new file mode 100644 index 00000000..0cbb5bb4 --- /dev/null +++ b/src/agent/history.py @@ -0,0 +1,46 @@ +"""How much of a conversation goes back to the model each turn. + +None of it was trimmed: every turn resent the whole thread to the rephraser, +the answer model and the live loop. A thread of 8,000-character messages +overflowed the context window at turn 47, short questions at turn 138, and a +message of 8,000 emoji -- about 24k tokens, under every character cap -- by +turn 6; after that every turn failed, and a logged-in thread resumed broken +(review, area 3). + +Recent turns are kept, by count and by size, and so is a handoff's seeded +first turn: it is the summary and data the whole conversation is about, and +the rules for reading them. +""" + +from collections.abc import Sequence + +from langchain_core.messages import BaseMessage + +MAX_MESSAGES = 40 +#: Characters, not tokens: no tokenizer is needed to bound it, and at ~1-4 +#: characters a token this stays well inside a 128k window. +MAX_CHARS = 60_000 +SEED_MARK = "reactome_analysis_seed" + + +def _size(message: BaseMessage) -> int: + return len(str(message.content)) + + +def recent(history: Sequence[BaseMessage] | None) -> list[BaseMessage]: + """The recent part of a conversation, plus a seeded first turn.""" + messages = list(history or []) + seeded = ( + messages[:2] + if any(getattr(m, "additional_kwargs", {}).get(SEED_MARK) for m in messages[:2]) + else [] + ) + rest = messages[len(seeded) :] + budget = MAX_CHARS - sum(_size(m) for m in seeded) + kept: list[BaseMessage] = [] + for message in reversed(rest): + if len(kept) >= MAX_MESSAGES or _size(message) > budget: + break + kept.append(message) + budget -= _size(message) + return seeded + kept[::-1] diff --git a/src/agent/models.py b/src/agent/models.py index cf7f74c9..d524c868 100644 --- a/src/agent/models.py +++ b/src/agent/models.py @@ -7,6 +7,9 @@ from langchain_openai.chat_models.base import ChatOpenAI from langchain_openai.embeddings import OpenAIEmbeddings +#: Seconds for one embedding request. +EMBEDDING_TIMEOUT_SECONDS = 30.0 + def get_embedding( provider: ( @@ -25,7 +28,13 @@ def get_embedding( if model is None: provider, model = provider.split("/", 1) if provider == "openai": - return OpenAIEmbeddings(model=model, base_url=base_url) + # With a timeout. The client default is none, and embedding calls run + # in the shared thread pool: a stalled endpoint pinned ~5 workers per + # question, which no cancellation could free, until every to_thread + # in the process queued behind them (review, area 3). + return OpenAIEmbeddings( + model=model, base_url=base_url, timeout=EMBEDDING_TIMEOUT_SECONDS + ) if provider == "huggingfacehub": return HuggingFaceEndpointEmbeddings(model=model) if provider == "huggingfacelocal": diff --git a/src/agent/profiles/base.py b/src/agent/profiles/base.py index 3045c7ee..ff13fed8 100644 --- a/src/agent/profiles/base.py +++ b/src/agent/profiles/base.py @@ -7,6 +7,7 @@ from langchain_core.runnables import Runnable, RunnableConfig from langgraph.graph.message import add_messages +from agent.history import recent from agent.tasks.detect_language import create_language_detector from agent.tasks.rephrase import create_rephrase_chain from agent.tasks.safety_checker import SafetyCheck, create_safety_checker @@ -54,7 +55,7 @@ async def preprocess(self, state: BaseState, config: RunnableConfig) -> BaseStat rephrased_input: str = await self.rephrase_chain.ainvoke( { "user_input": state["user_input"], - "chat_history": state.get("chat_history", []), + "chat_history": recent(state.get("chat_history")), }, config, ) diff --git a/src/agent/profiles/plantreactome.py b/src/agent/profiles/plantreactome.py index 8ebe7be0..ccf57e53 100644 --- a/src/agent/profiles/plantreactome.py +++ b/src/agent/profiles/plantreactome.py @@ -6,6 +6,7 @@ from langchain_core.runnables import Runnable, RunnableConfig from langgraph.graph.state import StateGraph +from agent.history import recent from agent.profiles.base import BaseGraphBuilder, BaseState from agent.tasks.unsafe_question import create_unsafe_answer_generator from retrievers.plantreactome.rag import create_plantreactome_rag @@ -85,7 +86,7 @@ async def call_model( # anything folded into it reaches BM25 and the query expander. "detected_language": state["detected_language"], "chat_history": ( - state["chat_history"] + recent(state["chat_history"]) if state["chat_history"] else [HumanMessage(state["user_input"])] ), diff --git a/src/agent/profiles/react_to_me.py b/src/agent/profiles/react_to_me.py index 3fcef3ee..a38ca9d8 100644 --- a/src/agent/profiles/react_to_me.py +++ b/src/agent/profiles/react_to_me.py @@ -8,6 +8,7 @@ from langchain_core.runnables import Runnable, RunnableConfig from langgraph.graph.state import StateGraph +from agent.history import recent from agent.profiles.base import BaseGraphBuilder, BaseState from agent.tasks.intent_classifier import ( QueryIntent, @@ -154,7 +155,7 @@ async def preprocess( self.rephrase_chain.ainvoke( { "user_input": state["user_input"], - "chat_history": state.get("chat_history", []), + "chat_history": recent(state.get("chat_history")), }, config, ), @@ -230,7 +231,7 @@ async def _answer_from_live_services( tools, state["rephrased_input"], language=state["detected_language"], - chat_history=state["chat_history"] or None, + chat_history=recent(state["chat_history"]) or None, config=config, report=report, ) @@ -293,7 +294,7 @@ async def generate_answer( # the query expander. "detected_language": state["detected_language"], "chat_history": ( - state["chat_history"] + recent(state["chat_history"]) if state["chat_history"] else [HumanMessage(state["user_input"])] ), diff --git a/src/retrievers/csv_chroma.py b/src/retrievers/csv_chroma.py index 19f38680..176a8d1c 100644 --- a/src/retrievers/csv_chroma.py +++ b/src/retrievers/csv_chroma.py @@ -29,6 +29,24 @@ logger = logging.getLogger(__name__) +#: Distinct query tokens BM25 scores. rank_bm25 scans every document once per +#: query token, repeats included: ~45 ms each over Release 97, so a 2,000- +#: character question cost ~16 s of CPU and an 8,000-character chat message +#: ~60 s, slowing every session (review, area 3). Questions are a few dozen +#: tokens; this bounds the hostile case and leaves real ones untouched. +MAX_QUERY_TOKENS = 64 + + +class BoundedBM25Retriever(BM25Retriever): + """BM25 with the query -- only the query -- deduplicated and capped.""" + + def _get_relevant_documents( + self, query: str, *, run_manager: CallbackManagerForRetrieverRun + ) -> list[Document]: + tokens = list(dict.fromkeys(self.preprocess_func(query)))[:MAX_QUERY_TOKENS] + return list(self.vectorizer.get_top_n(tokens, self.docs, n=self.k)) + + def chroma_settings() -> chromadb.config.Settings: """A *fresh* Settings object for every Chroma store. @@ -421,7 +439,7 @@ def from_subdirectory( file_path=str(csv_path), metadata_columns=_csv_column_names(csv_path) ) data = loader.load() - bm25_retriever = BM25Retriever.from_documents( + bm25_retriever = BoundedBM25Retriever.from_documents( data, preprocess_func=lambda text: word_tokenize( text.casefold(), language="english" diff --git a/tests/agent/test_history_and_resources.py b/tests/agent/test_history_and_resources.py new file mode 100644 index 00000000..9990fad8 --- /dev/null +++ b/tests/agent/test_history_and_resources.py @@ -0,0 +1,91 @@ +"""Bounds on what one conversation or one query can cost (review, area 3).""" + +import asyncio +from typing import Any + +import pytest +from langchain_core.documents import Document +from langchain_core.messages import AIMessage, BaseMessage, HumanMessage + +from agent import history +from agent.graph import AgentGraph + + +def test_recent_keeps_the_latest_turns_by_count() -> None: + turns = [HumanMessage(f"q{i}") for i in range(100)] + kept = history.recent(turns) + assert len(kept) == history.MAX_MESSAGES + assert kept[-1].content == "q99" + + +def test_recent_keeps_within_a_size_budget() -> None: + # 8,000-character messages overflowed the context window at turn 47. + turns = [HumanMessage("x" * 8000) for _ in range(30)] + kept = history.recent(turns) + assert sum(len(str(m.content)) for m in kept) <= history.MAX_CHARS + assert kept + assert kept[-1] is turns[-1] + + +def test_a_seeded_first_turn_is_always_kept() -> None: + seed: list[BaseMessage] = [ + HumanMessage("Summarise my analysis."), + AIMessage("The summary.", additional_kwargs={history.SEED_MARK: True}), + ] + later: list[BaseMessage] = [HumanMessage(f"q{i}") for i in range(100)] + kept = history.recent(seed + later) + assert kept[:2] == seed + assert kept[-1].content == "q99" + + +def test_recent_of_nothing_is_nothing() -> None: + assert history.recent(None) == [] + + +def test_bm25_scores_a_bounded_deduplicated_query() -> None: + # ~45 ms of CPU per query token over Release 97, repeats included. + from retrievers.csv_chroma import MAX_QUERY_TOKENS, BoundedBM25Retriever + + docs = [Document(page_content=t) for t in ("cdk5 tau", "apoptosis", "cell cycle")] + retriever = BoundedBM25Retriever.from_documents(docs, preprocess_func=str.split) + seen: list[list[str]] = [] + real = retriever.vectorizer.get_top_n + + def spy(tokens: list[str], documents: Any, n: int) -> Any: + seen.append(list(tokens)) + return real(tokens, documents, n=n) + + retriever.vectorizer.get_top_n = spy + words = " ".join(f"w{i}" for i in range(500)) + retriever.invoke("cdk5 cdk5 cdk5 " + words) + assert len(seen[-1]) == MAX_QUERY_TOKENS + assert seen[-1].count("cdk5") == 1 + # An ordinary question is scored exactly as before. + assert retriever.invoke("cdk5 tau")[0].page_content == "cdk5 tau" + + +def test_embedding_requests_have_a_timeout(monkeypatch: pytest.MonkeyPatch) -> None: + from agent.models import EMBEDDING_TIMEOUT_SECONDS, get_embedding + + monkeypatch.setenv("OPENAI_API_KEY", "sk-test") + embeddings = get_embedding("openai", "text-embedding-3-large") + assert getattr(embeddings, "request_timeout", None) == EMBEDDING_TIMEOUT_SECONDS + + +def test_the_graph_is_compiled_once_under_concurrency() -> None: + graph = AgentGraph.__new__(AgentGraph) + graph.graph = None + calls: list[int] = [] + + async def initialize() -> dict[str, Any]: + calls.append(1) + await asyncio.sleep(0.01) + return {"p": object()} + + graph.initialize = initialize # type: ignore[method-assign] + + async def many() -> None: + await asyncio.gather(*(graph._ensure_graph() for _ in range(5))) + + asyncio.run(many()) + assert calls == [1]