fix(agent-team): register QuestionSet with the langgraph checkpoint serializer (silence/avoid msgpack block)
This commit is contained in:
parent
b6a12af477
commit
2e09bb9a6e
2 changed files with 186 additions and 12 deletions
|
|
@ -40,6 +40,7 @@ from __future__ import annotations
|
||||||
|
|
||||||
import uuid
|
import uuid
|
||||||
from collections.abc import Callable
|
from collections.abc import Callable
|
||||||
|
from contextlib import contextmanager
|
||||||
from datetime import datetime, timedelta, timezone
|
from datetime import datetime, timedelta, timezone
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import TYPE_CHECKING, Any
|
from typing import TYPE_CHECKING, Any
|
||||||
|
|
@ -56,11 +57,29 @@ from agent_team.task_model import (
|
||||||
from agent_team.transport import QuestionSet
|
from agent_team.transport import QuestionSet
|
||||||
|
|
||||||
if TYPE_CHECKING: # pragma: no cover - typing only
|
if TYPE_CHECKING: # pragma: no cover - typing only
|
||||||
|
from collections.abc import Iterator
|
||||||
from contextlib import AbstractContextManager
|
from contextlib import AbstractContextManager
|
||||||
|
|
||||||
from langgraph.checkpoint.base import BaseCheckpointSaver
|
from langgraph.checkpoint.base import BaseCheckpointSaver
|
||||||
from langgraph.graph.state import CompiledStateGraph
|
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__ = [
|
__all__ = [
|
||||||
"APPROVED_ROUTE",
|
"APPROVED_ROUTE",
|
||||||
"BUILD_NODE",
|
"BUILD_NODE",
|
||||||
|
|
@ -73,6 +92,7 @@ __all__ = [
|
||||||
"PLAN",
|
"PLAN",
|
||||||
"REVIEW",
|
"REVIEW",
|
||||||
"VERIFY_NODE",
|
"VERIFY_NODE",
|
||||||
|
"build_checkpoint_serde",
|
||||||
"build_graph",
|
"build_graph",
|
||||||
"build_sqlite_checkpointer",
|
"build_sqlite_checkpointer",
|
||||||
"clarify_node",
|
"clarify_node",
|
||||||
|
|
@ -428,22 +448,61 @@ def build_graph(
|
||||||
return builder.compile(checkpointer=checkpointer)
|
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(
|
def build_sqlite_checkpointer(
|
||||||
db_path: Path | str,
|
db_path: Path | str,
|
||||||
) -> AbstractContextManager[BaseCheckpointSaver]:
|
) -> AbstractContextManager[BaseCheckpointSaver]:
|
||||||
"""Construct the production SQLite checkpointer over ``db_path`` (D9, §3.3).
|
"""Construct the production SQLite checkpointer over ``db_path`` (D9, §3.3).
|
||||||
|
|
||||||
Returns a **context manager**, not an entered saver: in
|
Returns a **context manager**, not an entered saver: the caller MUST enter
|
||||||
``langgraph-checkpoint-sqlite`` ``SqliteSaver.from_conn_string`` is a
|
it (``with`` it, or ``__enter__`` and retain it for the graph's lifetime)
|
||||||
``@contextmanager`` classmethod, so the caller MUST enter it (``with`` it,
|
before passing the yielded saver to :func:`build_graph`. The live
|
||||||
or ``__enter__`` and retain it for the graph's lifetime) before passing the
|
coordinator owns that lifecycle (it enters the CM at setup and holds it for
|
||||||
yielded saver to :func:`build_graph`. The live coordinator owns that
|
the daemon's life); passing the raw return value straight into
|
||||||
lifecycle (it enters the CM at setup and holds it for the daemon's life);
|
``build_graph`` would compile a graph whose checkpointer is an un-entered CM
|
||||||
passing the raw return value straight into ``build_graph`` would compile a
|
and break ``get_state`` / ``invoke`` at runtime.
|
||||||
graph whose checkpointer is an un-entered CM and break ``get_state`` /
|
|
||||||
``invoke`` at runtime. The earlier ``-> BaseCheckpointSaver`` annotation
|
We open the connection and construct ``SqliteSaver(conn, serde=...)``
|
||||||
mis-stated this contract (review FIX); the type now matches reality so a
|
ourselves rather than using ``SqliteSaver.from_conn_string`` because the
|
||||||
direct caller cannot be silently misled.
|
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
|
The import of ``langgraph.checkpoint.sqlite`` is deferred to call time so
|
||||||
this module imports cleanly in environments where that optional package is
|
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 = Path(db_path)
|
||||||
db_path.parent.mkdir(parents=True, exist_ok=True)
|
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). -----------------------------------------
|
# --- Driver seam (thread_id-keyed). -----------------------------------------
|
||||||
|
|
|
||||||
|
|
@ -29,6 +29,7 @@ from agent_team.graph import (
|
||||||
INTAKE,
|
INTAKE,
|
||||||
P1_PHASE_SEQUENCE,
|
P1_PHASE_SEQUENCE,
|
||||||
PLAN,
|
PLAN,
|
||||||
|
build_checkpoint_serde,
|
||||||
build_graph,
|
build_graph,
|
||||||
build_sqlite_checkpointer,
|
build_sqlite_checkpointer,
|
||||||
clarify_node,
|
clarify_node,
|
||||||
|
|
@ -293,6 +294,107 @@ def test_build_sqlite_checkpointer_builds_when_dep_present(tmp_path) -> None:
|
||||||
assert hasattr(saver, "get_next_version")
|
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. -------------------------------------------------
|
# --- P2 review-loop wiring. -------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
|
|
||||||
Reference in a new issue