Implement LangGraph reflection agent with structured critic, rewrite loop, max_rounds cap, and CLI demo logging.: update main.py
This commit is contained in:
@@ -1,117 +1,174 @@
|
|||||||
"""
|
import json
|
||||||
LangGraph Reflection Agent
|
import os
|
||||||
|
import re
|
||||||
Usage:
|
|
||||||
python main.py "Your question"
|
|
||||||
"""
|
|
||||||
import sys
|
import sys
|
||||||
from typing import TypedDict, Dict
|
from typing import Literal, TypedDict
|
||||||
|
|
||||||
|
from langchain_core.output_parsers import PydanticOutputParser
|
||||||
|
from langchain_openai import ChatOpenAI
|
||||||
|
from langgraph.graph import END, START, StateGraph
|
||||||
|
from pydantic import BaseModel, Field
|
||||||
|
|
||||||
|
|
||||||
|
DEFAULT_QUESTION = "Объясни студенту разницу между tool и resource в MCP"
|
||||||
|
|
||||||
|
|
||||||
# State definition
|
|
||||||
class ReflectState(TypedDict):
|
class ReflectState(TypedDict):
|
||||||
question: str
|
question: str
|
||||||
draft: str
|
draft: str
|
||||||
critique: str
|
critique: str
|
||||||
verdict: str # 'ok' or 'needs_revision'
|
verdict: str
|
||||||
round: int
|
round: int
|
||||||
max_rounds: int
|
max_rounds: int
|
||||||
|
|
||||||
# Node functions
|
|
||||||
async def draft_answer(state: Dict) -> Dict:
|
|
||||||
from langchain_openai import ChatOpenAI
|
|
||||||
llm = ChatOpenAI(model="gpt-3.5-turbo")
|
|
||||||
prompt = f"Write a concise answer (5–10 sentences) to the following question:\n\n{state['question']}"
|
|
||||||
response = await llm.invoke(prompt)
|
|
||||||
state["draft"] = response.content
|
|
||||||
return state
|
|
||||||
|
|
||||||
async def reflect(state: Dict) -> Dict:
|
class ReflectionResult(BaseModel):
|
||||||
from langchain_openai import ChatOpenAI
|
verdict: Literal["ok", "needs_revision"] = Field(
|
||||||
llm = ChatOpenAI(model="gpt-3.5-turbo")
|
description="Whether the draft is acceptable or needs one more revision."
|
||||||
prompt = (
|
|
||||||
f"You are a critic evaluating the draft answer for completeness, specificity, and lack of filler.\n"
|
|
||||||
f"Draft: {state['draft']}\n"
|
|
||||||
"Provide verdict (ok or needs_revision) and 2–3 bullet points of critique."
|
|
||||||
)
|
)
|
||||||
response = await llm.invoke(prompt)
|
critique: list[str] = Field(
|
||||||
# Simple parsing
|
description="Two or three short critique bullets in Russian."
|
||||||
text = response.content.strip()
|
|
||||||
if "needs_revision" in text.lower():
|
|
||||||
state["verdict"] = "needs_revision"
|
|
||||||
else:
|
|
||||||
state["verdict"] = "ok"
|
|
||||||
state["critique"] = text
|
|
||||||
return state
|
|
||||||
|
|
||||||
async def rewrite(state: Dict) -> Dict:
|
|
||||||
from langchain_openai import ChatOpenAI
|
|
||||||
llm = ChatOpenAI(model="gpt-3.5-turbo")
|
|
||||||
prompt = (
|
|
||||||
f"Rewrite the draft answer incorporating the following critique:\n"
|
|
||||||
f"Critique: {state['critique']}\n"
|
|
||||||
"Provide a revised concise answer (5–10 sentences)."
|
|
||||||
)
|
)
|
||||||
response = await llm.invoke(prompt)
|
|
||||||
state["draft"] = response.content
|
|
||||||
state["round"] += 1
|
|
||||||
return state
|
|
||||||
|
|
||||||
# Build graph
|
|
||||||
from langgraph.graph import StateGraph
|
|
||||||
|
|
||||||
graph_builder = StateGraph(ReflectState)
|
def build_llm() -> ChatOpenAI:
|
||||||
|
model = os.getenv("OPENAI_MODEL", "openai/gpt-oss-20b")
|
||||||
|
base_url = os.getenv("OPENAI_BASE_URL")
|
||||||
|
api_key = os.getenv("OPENAI_API_KEY", "dummy")
|
||||||
|
return ChatOpenAI(model=model, base_url=base_url, api_key=api_key, temperature=0)
|
||||||
|
|
||||||
graph_builder.add_node("draft_answer", draft_answer)
|
|
||||||
|
|
||||||
graph_builder.add_node("reflect", reflect)
|
def parse_reflection_result(raw_text: str) -> ReflectionResult:
|
||||||
|
match = re.search(r"\{.*\}", raw_text, re.DOTALL)
|
||||||
|
if not match:
|
||||||
|
raise ValueError(f"Could not find JSON object in critic output: {raw_text}")
|
||||||
|
data = json.loads(match.group(0))
|
||||||
|
return ReflectionResult.model_validate(data)
|
||||||
|
|
||||||
graph_builder.add_node("rewrite", rewrite)
|
|
||||||
|
|
||||||
# Connections
|
def draft_answer(state: ReflectState) -> dict:
|
||||||
start_edge = "draft_answer"
|
llm = build_llm()
|
||||||
end_edge = None # will be set in condition
|
prompt = (
|
||||||
|
"Ты пишешь краткий, содержательный учебный ответ на вопрос студента. "
|
||||||
|
"MCP здесь означает Model Context Protocol, а не Minecraft и не другие расшифровки. "
|
||||||
|
"Дай ответ ровно в формате обычного текста на 5-10 предложений, без таблиц, списков и markdown-оформления. "
|
||||||
|
"Обязательно объясни разницу между tool и resource именно в контексте Model Context Protocol. "
|
||||||
|
"Подсказка по смыслу: tool в MCP — это вызываемое действие/операция, у которой модель может передать аргументы и получить результат; "
|
||||||
|
"resource в MCP — это данные или контекст, доступные для чтения, часто адресуемые по URI, которые помогают модели, но сами ничего не выполняют. "
|
||||||
|
"Обязательно сравни их по назначению и приведи простой пример.\n\n"
|
||||||
|
f"Вопрос: {state['question']}"
|
||||||
|
)
|
||||||
|
response = llm.invoke(prompt)
|
||||||
|
return {"draft": response.content.strip()}
|
||||||
|
|
||||||
def should_end(state: Dict) -> str:
|
|
||||||
|
def reflect(state: ReflectState) -> dict:
|
||||||
|
llm = build_llm()
|
||||||
|
parser = PydanticOutputParser(pydantic_object=ReflectionResult)
|
||||||
|
prompt = (
|
||||||
|
"Ты отдельный узел-критик. Оцени черновик по трем критериям: полнота, "
|
||||||
|
"конкретика, отсутствие воды. Дополнительно проверь, что MCP интерпретирован "
|
||||||
|
"как Model Context Protocol и что ответ дан обычным текстом на 5-10 предложений, "
|
||||||
|
"без таблиц и списков. Также проверь смысловую точность: tool должен быть описан как вызываемое действие, "
|
||||||
|
"а resource — как читаемые данные или контекст. Если текст хорош, поставь verdict=ok. "
|
||||||
|
"Если есть проблемы, поставь verdict=needs_revision и дай 2-3 коротких "
|
||||||
|
"замечания.\n"
|
||||||
|
"Верни ответ строго в формате JSON по инструкции ниже.\n\n"
|
||||||
|
f"{parser.get_format_instructions()}\n\n"
|
||||||
|
f"Вопрос: {state['question']}\n\n"
|
||||||
|
f"Черновик:\n{state['draft']}"
|
||||||
|
)
|
||||||
|
response = llm.invoke(prompt)
|
||||||
|
result = parse_reflection_result(response.content)
|
||||||
|
critique_text = "\n".join(f"- {item}" for item in result.critique)
|
||||||
|
return {"verdict": result.verdict, "critique": critique_text}
|
||||||
|
|
||||||
|
|
||||||
|
def rewrite(state: ReflectState) -> dict:
|
||||||
|
llm = build_llm()
|
||||||
|
prompt = (
|
||||||
|
"Перепиши ответ с учетом замечаний критика. Сохрани формат краткого ответа "
|
||||||
|
"на 5-10 предложений обычным текстом без таблиц и списков, исправь недочеты "
|
||||||
|
"и не добавляй лишнюю воду. MCP здесь означает Model Context Protocol.\n\n"
|
||||||
|
f"Вопрос: {state['question']}\n\n"
|
||||||
|
f"Текущий черновик:\n{state['draft']}\n\n"
|
||||||
|
f"Замечания критика:\n{state['critique']}"
|
||||||
|
)
|
||||||
|
response = llm.invoke(prompt)
|
||||||
|
return {
|
||||||
|
"draft": response.content.strip(),
|
||||||
|
"round": state["round"] + 1,
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def next_step(state: ReflectState) -> str:
|
||||||
if state["verdict"] == "ok":
|
if state["verdict"] == "ok":
|
||||||
return "END"
|
return "finish"
|
||||||
if state["round"] >= state.get("max_rounds", 2):
|
if state["round"] >= state["max_rounds"]:
|
||||||
return "END"
|
return "finish"
|
||||||
return "rewrite"
|
return "rewrite"
|
||||||
|
|
||||||
# Add edges with condition
|
|
||||||
from langgraph.graph import END
|
|
||||||
|
|
||||||
graph_builder.set_entry_point(start_edge)
|
def build_graph():
|
||||||
|
builder = StateGraph(ReflectState)
|
||||||
graph_builder.add_conditional_edges(
|
builder.add_node("draft_answer", draft_answer)
|
||||||
|
builder.add_node("reflect", reflect)
|
||||||
|
builder.add_node("rewrite", rewrite)
|
||||||
|
builder.add_edge(START, "draft_answer")
|
||||||
|
builder.add_edge("draft_answer", "reflect")
|
||||||
|
builder.add_conditional_edges(
|
||||||
"reflect",
|
"reflect",
|
||||||
should_end,
|
next_step,
|
||||||
{
|
{
|
||||||
"rewrite": "rewrite",
|
"rewrite": "rewrite",
|
||||||
"END": END,
|
"finish": END,
|
||||||
},
|
},
|
||||||
)
|
)
|
||||||
# rewrite -> reflect
|
builder.add_edge("rewrite", "reflect")
|
||||||
|
return builder.compile()
|
||||||
|
|
||||||
graph_builder.add_edge("rewrite", "reflect")
|
|
||||||
|
|
||||||
graph = graph_builder.compile()
|
def run_with_log(question: str, max_rounds: int = 2) -> ReflectState:
|
||||||
|
graph = build_graph()
|
||||||
if __name__ == "__main__":
|
state: ReflectState = {
|
||||||
if len(sys.argv) < 2:
|
|
||||||
print("Usage: python main.py \"Your question\"")
|
|
||||||
sys.exit(1)
|
|
||||||
question = sys.argv[1]
|
|
||||||
initial_state: ReflectState = {
|
|
||||||
"question": question,
|
"question": question,
|
||||||
"draft": "",
|
"draft": "",
|
||||||
"critique": "",
|
"critique": "",
|
||||||
"verdict": "",
|
"verdict": "",
|
||||||
"round": 0,
|
"round": 0,
|
||||||
"max_rounds": 2,
|
"max_rounds": max_rounds,
|
||||||
}
|
}
|
||||||
result = graph.invoke(initial_state)
|
|
||||||
print("\n--- Final Answer ---")
|
print("=== Reflection Agent ===")
|
||||||
print(result["draft"])
|
print(f"Question: {question}")
|
||||||
print("\n--- Critique ---")
|
|
||||||
print(result["critique"])
|
for chunk in graph.stream(state, stream_mode="updates"):
|
||||||
|
for node_name, update in chunk.items():
|
||||||
|
state.update(update)
|
||||||
|
print()
|
||||||
|
if node_name == "draft_answer":
|
||||||
|
print("Draft:")
|
||||||
|
print(state["draft"])
|
||||||
|
elif node_name == "reflect":
|
||||||
|
print(f"Critic verdict: {state['verdict']}")
|
||||||
|
print("Critique:")
|
||||||
|
print(state["critique"] or "- no critique")
|
||||||
|
elif node_name == "rewrite":
|
||||||
|
print(f"Rewritten draft after round {state['round']}:")
|
||||||
|
print(state["draft"])
|
||||||
|
|
||||||
|
print()
|
||||||
|
print("Final answer:")
|
||||||
|
print(state["draft"])
|
||||||
|
print()
|
||||||
|
print(f"Revision rounds used: {state['round']} / {state['max_rounds']}")
|
||||||
|
return state
|
||||||
|
|
||||||
|
|
||||||
|
def main() -> None:
|
||||||
|
question = " ".join(sys.argv[1:]).strip() or DEFAULT_QUESTION
|
||||||
|
run_with_log(question)
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
main()
|
||||||
|
|||||||
Reference in New Issue
Block a user