diff --git a/libs/langgraph/langgraph/types.py b/libs/langgraph/langgraph/types.py index 2fcf5bef9..5e2a1a897 100644 --- a/libs/langgraph/langgraph/types.py +++ b/libs/langgraph/langgraph/types.py @@ -12,15 +12,23 @@ from typing import ( Generic, Literal, NamedTuple, - TypeVar, final, + overload, ) from warnings import warn from langchain_core.messages import AnyMessage from langchain_core.runnables import Runnable, RunnableConfig from langgraph.checkpoint.base import BaseCheckpointSaver, CheckpointMetadata -from typing_extensions import NotRequired, TypeAliasType, TypedDict, Unpack, deprecated +from pydantic import TypeAdapter +from typing_extensions import ( + NotRequired, + TypeAliasType, + TypedDict, + TypeVar, + Unpack, + deprecated, +) from xxhash import xxh3_128_hexdigest from langgraph._internal._cache import default_cache_key @@ -36,6 +44,7 @@ from langgraph.warnings import LangGraphDeprecatedSinceV10, LangGraphDeprecatedS # when used in standalone type aliases. StateT = TypeVar("StateT") OutputT = TypeVar("OutputT") +ResponseT = TypeVar("ResponseT", default=Any) if TYPE_CHECKING: from langgraph.pregel.protocol import PregelProtocol @@ -572,7 +581,7 @@ _DEFAULT_INTERRUPT_ID = "placeholder-id" @final @dataclass(init=False, slots=True) -class Interrupt: +class Interrupt(Generic[ResponseT]): """Information about an interrupt that occurred in a node. !!! version-added "Added in version 0.2.24" @@ -596,13 +605,22 @@ class Interrupt: id: str """The ID of the interrupt. Can be used to resume the interrupt directly.""" + response_schema: type[ResponseT] | dict[str, Any] | None = None + """Schema for the value expected when resuming this interrupt, if the graph provided one. + + A surfaced interrupt carries JSON Schema (a `dict`); `type[ResponseT]` records the + Python type at construction so `Interrupt[Decision]` is meaningful to type checkers.""" + def __init__( self, value: Any, id: str = _DEFAULT_INTERRUPT_ID, + *, + response_schema: type[ResponseT] | dict[str, Any] | None = None, **deprecated_kwargs: Unpack[DeprecatedKwargs], ) -> None: self.value = value + self.response_schema = response_schema if ( (ns := deprecated_kwargs.get("ns", MISSING)) is not MISSING @@ -614,8 +632,18 @@ class Interrupt: self.id = id @classmethod - def from_ns(cls, value: Any, ns: str) -> Interrupt: - return cls(value=value, id=xxh3_128_hexdigest(ns.encode())) + def from_ns( + cls, + value: Any, + ns: str, + *, + response_schema: type[ResponseT] | dict[str, Any] | None = None, + ) -> Interrupt[ResponseT]: + return cls( + value=value, + id=xxh3_128_hexdigest(ns.encode()), + response_schema=response_schema, + ) @property @deprecated("`interrupt_id` is deprecated. Use `id` instead.", category=None) @@ -848,7 +876,17 @@ class Command(Generic[N], ToolOutputMixin): PARENT: ClassVar[Literal["__parent__"]] = "__parent__" -def interrupt(value: Any) -> Any: +@overload +def interrupt(value: Any, *, response_schema: type[ResponseT]) -> ResponseT: ... + + +@overload +def interrupt(value: Any, *, response_schema: dict[str, Any] | None = None) -> Any: ... + + +def interrupt( + value: Any, *, response_schema: dict[str, Any] | type | None = None +) -> Any: """Interrupt the graph with a resumable exception from within a node. The `interrupt` function enables human-in-the-loop workflows by pausing graph @@ -918,7 +956,7 @@ def interrupt(value: Any) -> Any: for chunk in graph.stream({\"foo\": \"abc\"}, config): print(chunk) - # > {'__interrupt__': (Interrupt(value='what is your age?', id='45fda8478b2ef754419799e10992af06'),)} + # > {'__interrupt__': (Interrupt(value='what is your age?', id='45fda8478b2ef754419799e10992af06', response_schema=None),)} command = Command(resume=\"some input from a human!!!\") @@ -931,12 +969,20 @@ def interrupt(value: Any) -> Any: Args: value: The value to surface to the client when the graph is interrupted. + response_schema: Optional schema for the value expected on resume, surfaced + to clients so they can render a typed input form. Accepts a JSON Schema + `dict` (used as-is, resume values are not validated), or a Pydantic model + class, `TypedDict`, or dataclass, which are converted to JSON Schema for + clients and used to validate the resume value; the validated object is + what `interrupt` returns. Returns: - Any: On subsequent invocations within the same node (same task to be precise), returns the value provided during the first invocation + Any: On subsequent invocations within the same node (same task to be precise), returns the value provided during the first invocation, + validated against `response_schema` when one that supports validation was given. Raises: GraphInterrupt: On the first invocation within the node, halts execution and surfaces the provided value to the client. + pydantic.ValidationError: When a resume value does not match a Pydantic model, `TypedDict`, or dataclass `response_schema`. """ from langgraph._internal._constants import ( CONFIG_KEY_CHECKPOINT_NS, @@ -948,27 +994,36 @@ def interrupt(value: Any) -> Any: from langgraph.errors import GraphInterrupt conf = get_config()["configurable"] + adapter = ( + None + if response_schema is None or isinstance(response_schema, dict) + else TypeAdapter(response_schema) + ) # track interrupt index scratchpad = conf[CONFIG_KEY_SCRATCHPAD] idx = scratchpad.interrupt_counter() # find previous resume values if scratchpad.resume: if idx < len(scratchpad.resume): - conf[CONFIG_KEY_SEND]([(RESUME, scratchpad.resume)]) - return scratchpad.resume[idx] + v = scratchpad.resume[idx] + validated = adapter.validate_python(v) if adapter else v + conf[CONFIG_KEY_SEND]([(RESUME, scratchpad.resume[: idx + 1])]) + return validated # find current resume value v = scratchpad.get_null_resume(True) if v is not None: assert len(scratchpad.resume) == idx, (scratchpad.resume, idx) + validated = adapter.validate_python(v) if adapter else v scratchpad.resume.append(v) conf[CONFIG_KEY_SEND]([(RESUME, scratchpad.resume)]) - return v + return validated # no resume value found raise GraphInterrupt( ( Interrupt.from_ns( value=value, ns=conf[CONFIG_KEY_CHECKPOINT_NS], + response_schema=adapter.json_schema() if adapter else response_schema, ), ) ) diff --git a/libs/langgraph/tests/test_interruption.py b/libs/langgraph/tests/test_interruption.py index a484d74e5..f4d401e94 100644 --- a/libs/langgraph/tests/test_interruption.py +++ b/libs/langgraph/tests/test_interruption.py @@ -1,9 +1,14 @@ +from dataclasses import dataclass +from typing import Any + import pytest from langgraph.checkpoint.base import BaseCheckpointSaver +from pydantic import BaseModel, ValidationError from typing_extensions import TypedDict from langgraph.graph import END, START, StateGraph -from langgraph.types import Durability +from langgraph.types import Command, Durability, Interrupt, interrupt +from tests.any_str import AnyStr pytestmark = pytest.mark.anyio @@ -90,3 +95,150 @@ async def test_interruption_without_state_updates_async( assert (await graph.aget_state(thread)).next == () n_checkpoints = len([c async for c in graph.aget_state_history(thread)]) assert n_checkpoints == (5 if durability != "exit" else 3) + + +class Decision(BaseModel): + approved: bool + note: str | None = None + + +class DecisionDict(TypedDict): + approved: bool + + +@dataclass +class DecisionData: + approved: bool + + +RAW_SCHEMA = {"type": "object", "properties": {"approved": {"type": "boolean"}}} + + +@pytest.mark.parametrize( + ("response_schema", "expected_schema", "expected_answer"), + [ + (None, None, {"approved": True, "extra": 1}), + (RAW_SCHEMA, RAW_SCHEMA, {"approved": True, "extra": 1}), + (Decision, Decision.model_json_schema(), Decision(approved=True)), + ( + DecisionDict, + { + "properties": {"approved": {"title": "Approved", "type": "boolean"}}, + "required": ["approved"], + "title": "DecisionDict", + "type": "object", + }, + {"approved": True}, + ), + ( + DecisionData, + { + "properties": {"approved": {"title": "Approved", "type": "boolean"}}, + "required": ["approved"], + "title": "DecisionData", + "type": "object", + }, + DecisionData(approved=True), + ), + ], + ids=["none", "raw_dict", "pydantic", "typeddict", "dataclass"], +) +def test_interrupt_response_schema( + sync_checkpointer: BaseCheckpointSaver, + response_schema: Any, + expected_schema: dict[str, Any] | None, + expected_answer: Any, +) -> None: + class State(TypedDict): + answer: Any + + def node(state: State) -> State: + return { + "answer": interrupt( + {"question": "approve?"}, response_schema=response_schema + ) + } + + graph = ( + StateGraph(State) + .add_node("node", node) + .add_edge(START, "node") + .compile(checkpointer=sync_checkpointer) + ) + config = {"configurable": {"thread_id": "1"}} + expected = Interrupt( + value={"question": "approve?"}, id=AnyStr(), response_schema=expected_schema + ) + + assert list(graph.stream({"answer": None}, config)) == [ + {"__interrupt__": (expected,)} + ] + assert graph.get_state(config).tasks[0].interrupts == (expected,) + assert graph.invoke(Command(resume={"approved": True, "extra": 1}), config) == { + "answer": expected_answer + } + + +@pytest.mark.parametrize("resume_style", ["null", "map"]) +def test_interrupt_response_schema_rejects_invalid_resume( + sync_checkpointer: BaseCheckpointSaver, resume_style: str +) -> None: + class State(TypedDict): + answer: Any + + def node(state: State) -> State: + return {"answer": interrupt("approve?", response_schema=Decision)} + + graph = ( + StateGraph(State) + .add_node("node", node) + .add_edge(START, "node") + .compile(checkpointer=sync_checkpointer) + ) + config = {"configurable": {"thread_id": "1"}} + graph.invoke({"answer": None}, config) + [pending] = graph.get_state(config).tasks[0].interrupts + + def resume(value: dict[str, Any]) -> Command: + return Command(resume=value if resume_style == "null" else {pending.id: value}) + + with pytest.raises(ValidationError, match="approved"): + graph.invoke(resume({"approved": "nope"}), config) + + assert graph.invoke(resume({"approved": False}), config) == { + "answer": Decision(approved=False) + } + + +@pytest.mark.parametrize("resume_style", ["null", "id_map"]) +def test_interrupt_response_schema_invalid_resume_after_earlier_interrupt( + sync_checkpointer: BaseCheckpointSaver, resume_style: str +) -> None: + class State(TypedDict): + answer: Any + + def node(state: State) -> State: + first = interrupt("first") + second = interrupt("approve?", response_schema=Decision) + return {"answer": [first, second]} + + graph = ( + StateGraph(State) + .add_node("node", node) + .add_edge(START, "node") + .compile(checkpointer=sync_checkpointer) + ) + config = {"configurable": {"thread_id": "1"}} + graph.invoke({"answer": None}, config) + graph.invoke(Command(resume="ok"), config) + [pending] = graph.get_state(config).tasks[0].interrupts + + def resume(value: dict[str, Any]) -> Command: + return Command(resume=value if resume_style == "null" else {pending.id: value}) + + with pytest.raises(ValidationError, match="approved"): + graph.invoke(resume({"approved": "nope"}), config) + + assert graph.invoke(resume({"approved": True}), config) == { + "answer": ["ok", Decision(approved=True)] + } diff --git a/libs/langgraph/tests/test_pregel.py b/libs/langgraph/tests/test_pregel.py index c166c5837..2e6661e04 100644 --- a/libs/langgraph/tests/test_pregel.py +++ b/libs/langgraph/tests/test_pregel.py @@ -5583,6 +5583,7 @@ def test_falsy_return_from_task(sync_checkpointer: BaseCheckpointSaver): "interrupts": [ { "id": AnyStr(), + "response_schema": None, "value": "test", }, ], @@ -5627,6 +5628,7 @@ def test_falsy_return_from_task(sync_checkpointer: BaseCheckpointSaver): "interrupts": ( { "id": AnyStr(), + "response_schema": None, "value": "test", }, ), diff --git a/libs/sdk-py/langgraph_sdk/schema.py b/libs/sdk-py/langgraph_sdk/schema.py index 18b1b44f3..a16733e60 100644 --- a/libs/sdk-py/langgraph_sdk/schema.py +++ b/libs/sdk-py/langgraph_sdk/schema.py @@ -295,6 +295,8 @@ class Interrupt(TypedDict): """The value associated with the interrupt.""" id: str """The ID of the interrupt. Can be used to resume the interrupt.""" + response_schema: NotRequired[dict[str, Any]] + """JSON Schema for the value expected when resuming this interrupt, if the graph provided one.""" class Thread(TypedDict):