fix(agent-team): register QuestionSet with the langgraph checkpoint serializer #30

Merged
amoussa1229 merged 1 commit from fix/agent-team-msgpack-questionset into main 2026-06-22 20:14:27 +00:00
2 changed files with 186 additions and 12 deletions

View file

@ -40,6 +40,7 @@ from __future__ import annotations
import uuid
from collections.abc import Callable
from contextlib import contextmanager
from datetime import datetime, timedelta, timezone
from pathlib import Path
from typing import TYPE_CHECKING, Any
@ -56,11 +57,29 @@ from agent_team.task_model import (
from agent_team.transport import QuestionSet
if TYPE_CHECKING: # pragma: no cover - typing only
from collections.abc import Iterator
from contextlib import AbstractContextManager
from langgraph.checkpoint.base import BaseCheckpointSaver
from langgraph.graph.state import CompiledStateGraph
# The custom (non-builtin) types that travel inside a LangGraph checkpoint and
# must therefore be on the msgpack deserialization allowlist (D9, §3.3.1). Today
# the only such type is the clarifier's QuestionSet, which rides in the
# ``interrupt()`` payload (graph.clarify_node) and so is msgpack-encoded by the
# checkpoint serializer as ``(module, name, kwargs)`` and reconstructed on
# resume. Without an explicit allowlist the serializer runs in permissive mode
# and logs "Deserializing unregistered type agent_team.transport.base.QuestionSet
# ... will be blocked in a future version" on every checkpoint load; passing the
# allowlist registers the type so it deserializes silently AND keeps working when
# LangGraph flips the future default to block-unregistered. Every PipelineState
# field is a primitive/dict/list and TaskStatus/Phase are stored as ``.value``
# strings, so QuestionSet is the complete set; add any new custom checkpoint type
# here when one is introduced (otherwise it would be blocked under strict mode).
_CHECKPOINT_ALLOWED_MSGPACK_TYPES: tuple[tuple[str, str], ...] = (
("agent_team.transport.base", "QuestionSet"),
)
__all__ = [
"APPROVED_ROUTE",
"BUILD_NODE",
@ -73,6 +92,7 @@ __all__ = [
"PLAN",
"REVIEW",
"VERIFY_NODE",
"build_checkpoint_serde",
"build_graph",
"build_sqlite_checkpointer",
"clarify_node",
@ -428,22 +448,61 @@ def build_graph(
return builder.compile(checkpointer=checkpointer)
def build_checkpoint_serde() -> Any:
"""Build the checkpoint serializer with the QuestionSet msgpack allowlist.
Returns a :class:`~langgraph.checkpoint.serde.jsonplus.JsonPlusSerializer`
constructed with an explicit ``allowed_msgpack_modules`` covering every
custom type that rides inside a checkpoint (today only
:class:`~agent_team.transport.QuestionSet`; see
:data:`_CHECKPOINT_ALLOWED_MSGPACK_TYPES`).
Why this exists: the default serializer runs in *permissive* msgpack mode
(``allowed_msgpack_modules=True``), which deserializes any type but logs
``"Deserializing unregistered type agent_team.transport.base.QuestionSet ...
This will be blocked in a future version"`` on every checkpoint load — and a
future LangGraph release will turn that into a hard block, breaking durable
resume. Passing the explicit allowlist is the registration path that warning
recommends: the listed type deserializes silently, and the config is
already block-clean for when the default flips. (Per
``langgraph/checkpoint/serde/jsonplus.py``: ``_create_msgpack_ext_hook``
only emits the warning while the allowlist is the ``True`` sentinel; once an
explicit collection is supplied, an allowlisted ``(module, name)`` returns
True with no warning.)
The ``langgraph.checkpoint.serde.jsonplus`` import is deferred to call time
so this module still imports cleanly where the optional checkpoint package
is absent (pre-deploy scaffolding).
"""
from langgraph.checkpoint.serde.jsonplus import JsonPlusSerializer
return JsonPlusSerializer(
allowed_msgpack_modules=list(_CHECKPOINT_ALLOWED_MSGPACK_TYPES)
)
def build_sqlite_checkpointer(
db_path: Path | str,
) -> AbstractContextManager[BaseCheckpointSaver]:
"""Construct the production SQLite checkpointer over ``db_path`` (D9, §3.3).
Returns a **context manager**, not an entered saver: in
``langgraph-checkpoint-sqlite`` ``SqliteSaver.from_conn_string`` is a
``@contextmanager`` classmethod, so the caller MUST enter it (``with`` it,
or ``__enter__`` and retain it for the graph's lifetime) before passing the
yielded saver to :func:`build_graph`. The live coordinator owns that
lifecycle (it enters the CM at setup and holds it for the daemon's life);
passing the raw return value straight into ``build_graph`` would compile a
graph whose checkpointer is an un-entered CM and break ``get_state`` /
``invoke`` at runtime. The earlier ``-> BaseCheckpointSaver`` annotation
mis-stated this contract (review FIX); the type now matches reality so a
direct caller cannot be silently misled.
Returns a **context manager**, not an entered saver: the caller MUST enter
it (``with`` it, or ``__enter__`` and retain it for the graph's lifetime)
before passing the yielded saver to :func:`build_graph`. The live
coordinator owns that lifecycle (it enters the CM at setup and holds it for
the daemon's life); passing the raw return value straight into
``build_graph`` would compile a graph whose checkpointer is an un-entered CM
and break ``get_state`` / ``invoke`` at runtime.
We open the connection and construct ``SqliteSaver(conn, serde=...)``
ourselves rather than using ``SqliteSaver.from_conn_string`` because the
latter has no seam to inject a serializer (it always builds the default,
warning-emitting one). The injected serde is :func:`build_checkpoint_serde`,
whose msgpack allowlist registers :class:`QuestionSet` so resume no longer
logs the "unregistered type" warning and stays forward-compatible with
LangGraph's coming block-by-default. The connection is opened with
``check_same_thread=False`` (matching ``from_conn_string``) and closed when
the context manager exits.
The import of ``langgraph.checkpoint.sqlite`` is deferred to call time so
this module imports cleanly in environments where that optional package is
@ -466,7 +525,20 @@ def build_sqlite_checkpointer(
db_path = Path(db_path)
db_path.parent.mkdir(parents=True, exist_ok=True)
return SqliteSaver.from_conn_string(str(db_path))
serde = build_checkpoint_serde()
@contextmanager
def _saver_cm() -> Iterator[BaseCheckpointSaver]:
import sqlite3
from contextlib import closing
# Mirror SqliteSaver.from_conn_string's connection settings, but build
# the saver with our allowlisted serde (from_conn_string offers no serde
# seam). closing() guarantees the connection is released on exit.
with closing(sqlite3.connect(str(db_path), check_same_thread=False)) as conn:
yield SqliteSaver(conn, serde=serde)
return _saver_cm()
# --- Driver seam (thread_id-keyed). -----------------------------------------

View file

@ -29,6 +29,7 @@ from agent_team.graph import (
INTAKE,
P1_PHASE_SEQUENCE,
PLAN,
build_checkpoint_serde,
build_graph,
build_sqlite_checkpointer,
clarify_node,
@ -293,6 +294,107 @@ def test_build_sqlite_checkpointer_builds_when_dep_present(tmp_path) -> None:
assert hasattr(saver, "get_next_version")
# --- Checkpoint serializer / QuestionSet msgpack allowlist (D9). ------------
# QuestionSet rides in the clarifier interrupt payload and so is msgpack-encoded
# into every checkpoint. The default serializer deserializes it but logs
# "Deserializing unregistered type agent_team.transport.base.QuestionSet ... will
# be blocked in a future version" on each load, and the coming LangGraph default
# turns that warning into a hard block (breaking durable resume). These pin that
# build_checkpoint_serde registers the type so it round-trips silently AND is
# already block-clean.
_QSET_MSGPACK_KEY = ("agent_team.transport.base", "QuestionSet")
def _capture_serde_warnings():
"""Attach a capturing handler to the serde logger; return (handler, buffer)."""
import io
import logging
buf = io.StringIO()
handler = logging.StreamHandler(buf)
handler.setLevel(logging.WARNING)
logger = logging.getLogger("langgraph.checkpoint.serde.jsonplus")
logger.addHandler(handler)
logger.setLevel(logging.WARNING)
return logger, handler, buf
def _reset_serde_warning_dedup() -> None:
"""Clear the serializer's process-wide warn-once dedup set.
jsonplus dedups "unregistered type" warnings across the process lifetime, so
an earlier test (or the assertion below) could mask a regression. Clearing the
set makes each assertion observe the live behavior, not a stale dedup.
"""
from langgraph.checkpoint.serde import jsonplus as _jp
_jp._warned_unregistered_types.clear()
_jp._warned_blocked_types.clear()
def test_build_checkpoint_serde_allowlists_questionset() -> None:
# The serde must carry QuestionSet on its msgpack allowlist (an explicit
# collection, NOT the permissive ``True`` sentinel) — that is the registration
# path the warning recommends and what makes resume forward-compatible.
pytest.importorskip("langgraph.checkpoint.serde.jsonplus")
serde = build_checkpoint_serde()
allowed = serde._allowed_msgpack_modules
assert allowed is not True # not the warn-on-everything permissive default
assert _QSET_MSGPACK_KEY in allowed
def test_questionset_round_trips_through_serde_without_warning() -> None:
# The configured serializer must round-trip a QuestionSet AND emit no
# "unregistered type" warning on deserialize.
pytest.importorskip("langgraph.checkpoint.serde.jsonplus")
_reset_serde_warning_dedup()
logger, handler, buf = _capture_serde_warnings()
try:
serde = build_checkpoint_serde()
qset = QuestionSet(
thread_id="t-1",
question_id="q-1",
turn=0,
questions=["What is in scope?"],
context={"phase": Phase.CLARIFY.value},
)
encoded = serde.dumps_typed(qset)
restored = serde.loads_typed(encoded)
finally:
logger.removeHandler(handler)
assert restored == qset
assert "unregistered type" not in buf.getvalue()
def test_questionset_round_trips_through_sqlite_checkpointer_without_warning(
tmp_path,
) -> None:
# End-to-end against the REAL production saver: drive the graph to the
# clarifier suspend (which checkpoints a QuestionSet), force a checkpoint
# load, and resume — asserting no "unregistered type" warning prints and the
# task still completes. This is the runtime regression guard for the warning.
pytest.importorskip("langgraph.checkpoint.sqlite")
_reset_serde_warning_dedup()
logger, handler, buf = _capture_serde_warnings()
try:
cm = build_sqlite_checkpointer(tmp_path / "state.db")
with cm as saver:
graph = build_graph(saver)
thread_id, _ = start_task(graph, transport="slack")
pending = pending_question(graph, thread_id=thread_id)
assert isinstance(pending["question_set"], QuestionSet)
# Force a fresh checkpoint deserialize (the warning's trigger point).
get_pipeline_state(graph, thread_id=thread_id)
resumed = resume_task(graph, thread_id=thread_id, answer="scope it")
finally:
logger.removeHandler(handler)
assert resumed["status"] == TaskStatus.DONE.value
assert "unregistered type" not in buf.getvalue()
# --- P2 review-loop wiring. -------------------------------------------------