add graph.py
This commit is contained in:
@@ -0,0 +1,105 @@
|
|||||||
|
"""
|
||||||
|
Graph definition for the Human‑in‑the‑loop example.
|
||||||
|
|
||||||
|
The graph is intentionally tiny – two nodes:
|
||||||
|
|
||||||
|
* ``ask_user`` – triggers an interrupt that pauses execution.
|
||||||
|
* ``process_answer`` – receives the resume payload and stores the user answer in state.
|
||||||
|
|
||||||
|
Both functions are async to match LangGraph’s expectations. The module exposes a single helper
|
||||||
|
`build_graph()` which returns a compiled :class:`langgraph.graph.StateGraph` instance.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from typing import Dict, Any
|
||||||
|
|
||||||
|
import questionary
|
||||||
|
from langgraph.constants import START
|
||||||
|
from langgraph.types import interrupt, Command
|
||||||
|
from langgraph.checkpoint.memory import InMemorySaver
|
||||||
|
from langgraph.graph import StateGraph
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# State definition – used by the main module.
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
class GraphState(dict): # Simple dict‑based state for brevity.
|
||||||
|
"""Minimal graph state.
|
||||||
|
|
||||||
|
Attributes:
|
||||||
|
human_value (str | None) – value supplied by the user during interrupt.
|
||||||
|
foo (int) – placeholder for any initial data.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self, *, human_value: str | None = None, foo: int = 0):
|
||||||
|
super().__init__(human_value=human_value, foo=foo)
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# Node that triggers an interrupt.
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
async def ask_user(state: GraphState) -> Command:
|
||||||
|
"""Pause the graph and ask the user for confirmation.
|
||||||
|
|
||||||
|
The node returns a :class:`Command` with ``resume=interrupt(payload)``. Execution stops until
|
||||||
|
the caller sends a resume command.
|
||||||
|
"""
|
||||||
|
payload: Dict[str, Any] = {
|
||||||
|
"type": "confirm",
|
||||||
|
"question": "Do you want to continue?",
|
||||||
|
"options": ["yes", "no"],
|
||||||
|
}
|
||||||
|
return Command(resume=interrupt(payload))
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# Node that receives the user's answer.
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
async def process_answer(state: GraphState) -> GraphState:
|
||||||
|
"""Handle the resume payload and store the user answer in state."""
|
||||||
|
interrupt_obj = state.get("__interrupt__", [])[0]
|
||||||
|
payload: Dict[str, Any] = interrupt_obj.value # type: ignore[assignment]
|
||||||
|
|
||||||
|
answer = payload.get("answer")
|
||||||
|
if not isinstance(answer, str):
|
||||||
|
raise ValueError("Interrupt payload missing 'answer' field")
|
||||||
|
|
||||||
|
state["human_value"] = answer
|
||||||
|
return state
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# Build and compile the graph.
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
def build_graph() -> StateGraph:
|
||||||
|
"""Return a compiled :class:`StateGraph` with interrupt support."""
|
||||||
|
graph = StateGraph(GraphState)
|
||||||
|
graph.add_node("ask", ask_user)
|
||||||
|
graph.add_node("answer", process_answer)
|
||||||
|
graph.set_entry_point("ask")
|
||||||
|
graph.add_edge(START, "ask")
|
||||||
|
graph.add_edge("ask", "answer")
|
||||||
|
|
||||||
|
return graph.compile(checkpointer=InMemorySaver())
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# If run as a script, demonstrate the graph.
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
if __name__ == "__main__": # pragma: no cover
|
||||||
|
import asyncio
|
||||||
|
|
||||||
|
async def demo():
|
||||||
|
g = build_graph()
|
||||||
|
config = {"configurable": {"thread_id": "demo-thread"}}
|
||||||
|
init_state = GraphState(human_value=None, foo=42)
|
||||||
|
async for chunk in g.stream(init_state, config):
|
||||||
|
if "__interrupt__" in chunk:
|
||||||
|
interrupt_obj = chunk["__interrupt__"][0]
|
||||||
|
payload: Dict[str, Any] = interrupt_obj.value # type: ignore[assignment]
|
||||||
|
answer = questionary.select("Choose an option:", choices=payload.get("options", [])).ask()
|
||||||
|
if answer is None:
|
||||||
|
raise RuntimeError("User cancelled")
|
||||||
|
payload["answer"] = answer
|
||||||
|
await g.stream(Command(resume=payload), config)
|
||||||
|
else:
|
||||||
|
print(chunk)
|
||||||
|
print("Final state:", g.checkpointer.get_state(config))
|
||||||
|
|
||||||
|
asyncio.run(demo())
|
||||||
Reference in New Issue
Block a user