add main.py
This commit is contained in:
@@ -0,0 +1,207 @@
|
||||
"""Самокорректирующийся LangGraph-агент: execute → verify (LLM) → retry."""
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
import random
|
||||
import re
|
||||
import sys
|
||||
from typing import Literal, TypedDict
|
||||
|
||||
from dotenv import load_dotenv
|
||||
from langchain_core.messages import HumanMessage, SystemMessage
|
||||
from langchain_openai import ChatOpenAI
|
||||
from langgraph.checkpoint.memory import InMemorySaver
|
||||
from langgraph.graph import END, START, StateGraph
|
||||
|
||||
load_dotenv()
|
||||
|
||||
BROJS_INFERENCE_URL = "https://platform.brojs.ru/jrnl-bh/api/inference/v1"
|
||||
DEFAULT_MODEL = "openai/gpt-oss-20b:free"
|
||||
DEFAULT_TASK = "Вычисли 2+2"
|
||||
DEFAULT_MAX_ATTEMPTS = 5
|
||||
|
||||
|
||||
class AgentState(TypedDict, total=False):
|
||||
task: str
|
||||
result: str
|
||||
attempts: int
|
||||
status: str # pending | success | failed | max_attempts
|
||||
error: str | None
|
||||
max_attempts: int
|
||||
|
||||
|
||||
def _api_key() -> str:
|
||||
return (
|
||||
os.getenv("JOURNAL_MCP_PAT")
|
||||
or os.getenv("JOURNAL_TOKEN")
|
||||
or os.getenv("OPENAI_API_KEY")
|
||||
or ""
|
||||
)
|
||||
|
||||
|
||||
def _base_url() -> str:
|
||||
if os.getenv("OPENAI_BASE_URL"):
|
||||
return os.environ["OPENAI_BASE_URL"]
|
||||
if os.getenv("OPENAI_API_KEY") and not os.getenv("JOURNAL_MCP_PAT"):
|
||||
return os.getenv("OPENROUTER_BASE_URL", "https://openrouter.ai/api/v1")
|
||||
return BROJS_INFERENCE_URL
|
||||
|
||||
|
||||
def build_llm() -> ChatOpenAI:
|
||||
return ChatOpenAI(
|
||||
model=os.getenv("OPENAI_MODEL", DEFAULT_MODEL),
|
||||
base_url=_base_url(),
|
||||
api_key=_api_key(),
|
||||
temperature=0.2,
|
||||
)
|
||||
|
||||
|
||||
def unreliable_tool(task: str) -> str:
|
||||
"""Тестовый инструмент: ~30% бросает ValueError для демонстрации retry."""
|
||||
if random.random() < 0.3:
|
||||
raise ValueError("Случайный сбой unreliable_tool")
|
||||
task_l = task.lower()
|
||||
if "2+2" in task_l or "2 + 2" in task_l:
|
||||
return "4"
|
||||
return f"Выполнено: {task}"
|
||||
|
||||
|
||||
def execute_task(state: AgentState) -> dict:
|
||||
task = state.get("task", DEFAULT_TASK)
|
||||
attempts = int(state.get("attempts", 0))
|
||||
try:
|
||||
result = unreliable_tool(task)
|
||||
return {
|
||||
"result": result,
|
||||
"error": None,
|
||||
"status": "pending",
|
||||
"attempts": attempts,
|
||||
}
|
||||
except Exception as exc:
|
||||
return {
|
||||
"result": "",
|
||||
"error": str(exc),
|
||||
"status": "failed",
|
||||
"attempts": attempts,
|
||||
}
|
||||
|
||||
|
||||
def verify_result(state: AgentState, llm: ChatOpenAI) -> dict:
|
||||
task = state.get("task", DEFAULT_TASK)
|
||||
result = state.get("result", "")
|
||||
error = state.get("error")
|
||||
attempts = int(state.get("attempts", 0))
|
||||
|
||||
if error:
|
||||
return {"status": "failed", "attempts": attempts}
|
||||
|
||||
prompt = (
|
||||
"Ты судья качества ответа. Оцени, решена ли задача.\n"
|
||||
f"Задача: {task}\n"
|
||||
f"Результат: {result}\n\n"
|
||||
"Ответь одним словом: success или failed."
|
||||
)
|
||||
raw = llm.invoke([SystemMessage(content=prompt)]).content or ""
|
||||
verdict = _parse_verdict(str(raw))
|
||||
|
||||
if verdict == "success":
|
||||
return {"status": "success", "attempts": attempts}
|
||||
return {"status": "failed", "attempts": attempts}
|
||||
|
||||
|
||||
def _parse_verdict(text: str) -> Literal["success", "failed"]:
|
||||
t = text.lower().strip()
|
||||
if re.search(r"\bsuccess\b", t):
|
||||
return "success"
|
||||
if re.search(r"\bfailed\b", t):
|
||||
return "failed"
|
||||
if "успех" in t or "верно" in t or "правиль" in t:
|
||||
return "success"
|
||||
return "failed"
|
||||
|
||||
|
||||
def handle_error(state: AgentState) -> dict:
|
||||
attempts = int(state.get("attempts", 0)) + 1
|
||||
max_attempts = int(state.get("max_attempts", DEFAULT_MAX_ATTEMPTS))
|
||||
if attempts >= max_attempts:
|
||||
return {"attempts": attempts, "status": "max_attempts"}
|
||||
return {"attempts": attempts, "status": "pending", "error": None}
|
||||
|
||||
|
||||
def route_after_verify(state: AgentState) -> str:
|
||||
status = state.get("status", "pending")
|
||||
attempts = int(state.get("attempts", 0))
|
||||
max_attempts = int(state.get("max_attempts", DEFAULT_MAX_ATTEMPTS))
|
||||
|
||||
if status == "success":
|
||||
return "end"
|
||||
if attempts >= max_attempts:
|
||||
return "end"
|
||||
return "retry"
|
||||
|
||||
|
||||
def build_graph(llm: ChatOpenAI | None = None):
|
||||
llm = llm or build_llm()
|
||||
|
||||
def _verify(state: AgentState) -> dict:
|
||||
return verify_result(state, llm)
|
||||
|
||||
graph = StateGraph(AgentState)
|
||||
graph.add_node("execute_task", execute_task)
|
||||
graph.add_node("verify_result", _verify)
|
||||
graph.add_node("handle_error", handle_error)
|
||||
|
||||
graph.add_edge(START, "execute_task")
|
||||
graph.add_edge("execute_task", "verify_result")
|
||||
graph.add_conditional_edges(
|
||||
"verify_result",
|
||||
route_after_verify,
|
||||
{"end": END, "retry": "handle_error"},
|
||||
)
|
||||
graph.add_edge("handle_error", "execute_task")
|
||||
|
||||
return graph.compile(checkpointer=InMemorySaver())
|
||||
|
||||
|
||||
def run_demo(task: str = DEFAULT_TASK, max_attempts: int = DEFAULT_MAX_ATTEMPTS) -> AgentState:
|
||||
app = build_graph()
|
||||
config = {"configurable": {"thread_id": "self-correcting-demo"}}
|
||||
initial: AgentState = {
|
||||
"task": task,
|
||||
"result": "",
|
||||
"attempts": 0,
|
||||
"status": "pending",
|
||||
"error": None,
|
||||
"max_attempts": max_attempts,
|
||||
}
|
||||
|
||||
print(f"Задача: {task}")
|
||||
final: AgentState = initial
|
||||
for event in app.stream(initial, config=config, stream_mode="updates"):
|
||||
for node, update in event.items():
|
||||
if not isinstance(update, dict):
|
||||
continue
|
||||
final = {**final, **update}
|
||||
attempt = final.get("attempts", 0)
|
||||
if node == "execute_task":
|
||||
if final.get("error"):
|
||||
print(f"Попытка {attempt + 1}: Error → verify: failed")
|
||||
else:
|
||||
print(f"Попытка {attempt + 1}: результат {final.get('result')} → verify: ...")
|
||||
elif node == "verify_result":
|
||||
print(f" verify: {final.get('status')}")
|
||||
elif node == "handle_error":
|
||||
print(f" retry (attempts={final.get('attempts')})")
|
||||
|
||||
print(f"Итог: {final.get('status')} за {final.get('attempts', 0)} попыток")
|
||||
return final
|
||||
|
||||
|
||||
def main() -> int:
|
||||
task = sys.argv[1] if len(sys.argv) > 1 else DEFAULT_TASK
|
||||
run_demo(task)
|
||||
return 0
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
raise SystemExit(main())
|
||||
Reference in New Issue
Block a user