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
33 changes: 26 additions & 7 deletions src/agent/graph.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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})
Expand All @@ -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)

Expand All @@ -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
Expand Down Expand Up @@ -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(
Expand All @@ -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
Expand Down
46 changes: 46 additions & 0 deletions src/agent/history.py
Original file line number Diff line number Diff line change
@@ -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]
11 changes: 10 additions & 1 deletion src/agent/models.py
Original file line number Diff line number Diff line change
Expand Up @@ -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: (
Expand All @@ -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":
Expand Down
3 changes: 2 additions & 1 deletion src/agent/profiles/base.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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,
)
Expand Down
3 changes: 2 additions & 1 deletion src/agent/profiles/plantreactome.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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"])]
),
Expand Down
7 changes: 4 additions & 3 deletions src/agent/profiles/react_to_me.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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,
),
Expand Down Expand Up @@ -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,
)
Expand Down Expand Up @@ -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"])]
),
Expand Down
20 changes: 19 additions & 1 deletion src/retrievers/csv_chroma.py
Original file line number Diff line number Diff line change
Expand Up @@ -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.

Expand Down Expand Up @@ -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"
Expand Down
91 changes: 91 additions & 0 deletions tests/agent/test_history_and_resources.py
Original file line number Diff line number Diff line change
@@ -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]
Loading