diff --git a/main.py b/main.py new file mode 100644 index 0000000..73bdce0 --- /dev/null +++ b/main.py @@ -0,0 +1,123 @@ +""" +Self‑correcting LangGraph agent. + +Run with: + python main.py "Вычисли 2+2" + +The script will keep retrying until the LLM judge says `success` or the maximum number of attempts is reached. +""" +import random +from typing import TypedDict, Dict + +# LangGraph imports +from langgraph.graph import StateGraph +from langgraph.checkpoint.memory import InMemorySaver +from langchain_openai import ChatOpenAI +from langchain_core.messages import HumanMessage, AIMessage + +# ---------- 1. State definition ------------------------------------------- +class AgentState(TypedDict): + task: str + result: str + attempts: int + status: str # pending | success | failed | max_attempts + error: str | None + max_attempts: int + +# ---------- 2. Tool ------------------------------------------------------- +class UnreliableTool: + """Simulates a tool that fails with ~30% probability.""" + + def __call__(self, input_: str) -> str: + if random.random() < 0.3: + raise ValueError("Simulated tool failure") + # Very simple evaluation: try to compute arithmetic expression + try: + return str(eval(input_)) + except Exception as e: + raise ValueError(f"Evaluation error: {e}") + +unreliable_tool = UnreliableTool() + +# ---------- 3. Nodes ----------------------------------------------------- +async def execute_task(state: AgentState) -> Dict[str, str]: + """Runs the task using the unreliable tool.""" + try: + result = unreliable_tool(state["task"]) + return {"result": result, "error": None} + except Exception as e: + return {"result": "", "error": str(e)} + +async def verify_result(state: AgentState) -> Dict[str, str]: + """LLM judge that decides success or failed.""" + llm = ChatOpenAI(temperature=0) + # Ask the model to output only 'success' or 'failed' + prompt = ( + f"Task: {state['task']}\n" + f"Result: {state['result']}\n" + "Is this result correct? Respond with either 'success' or 'failed'." + ) + response = await llm.ainvoke(HumanMessage(content=prompt)) + verdict = response.content.strip().lower() + if verdict not in {"success", "failed"}: + # Fallback: treat as failed + verdict = "failed" + return {"status": verdict} + +async def handle_error(state: AgentState) -> Dict[str, str]: + """Increment attempts and prepare for retry.""" + new_attempts = state["attempts"] + 1 + if new_attempts >= state["max_attempts"]: + return {"status": "max_attempts", "attempts": new_attempts} + return {"attempts": new_attempts, "status": "pending"} + +# ---------- 4. Graph ----------------------------------------------------- +builder = StateGraph(AgentState) +builder.add_node("execute_task", execute_task) +builder.add_node("verify_result", verify_result) +builder.add_node("handle_error", handle_error) + +builder.set_entry_point("execute_task") +builder.add_edge("execute_task", "verify_result") +builder.add_conditional_edges( + "verify_result", + lambda x: x["status"], + { + "success": "END", + "failed": "handle_error", + "max_attempts": "END", + }, +) +builder.add_edge("handle_error", "execute_task") + +graph = builder.compile(checkpointer=InMemorySaver()) + +# ---------- 5. CLI ------------------------------------------------------- +if __name__ == "__main__": + import sys + + if len(sys.argv) < 2: + print("Usage: python main.py ''") + sys.exit(1) + + task_input = sys.argv[1] + initial_state: AgentState = { + "task": task_input, + "result": "", + "attempts": 0, + "status": "pending", + "error": None, + "max_attempts": 5, + } + + result = graph.invoke(initial_state) + final_status = result["status"] + attempts = result.get("attempts", 0) + 1 # include last attempt + print(f"Задача: {task_input}") + if final_status == "success": + print(f"Итог: success за {attempts} попытки{'и' if attempts>1 else ''}") + elif final_status == "max_attempts": + print(f"Не удалось достичь успеха после {attempts} попыток.") + else: + print("Непредвиденный статус", final_status) +"""