From 5b28d55a4dd2915a369291884137c850822e45e3 Mon Sep 17 00:00:00 2001 From: Adam Wright Date: Sat, 3 Oct 2026 12:09:30 +0000 Subject: [PATCH] Robustness fixes from review (area 1a) - /api/answer deletes its one-shot graph thread afterwards. Each answer left ~15 KiB of checkpoints for the life of the process, which also serves the chat (measured: 600 answers grew RSS by 22 MiB). - verify() refuses every malformed token. InvalidKeyError -- a header naming the allowed algorithm that does not match the key -- is not an InvalidTokenError and became an unauthenticated 500 on three routes. Tokens over 4,096 characters are refused before parsing. - One handoff claim per session. The window-message channel accepts posts from any page and has no rate limit; every distinct id was redeemed, each costing a message, a log line and state. The claimed id is kept as a string: a set in user_session does not survive being saved as JSON. - A claim runs as a task, like a message, so sending is blocked while the analysis data loads. A question asked meanwhile ran on the same thread and the seed was lost, while the chat said it was continuing. - Claim ids are fullmatched; `^...$` with match() accepted a newline. - Handoff refusals and claims log their reason in the message: the only formatter does not print `extra=`, so every refusal read the same. Co-Authored-By: Claude Opus 5.5 --- bin/chat-chainlit.py | 29 ++++++++++++++++----------- src/agent/graph.py | 13 ++++++++++++ src/api/answer.py | 17 +++++++++++++++- src/api/handoff.py | 4 +++- src/handoff/window.py | 5 +++-- src/util/caller_token.py | 12 +++++++++++ tests/api/test_answer_endpoint.py | 30 ++++++++++++++++++++++++++-- tests/handoff/test_handoff_window.py | 7 +++++++ tests/util/test_caller_token.py | 24 ++++++++++++++++++++++ 9 files changed, 123 insertions(+), 18 deletions(-) diff --git a/bin/chat-chainlit.py b/bin/chat-chainlit.py index 4eb7775..3db4a54 100644 --- a/bin/chat-chainlit.py +++ b/bin/chat-chainlit.py @@ -255,19 +255,17 @@ async def continue_from_handoff(handoff_id: str) -> None: logger.exception("handoff seeding failed") seeded = False if not seeded: - logger.warning("handoff not seeded", extra={"profile": profile}) + logger.warning("handoff not seeded: profile %r", profile) await cl.Message(content=seed.UNAVAILABLE).send() return if isinstance(handoff, AnalysisHandoff): cl.user_session.set("analysis_seeded", True) logger.info( - "handoff claimed", - extra={ - "kind": handoff.kind, - "tier": getattr(handoff, "tier", None), - "with_data": data is not None, - }, + "handoff claimed: %s, tier %s, with data %s", + handoff.kind, + getattr(handoff, "tier", None), + data is not None, ) await cl.Message(content=seed.shown_to_reader(handoff)).send() @@ -460,11 +458,18 @@ async def on_window_message(message: object) -> None: if handoff_id is None: return - claimed: set[str] = cl.user_session.get("handoff_claimed") or set() - if handoff_id not in claimed: - claimed.add(handoff_id) - cl.user_session.set("handoff_claimed", claimed) - await continue_from_handoff(handoff_id) + # One claim per session: a tab opens on one handoff. Any page can post + # into this channel and it has no rate limit, so redeeming every distinct + # id let one session send unbounded messages, log lines and state (review, + # area 1a). Later ids are acknowledged -- so the tab's retries stop -- and + # otherwise ignored. A string, not a set: user_session is saved as JSON. + if cl.user_session.get("handoff_claimed") is None: + cl.user_session.set("handoff_claimed", handoff_id) + # As a task, like a message: sending is blocked and Stop is shown while + # the summary's data loads. Otherwise a question asked meanwhile ran on + # the same thread, and the seed landed beside it and was lost -- while + # the chat still said it was continuing from the summary. + run_as_task(lambda: continue_from_handoff(handoff_id)) await cl.send_window_message(acknowledgement(handoff_id)) diff --git a/src/agent/graph.py b/src/agent/graph.py index e770430..574ef18 100644 --- a/src/agent/graph.py +++ b/src/agent/graph.py @@ -328,6 +328,7 @@ def __init__( # The following are set asynchronously by calling initialize() self.graph: dict[str, CompiledStateGraph] | None = None self.pool: AsyncConnectionPool[AsyncConnection[dict[str, Any]]] | None = None + self.checkpointer: BaseCheckpointSaver[str] | None = None def __del__(self) -> None: """Close the connection pool if nothing else did. @@ -370,6 +371,7 @@ def __del__(self) -> None: async def initialize(self) -> dict[str, CompiledStateGraph]: checkpointer: BaseCheckpointSaver[str] = await self.create_checkpointer() + self.checkpointer = checkpointer return { profile: graph.compile(checkpointer=checkpointer) for profile, graph in self.uncompiled_graph.items() @@ -400,6 +402,17 @@ async def close_pool(self) -> None: if self.pool: await self.pool.close() + async def forget_thread(self, thread_id: str) -> None: + """Delete a one-shot thread's checkpoints. + + The search page's answers each run on a fresh thread that nothing + reads again. Left in place, every answer kept ~15 KiB of checkpoints + in the MemorySaver for the life of the process, which also serves the + chat (review, area 1a: 600 answers grew RSS by 22 MiB). + """ + if self.checkpointer is not None: + await self.checkpointer.adelete_thread(thread_id) + async def astream_answer( self, user_input: str, diff --git a/src/api/answer.py b/src/api/answer.py index 71855a9..92ada65 100644 --- a/src/api/answer.py +++ b/src/api/answer.py @@ -40,6 +40,16 @@ router = APIRouter() +#: Deletions in flight, so they are not garbage-collected before they run. +_forgetting: set[asyncio.Task[None]] = set() + + +def _forget(graph: Any, thread_id: str) -> None: + task = asyncio.get_running_loop().create_task(graph.forget_thread(thread_id)) + _forgetting.add(task) + task.add_done_callback(_forgetting.discard) + + PROFILE = "react-to-me" # One limiter for the process, built at import so the window is not reset by a @@ -135,12 +145,13 @@ async def stream() -> AsyncIterator[str]: # text after stripping, because that is what the page renders. shown: list[str] = [] cited: list[tuple[str, str]] = [] + thread_id = f"search-{uuid.uuid4()}" try: async with asyncio.timeout(ANSWER_TIMEOUT_SECONDS): async for event in graph.astream_answer( body.question, PROFILE, - thread_id=f"search-{uuid.uuid4()}", + thread_id=thread_id, # The postprocess node runs a Tavily web search after the # answer, and `astream_answer` has no event to carry the # result -- so on this path it was paid for and discarded, @@ -217,6 +228,10 @@ async def stream() -> AsyncIterator[str]: # terminal event to stop waiting. logger.exception("answering %r failed", body.question[:80]) state = "failed" + finally: + # Nothing reads this thread again. A task, not an await: on a + # hang-up this runs during cancellation, where awaiting is unsafe. + _forget(graph, thread_id) held = sources.feed(stripper.flush()) + sources.flush() if held: shown.append(held) diff --git a/src/api/handoff.py b/src/api/handoff.py index dd2b7d2..c696341 100644 --- a/src/api/handoff.py +++ b/src/api/handoff.py @@ -79,7 +79,9 @@ class SearchHandoffRequest(BaseModel): def _refuse(status: int, reason: str, log: str) -> JSONResponse: - logger.info("handoff refused", extra={"reason": reason, "detail": log}) + # In the message, not `extra=`: the only formatter prints the message, so + # every refusal logged a bare "handoff refused" (review, area 1a). + logger.info("handoff refused: %s (%s)", reason, log) return JSONResponse(status_code=status, content={"reason": reason}) diff --git a/src/handoff/window.py b/src/handoff/window.py index 59ebfd0..cd777b2 100644 --- a/src/handoff/window.py +++ b/src/handoff/window.py @@ -33,7 +33,7 @@ #: At least 128 bits of URL-safe randomness (22 base64url characters), and a #: ceiling so a hostile page cannot post something enormous. -_ID = re.compile(r"^[A-Za-z0-9_-]{22,128}$") +_ID = re.compile(r"[A-Za-z0-9_-]{22,128}") def claimed_id(message: Any) -> str | None: @@ -47,7 +47,8 @@ def claimed_id(message: Any) -> str | None: if message.get("type") != CLAIM_TYPE: return None value = message.get("id") - if not isinstance(value, str) or not _ID.match(value): + # fullmatch: `^...$` with match() accepted a trailing newline. + if not isinstance(value, str) or not _ID.fullmatch(value): return None return value diff --git a/src/util/caller_token.py b/src/util/caller_token.py index 7d697ab..cb08300 100644 --- a/src/util/caller_token.py +++ b/src/util/caller_token.py @@ -42,6 +42,9 @@ import jwt ALGORITHMS = ["EdDSA", "RS256"] +#: A real caller token is a few hundred characters. Bounding it bounds the +#: parse an unauthenticated request can force. +MAX_TOKEN_CHARS = 4096 """What the website may sign with. Both are asymmetric; no HS* symmetric option is offered, because accepting one would let a leaked verifying key mint tokens.""" @@ -108,6 +111,8 @@ def verify(token: str, verifying_key: str, *, audience: str | None = None) -> di """ if not token: raise TokenRejectedError("no token presented") + if len(token) > MAX_TOKEN_CHARS: + raise TokenRejectedError("token too long") expected = audience or expected_audience() try: return dict( @@ -142,6 +147,13 @@ def verify(token: str, verifying_key: str, *, audience: str | None = None) -> di ) from exc except jwt.InvalidTokenError as exc: raise TokenRejectedError(f"invalid token: {type(exc).__name__}") from exc + except (jwt.PyJWTError, ValueError, TypeError, RecursionError) as exc: + # Not every rejection is an InvalidTokenError: a header naming an + # allowed algorithm that does not match the key raises InvalidKeyError, + # and a deeply nested header RecursionError. Uncaught, an + # unauthenticated request turned into a 500 on all three routes + # (review, area 1a). Every failure path refuses. + raise TokenRejectedError(f"unusable token: {type(exc).__name__}") from exc # --- human presence, for the analysis-summary endpoint ---------------------- diff --git a/tests/api/test_answer_endpoint.py b/tests/api/test_answer_endpoint.py index 2cb7c18..c119c0c 100644 --- a/tests/api/test_answer_endpoint.py +++ b/tests/api/test_answer_endpoint.py @@ -37,6 +37,8 @@ class _StubGraph: def __init__(self, events: list[AnswerEvent] | None = None) -> None: self.calls = 0 + self.threads: list[str] = [] + self.forgotten: list[str] = [] self._events = events or [ AnswerEvent(kind="citation", st_id="R-HSA-1", display_name="Apoptosis"), AnswerEvent(kind="token", text="CDK5 "), @@ -44,15 +46,20 @@ def __init__(self, events: list[AnswerEvent] | None = None) -> None: AnswerEvent(kind="done", state="answered"), ] - async def astream_answer(self, *_a: Any, **_k: Any) -> AsyncIterator[AnswerEvent]: + async def astream_answer(self, *_a: Any, **k: Any) -> AsyncIterator[AnswerEvent]: self.calls += 1 + self.threads.append(k["thread_id"]) for event in self._events: yield event + async def forget_thread(self, thread_id: str) -> None: + self.forgotten.append(thread_id) + class _ExplodingGraph(_StubGraph): - async def astream_answer(self, *_a: Any, **_k: Any) -> AsyncIterator[AnswerEvent]: + async def astream_answer(self, *_a: Any, **k: Any) -> AsyncIterator[AnswerEvent]: self.calls += 1 + self.threads.append(k["thread_id"]) yield AnswerEvent(kind="token", text="partial ") raise RuntimeError("upstream died mid-answer") @@ -612,3 +619,22 @@ def test_a_failure_gets_no_id( done, _ = self.ask(keys) assert done["state"] == "failed" assert "answer_id" not in done + + +@pytest.mark.parametrize("graph_class", [_StubGraph, _ExplodingGraph]) +def test_the_one_shot_thread_is_deleted_afterwards( + keys: tuple[str, str], + monkeypatch: pytest.MonkeyPatch, + graph_class: type[_StubGraph], +) -> None: + # Each answer runs on a fresh thread nothing reads again. Kept, every one + # stayed in the checkpointer for the life of the process (review, 1a). + graph = graph_class() + monkeypatch.setattr("api.answer.get_graph", lambda: graph) + private, public = keys + _client(public).post( + f"{PREFIX}/answer", + json={"question": "what is CDK5", "caller_token": _token(private)}, + ) + assert graph.threads, "the graph was never asked" + assert graph.forgotten == graph.threads diff --git a/tests/handoff/test_handoff_window.py b/tests/handoff/test_handoff_window.py index c7fa14a..fa726e1 100644 --- a/tests/handoff/test_handoff_window.py +++ b/tests/handoff/test_handoff_window.py @@ -62,3 +62,10 @@ def test_the_acknowledgement_names_the_id_it_answers() -> None: # Two tabs may be retrying at once; each stops only on its own ack. ack = window.acknowledgement(VALID) assert ack == {"type": window.ACK_TYPE, "id": VALID} + + +def test_a_trailing_newline_is_not_part_of_an_id() -> None: + # `^...$` with match() let "id\n" through (review, area 1a). + good = "A" * 22 + assert window.claimed_id({"type": window.CLAIM_TYPE, "id": good}) == good + assert window.claimed_id({"type": window.CLAIM_TYPE, "id": good + "\n"}) is None diff --git a/tests/util/test_caller_token.py b/tests/util/test_caller_token.py index df747e9..bb8678b 100644 --- a/tests/util/test_caller_token.py +++ b/tests/util/test_caller_token.py @@ -258,3 +258,27 @@ def test_the_detail_admits_when_it_has_no_cause_rather_than_inventing_one() -> N valid = {"human": True, "human_iat": now - 10} assert human_presence_reason(valid, now) is None assert "out of step" in human_presence_detail(valid, now) + + +def _b64(data: object) -> str: + return base64.urlsafe_b64encode(json.dumps(data).encode()).rstrip(b"=").decode() + + +def test_a_header_naming_the_other_algorithm_is_refused_not_raised( + keys: tuple[str, str], +) -> None: + # RS256 is allowed, the key is Ed25519: PyJWT raises InvalidKeyError, which + # is not an InvalidTokenError. Uncaught, an unauthenticated request was a + # 500 on all three routes (review, area 1a). + _, public = keys + forged = f'{_b64({"alg": "RS256", "typ": "JWT"})}.{_b64({"aud": DEFAULT_AUDIENCE, "exp": 9999999999})}.AAAA' + with pytest.raises(TokenRejectedError): + verify(forged, public) + + +def test_an_oversized_token_is_refused_before_it_is_parsed( + keys: tuple[str, str], +) -> None: + _, public = keys + with pytest.raises(TokenRejectedError, match="too long"): + verify("a" * 5000, public)