115 lines
3.3 KiB
Python
115 lines
3.3 KiB
Python
"""
|
|
LangGraph planning agent demo.
|
|
|
|
Run with:
|
|
python main.py "Сравни Python и JavaScript"
|
|
|
|
Requires OPENAI_API_KEY environment variable.
|
|
"""
|
|
import os
|
|
from typing import TypedDict, List, Optional
|
|
|
|
from langchain_openai import ChatOpenAI
|
|
from langgraph.graph import StateGraph, END
|
|
from langgraph.prebuilt import create_chat_agent
|
|
|
|
# ---------- State definition ----------
|
|
class PlanningState(TypedDict):
|
|
task: str
|
|
plan: Optional[List[str]]
|
|
current_step: int
|
|
results: List[str]
|
|
|
|
# ---------- LLM ----------
|
|
llm = ChatOpenAI(model="gpt-4o-mini", temperature=0)
|
|
|
|
# ---------- Planning node ----------
|
|
def planning(state: PlanningState) -> PlanningState:
|
|
prompt = (
|
|
"Разбей задачу на 3–6 конкретных шагов.\n"
|
|
"Ответ в формате JSON массива строк, например:\n"
|
|
"[""\n"
|
|
" \"Шаг 1: ...\",\n"
|
|
" \"Шаг 2: ...\"\n"
|
|
"]\n"
|
|
f"Задача: {state['task']}"
|
|
)
|
|
response = llm.invoke(prompt)
|
|
import json
|
|
try:
|
|
plan = json.loads(response.content)
|
|
if not isinstance(plan, list):
|
|
raise ValueError("Not a list")
|
|
except Exception as e:
|
|
# fallback: split by lines starting with digits
|
|
plan = []
|
|
for line in response.content.splitlines():
|
|
line = line.strip()
|
|
if line and (line[0].isdigit() or line.startswith("Шаг")):
|
|
plan.append(line)
|
|
return {
|
|
**state,
|
|
"plan": plan,
|
|
"current_step": 0,
|
|
"results": [],
|
|
}
|
|
|
|
# ---------- Execution node ----------
|
|
def execution(state: PlanningState) -> PlanningState:
|
|
step = state["plan"][state["current_step"]]
|
|
prompt = f"Выполни шаг:\n{step}\nОтвет в одном абзаце."
|
|
response = llm.invoke(prompt)
|
|
result = response.content.strip()
|
|
new_results = state["results"].copy()
|
|
new_results.append(result)
|
|
return {
|
|
**state,
|
|
"current_step": state["current_step"] + 1,
|
|
"results": new_results,
|
|
}
|
|
|
|
# ---------- Should continue node ----------
|
|
def should_continue(state: PlanningState) -> str:
|
|
if state["current_step"] >= len(state["plan"]):
|
|
return "finish"
|
|
return "execute"
|
|
|
|
# ---------- Graph construction ----------
|
|
builder = StateGraph(PlanningState)
|
|
builder.add_node("planning", planning)
|
|
builder.add_node("execution", execution)
|
|
builder.add_conditional_edges(
|
|
"execution",
|
|
should_continue,
|
|
{
|
|
"finish": END,
|
|
"execute": "execution",
|
|
},
|
|
)
|
|
builder.set_entry_point("planning")
|
|
graph = builder.compile()
|
|
|
|
# ---------- Run ----------
|
|
if __name__ == "__main__":
|
|
import sys
|
|
if len(sys.argv) < 2:
|
|
print("Usage: python main.py '<task>'")
|
|
sys.exit(1)
|
|
task = sys.argv[1]
|
|
initial_state: PlanningState = {
|
|
"task": task,
|
|
"plan": None,
|
|
"current_step": 0,
|
|
"results": [],
|
|
}
|
|
result = graph.invoke(initial_state)
|
|
plan = result["plan"] or []
|
|
print(f"\nЗадача: {task}\n")
|
|
print("План:")
|
|
for i, step in enumerate(plan, 1):
|
|
print(f"{i}. {step}")
|
|
print("\n[Шаги]")
|
|
for i, res in enumerate(result["results"], 1):
|
|
print(f"[{i}] {res}\n")
|
|
print("Итог: ", "\n".join(result["results"]))
|