mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-09-13 13:17:52 +02:00
Compare commits
3
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
375dafb6c6 | ||
|
|
c0e1399e22 | ||
|
|
3913144bdd |
@@ -14,12 +14,14 @@ from typing import (
|
|||||||
NamedTuple,
|
NamedTuple,
|
||||||
TypeVar,
|
TypeVar,
|
||||||
final,
|
final,
|
||||||
|
overload,
|
||||||
)
|
)
|
||||||
from warnings import warn
|
from warnings import warn
|
||||||
|
|
||||||
from langchain_core.messages import AnyMessage
|
from langchain_core.messages import AnyMessage
|
||||||
from langchain_core.runnables import Runnable, RunnableConfig
|
from langchain_core.runnables import Runnable, RunnableConfig
|
||||||
from langgraph.checkpoint.base import BaseCheckpointSaver, CheckpointMetadata
|
from langgraph.checkpoint.base import BaseCheckpointSaver, CheckpointMetadata
|
||||||
|
from pydantic import TypeAdapter
|
||||||
from typing_extensions import NotRequired, TypeAliasType, TypedDict, Unpack, deprecated
|
from typing_extensions import NotRequired, TypeAliasType, TypedDict, Unpack, deprecated
|
||||||
from xxhash import xxh3_128_hexdigest
|
from xxhash import xxh3_128_hexdigest
|
||||||
|
|
||||||
@@ -36,6 +38,7 @@ from langgraph.warnings import LangGraphDeprecatedSinceV10, LangGraphDeprecatedS
|
|||||||
# when used in standalone type aliases.
|
# when used in standalone type aliases.
|
||||||
StateT = TypeVar("StateT")
|
StateT = TypeVar("StateT")
|
||||||
OutputT = TypeVar("OutputT")
|
OutputT = TypeVar("OutputT")
|
||||||
|
ResponseT = TypeVar("ResponseT")
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from langgraph.pregel.protocol import PregelProtocol
|
from langgraph.pregel.protocol import PregelProtocol
|
||||||
@@ -596,13 +599,19 @@ class Interrupt:
|
|||||||
id: str
|
id: str
|
||||||
"""The ID of the interrupt. Can be used to resume the interrupt directly."""
|
"""The ID of the interrupt. Can be used to resume the interrupt directly."""
|
||||||
|
|
||||||
|
response_schema: dict[str, Any] | None = None
|
||||||
|
"""JSON Schema for the value expected when resuming this interrupt, if the graph provided one."""
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
value: Any,
|
value: Any,
|
||||||
id: str = _DEFAULT_INTERRUPT_ID,
|
id: str = _DEFAULT_INTERRUPT_ID,
|
||||||
|
*,
|
||||||
|
response_schema: dict[str, Any] | None = None,
|
||||||
**deprecated_kwargs: Unpack[DeprecatedKwargs],
|
**deprecated_kwargs: Unpack[DeprecatedKwargs],
|
||||||
) -> None:
|
) -> None:
|
||||||
self.value = value
|
self.value = value
|
||||||
|
self.response_schema = response_schema
|
||||||
|
|
||||||
if (
|
if (
|
||||||
(ns := deprecated_kwargs.get("ns", MISSING)) is not MISSING
|
(ns := deprecated_kwargs.get("ns", MISSING)) is not MISSING
|
||||||
@@ -614,8 +623,14 @@ class Interrupt:
|
|||||||
self.id = id
|
self.id = id
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def from_ns(cls, value: Any, ns: str) -> Interrupt:
|
def from_ns(
|
||||||
return cls(value=value, id=xxh3_128_hexdigest(ns.encode()))
|
cls, value: Any, ns: str, *, response_schema: dict[str, Any] | None = None
|
||||||
|
) -> Interrupt:
|
||||||
|
return cls(
|
||||||
|
value=value,
|
||||||
|
id=xxh3_128_hexdigest(ns.encode()),
|
||||||
|
response_schema=response_schema,
|
||||||
|
)
|
||||||
|
|
||||||
@property
|
@property
|
||||||
@deprecated("`interrupt_id` is deprecated. Use `id` instead.", category=None)
|
@deprecated("`interrupt_id` is deprecated. Use `id` instead.", category=None)
|
||||||
@@ -848,7 +863,17 @@ class Command(Generic[N], ToolOutputMixin):
|
|||||||
PARENT: ClassVar[Literal["__parent__"]] = "__parent__"
|
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[Any] | None = None
|
||||||
|
) -> Any:
|
||||||
"""Interrupt the graph with a resumable exception from within a node.
|
"""Interrupt the graph with a resumable exception from within a node.
|
||||||
|
|
||||||
The `interrupt` function enables human-in-the-loop workflows by pausing graph
|
The `interrupt` function enables human-in-the-loop workflows by pausing graph
|
||||||
@@ -918,7 +943,7 @@ def interrupt(value: Any) -> Any:
|
|||||||
for chunk in graph.stream({\"foo\": \"abc\"}, config):
|
for chunk in graph.stream({\"foo\": \"abc\"}, config):
|
||||||
print(chunk)
|
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!!!\")
|
command = Command(resume=\"some input from a human!!!\")
|
||||||
|
|
||||||
@@ -931,12 +956,20 @@ def interrupt(value: Any) -> Any:
|
|||||||
|
|
||||||
Args:
|
Args:
|
||||||
value: The value to surface to the client when the graph is interrupted.
|
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:
|
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:
|
Raises:
|
||||||
GraphInterrupt: On the first invocation within the node, halts execution and surfaces the provided value to the client.
|
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 (
|
from langgraph._internal._constants import (
|
||||||
CONFIG_KEY_CHECKPOINT_NS,
|
CONFIG_KEY_CHECKPOINT_NS,
|
||||||
@@ -948,27 +981,36 @@ def interrupt(value: Any) -> Any:
|
|||||||
from langgraph.errors import GraphInterrupt
|
from langgraph.errors import GraphInterrupt
|
||||||
|
|
||||||
conf = get_config()["configurable"]
|
conf = get_config()["configurable"]
|
||||||
|
adapter = (
|
||||||
|
None
|
||||||
|
if response_schema is None or isinstance(response_schema, dict)
|
||||||
|
else TypeAdapter(response_schema)
|
||||||
|
)
|
||||||
# track interrupt index
|
# track interrupt index
|
||||||
scratchpad = conf[CONFIG_KEY_SCRATCHPAD]
|
scratchpad = conf[CONFIG_KEY_SCRATCHPAD]
|
||||||
idx = scratchpad.interrupt_counter()
|
idx = scratchpad.interrupt_counter()
|
||||||
# find previous resume values
|
# find previous resume values
|
||||||
if scratchpad.resume:
|
if scratchpad.resume:
|
||||||
if idx < len(scratchpad.resume):
|
if idx < len(scratchpad.resume):
|
||||||
|
v = scratchpad.resume[idx]
|
||||||
|
validated = adapter.validate_python(v) if adapter else v
|
||||||
conf[CONFIG_KEY_SEND]([(RESUME, scratchpad.resume)])
|
conf[CONFIG_KEY_SEND]([(RESUME, scratchpad.resume)])
|
||||||
return scratchpad.resume[idx]
|
return validated
|
||||||
# find current resume value
|
# find current resume value
|
||||||
v = scratchpad.get_null_resume(True)
|
v = scratchpad.get_null_resume(True)
|
||||||
if v is not None:
|
if v is not None:
|
||||||
assert len(scratchpad.resume) == idx, (scratchpad.resume, idx)
|
assert len(scratchpad.resume) == idx, (scratchpad.resume, idx)
|
||||||
|
validated = adapter.validate_python(v) if adapter else v
|
||||||
scratchpad.resume.append(v)
|
scratchpad.resume.append(v)
|
||||||
conf[CONFIG_KEY_SEND]([(RESUME, scratchpad.resume)])
|
conf[CONFIG_KEY_SEND]([(RESUME, scratchpad.resume)])
|
||||||
return v
|
return validated
|
||||||
# no resume value found
|
# no resume value found
|
||||||
raise GraphInterrupt(
|
raise GraphInterrupt(
|
||||||
(
|
(
|
||||||
Interrupt.from_ns(
|
Interrupt.from_ns(
|
||||||
value=value,
|
value=value,
|
||||||
ns=conf[CONFIG_KEY_CHECKPOINT_NS],
|
ns=conf[CONFIG_KEY_CHECKPOINT_NS],
|
||||||
|
response_schema=adapter.json_schema() if adapter else response_schema,
|
||||||
),
|
),
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -1,9 +1,14 @@
|
|||||||
|
from dataclasses import dataclass
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
from langgraph.checkpoint.base import BaseCheckpointSaver
|
from langgraph.checkpoint.base import BaseCheckpointSaver
|
||||||
|
from pydantic import BaseModel, ValidationError
|
||||||
from typing_extensions import TypedDict
|
from typing_extensions import TypedDict
|
||||||
|
|
||||||
from langgraph.graph import END, START, StateGraph
|
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
|
pytestmark = pytest.mark.anyio
|
||||||
|
|
||||||
@@ -90,3 +95,116 @@ async def test_interruption_without_state_updates_async(
|
|||||||
assert (await graph.aget_state(thread)).next == ()
|
assert (await graph.aget_state(thread)).next == ()
|
||||||
n_checkpoints = len([c async for c in graph.aget_state_history(thread)])
|
n_checkpoints = len([c async for c in graph.aget_state_history(thread)])
|
||||||
assert n_checkpoints == (5 if durability != "exit" else 3)
|
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)
|
||||||
|
}
|
||||||
|
|||||||
@@ -5583,6 +5583,7 @@ def test_falsy_return_from_task(sync_checkpointer: BaseCheckpointSaver):
|
|||||||
"interrupts": [
|
"interrupts": [
|
||||||
{
|
{
|
||||||
"id": AnyStr(),
|
"id": AnyStr(),
|
||||||
|
"response_schema": None,
|
||||||
"value": "test",
|
"value": "test",
|
||||||
},
|
},
|
||||||
],
|
],
|
||||||
@@ -5627,6 +5628,7 @@ def test_falsy_return_from_task(sync_checkpointer: BaseCheckpointSaver):
|
|||||||
"interrupts": (
|
"interrupts": (
|
||||||
{
|
{
|
||||||
"id": AnyStr(),
|
"id": AnyStr(),
|
||||||
|
"response_schema": None,
|
||||||
"value": "test",
|
"value": "test",
|
||||||
},
|
},
|
||||||
),
|
),
|
||||||
|
|||||||
@@ -295,6 +295,8 @@ class Interrupt(TypedDict):
|
|||||||
"""The value associated with the interrupt."""
|
"""The value associated with the interrupt."""
|
||||||
id: str
|
id: str
|
||||||
"""The ID of the interrupt. Can be used to resume the interrupt."""
|
"""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):
|
class Thread(TypedDict):
|
||||||
|
|||||||
Reference in New Issue
Block a user