fix(agent-team): register QuestionSet with the langgraph checkpoint serializer (silence/avoid msgpack block) (#30)
This commit is contained in:
parent
a6275b4000
commit
c671e6bdbb
2 changed files with 186 additions and 12 deletions
|
|
@ -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). -----------------------------------------
|
||||
|
|
|
|||
|
|
@ -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. -------------------------------------------------
|
||||
|
||||
|
||||
|
|
|
|||
Reference in a new issue