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
29 changes: 17 additions & 12 deletions bin/chat-chainlit.py
Original file line number Diff line number Diff line change
Expand Up @@ -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()

Expand Down Expand Up @@ -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))

Expand Down
13 changes: 13 additions & 0 deletions src/agent/graph.py
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand Down Expand Up @@ -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()
Expand Down Expand Up @@ -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,
Expand Down
17 changes: 16 additions & 1 deletion src/api/answer.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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)
Expand Down
4 changes: 3 additions & 1 deletion src/api/handoff.py
Original file line number Diff line number Diff line change
Expand Up @@ -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})


Expand Down
5 changes: 3 additions & 2 deletions src/handoff/window.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand All @@ -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

Expand Down
12 changes: 12 additions & 0 deletions src/util/caller_token.py
Original file line number Diff line number Diff line change
Expand Up @@ -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."""

Expand Down Expand Up @@ -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(
Expand Down Expand Up @@ -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 ----------------------
Expand Down
30 changes: 28 additions & 2 deletions tests/api/test_answer_endpoint.py
Original file line number Diff line number Diff line change
Expand Up @@ -37,22 +37,29 @@ 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 "),
AnswerEvent(kind="token", text="phosphorylates tau."),
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")

Expand Down Expand Up @@ -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
7 changes: 7 additions & 0 deletions tests/handoff/test_handoff_window.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
24 changes: 24 additions & 0 deletions tests/util/test_caller_token.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Loading