Files

132 lines
5.6 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
# main.py
"""
Пример реализации Humanintheloop (HITL) в LangGraph.
Граф состоит из одного узла, который вызывает кастомное прерывание,
показывает пользователю вопрос и варианты ответа через questionary,
получает ответ и возобновляет выполнение графа.
Требования:
- Python 3.10+
- langgraph
- questionary
"""
from typing import TypedDict, List, Dict, Any
import questionary
from langgraph.graph import StateGraph, START
from langgraph.constants import interrupt
from langgraph.types import Command
from langgraph.checkpoint.memory import InMemorySaver
# ----------------------------------------------------------------------
# 1. Состояние графа
# ----------------------------------------------------------------------
class GraphState(TypedDict):
"""
Хранит данные, которые будут заполнены пользователем.
- human_value: ответ пользователя на вопрос из прерывания.
- foo: произвольное начальное значение (необязательно).
"""
human_value: str | None
foo: int
# ----------------------------------------------------------------------
# 2. Узел с кастомным прерыванием
# ----------------------------------------------------------------------
def interrupt_node(state: GraphState) -> GraphState:
"""
Вызывает `interrupt` и возвращает обновлённое состояние после возобновления.
"""
# Создаём словарь‑payload для прерывания
payload = {
"type": "confirm",
"question": "Уверены, что хотите продолжить?",
"allow_responds": ["approve", "reject"],
}
# Вызываем прерывание. Пока граф не возобновят,
# выполнение из этого узла не вернётся.
interrupt(payload)
# После возобновления `state` будет обновлён функцией
# обработчика в цикле ниже (см. `resume_node`).
return state
# ----------------------------------------------------------------------
# 3. Узел, который сохраняет ответ пользователя
# ----------------------------------------------------------------------
def resume_node(state: GraphState) -> GraphState:
"""
Этот узел вызывается после возобновления графа.
Он получает в состоянии поле `human_value`, которое было добавлено
в payload прерывания пользователем.
"""
# В state уже должно быть ключ 'human_value', если всё прошло правильно
return state
# ----------------------------------------------------------------------
# 4. Сборка графа
# ----------------------------------------------------------------------
builder = StateGraph(GraphState)
# Добавляем узлы и переходы
builder.add_node("interrupt", interrupt_node)
builder.add_node("resume", resume_node)
builder.set_entry_point("interrupt")
builder.add_edge(START, "interrupt")
builder.add_edge("interrupt", "resume")
# Создаём чекпоинтер (InMemorySaver) для возобновления после прерывания
checkpoint = InMemorySaver()
graph = builder.compile(checkpointer=checkpoint)
# ----------------------------------------------------------------------
# 5. Цикл запуска с обработкой прерываний
# ----------------------------------------------------------------------
def run_graph():
"""
Запускает граф, обрабатывает кастомные прерывания и возобновляет выполнение.
"""
# Инициализируем состояние (можно задать начальные данные)
init_state: GraphState = {"human_value": None, "foo": 42}
# Уникальный идентификатор потока
thread_id = "demo_thread"
# Конфигурация для stream()
config = {"configurable": {"thread_id": thread_id}}
# Запускаем поток генерации чанков
for chunk in graph.stream(init_state, config):
# Если в чанк пришло прерывание – обрабатываем его
if "__interrupt__" in chunk:
interrupt_payload = chunk["__interrupt__"][0].value # dict
# Показываем пользователю вопрос и варианты ответа
answer = questionary.select(
interrupt_payload["question"],
choices=interrupt_payload["allow_responds"]
).ask()
# Добавляем ответ в payload
interrupt_payload["answer"] = answer
# Возобновляем граф, передавая обновлённый payload
resume_cmd = Command(resume=interrupt_payload)
for _ in graph.stream(resume_cmd, config):
pass # продолжаем до конца
else:
# Выводим обычные чанки (можно логировать или обрабатывать иначе)
print(chunk)
if __name__ == "__main__":
run_graph()