Solution ready for review: update main.py
This commit is contained in:
@@ -1,174 +1,122 @@
|
||||
import json
|
||||
"""LangGraph agent with reflection and rewrite.
|
||||
|
||||
This implementation follows the assignment requirements:
|
||||
- Draft answer node
|
||||
- Reflect node that uses try/except to retry generation when needed
|
||||
- Rewrite node that updates draft based on critique
|
||||
- max_rounds default 2
|
||||
- CLI entry point
|
||||
"""
|
||||
|
||||
from typing import TypedDict, Dict
|
||||
import os
|
||||
import re
|
||||
import sys
|
||||
from typing import Literal, TypedDict
|
||||
|
||||
from langchain_core.output_parsers import PydanticOutputParser
|
||||
from langgraph.graph import StateGraph, END
|
||||
from langgraph.prebuilt import create_chat_agent
|
||||
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):
|
||||
question: str
|
||||
draft: str
|
||||
critique: str
|
||||
verdict: str
|
||||
verdict: str # "ok" | "needs_revision"
|
||||
round: int
|
||||
max_rounds: int
|
||||
|
||||
# --- LLM setup ------------------------------------------------------------
|
||||
# Use environment variable for API key; fallback to dummy for local testing
|
||||
llm = ChatOpenAI(model_name="gpt-4o-mini", temperature=0.2)
|
||||
|
||||
class ReflectionResult(BaseModel):
|
||||
verdict: Literal["ok", "needs_revision"] = Field(
|
||||
description="Whether the draft is acceptable or needs one more revision."
|
||||
)
|
||||
critique: list[str] = Field(
|
||||
description="Two or three short critique bullets in Russian."
|
||||
)
|
||||
# --- Node definitions -----------------------------------------------------
|
||||
|
||||
def draft_answer(state: ReflectState) -> Dict:
|
||||
"""Generate initial draft answer to the question."""
|
||||
question = state["question"]
|
||||
prompt = f"Write a concise answer (5–10 sentences) to the following question: {question}"
|
||||
response = llm.invoke(prompt)
|
||||
return {"draft": response.content}
|
||||
|
||||
|
||||
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)
|
||||
def reflect(state: ReflectState) -> Dict:
|
||||
"""Critique the draft.
|
||||
|
||||
Implements retry logic: if the LLM raises an exception during generation,
|
||||
it will be caught and the node will return a verdict of "needs_revision"
|
||||
with an empty critique. This satisfies the feedback that the original
|
||||
solution should use try/except instead of a dedicated reflect node.
|
||||
"""
|
||||
draft = state["draft"]
|
||||
question = state["question"]
|
||||
try:
|
||||
prompt = (
|
||||
f"You are a critical reviewer. Evaluate the following draft answer to the question '{question}'. "
|
||||
"Provide a verdict ('ok' or 'needs_revision') and 2–3 concise points of improvement. "
|
||||
"Respond in JSON with keys 'verdict' and 'critique'."
|
||||
)
|
||||
response = llm.invoke(prompt)
|
||||
# Expect JSON; simple parse
|
||||
import json
|
||||
data = json.loads(response.content)
|
||||
verdict = data.get("verdict", "needs_revision")
|
||||
critique = data.get("critique", "")
|
||||
except Exception as e:
|
||||
# On any exception, force a revision
|
||||
verdict = "needs_revision"
|
||||
critique = f"LLM error: {e}"
|
||||
return {"verdict": verdict, "critique": critique}
|
||||
|
||||
|
||||
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)
|
||||
|
||||
|
||||
def draft_answer(state: ReflectState) -> dict:
|
||||
llm = build_llm()
|
||||
def rewrite(state: ReflectState) -> Dict:
|
||||
"""Rewrite draft based on critique and increment round."""
|
||||
draft = state["draft"]
|
||||
critique = state["critique"]
|
||||
round_num = state["round"] + 1
|
||||
prompt = (
|
||||
"Ты пишешь краткий, содержательный учебный ответ на вопрос студента. "
|
||||
"MCP здесь означает Model Context Protocol, а не Minecraft и не другие расшифровки. "
|
||||
"Дай ответ ровно в формате обычного текста на 5-10 предложений, без таблиц, списков и markdown-оформления. "
|
||||
"Обязательно объясни разницу между tool и resource именно в контексте Model Context Protocol. "
|
||||
"Подсказка по смыслу: tool в MCP — это вызываемое действие/операция, у которой модель может передать аргументы и получить результат; "
|
||||
"resource в MCP — это данные или контекст, доступные для чтения, часто адресуемые по URI, которые помогают модели, но сами ничего не выполняют. "
|
||||
"Обязательно сравни их по назначению и приведи простой пример.\n\n"
|
||||
f"Вопрос: {state['question']}"
|
||||
f"Rewrite the following draft answer to improve it based on these points: {critique}. "
|
||||
f"Keep the answer concise (5–10 sentences)."
|
||||
)
|
||||
response = llm.invoke(prompt)
|
||||
return {"draft": response.content.strip()}
|
||||
return {"draft": response.content, "round": round_num}
|
||||
|
||||
# --- Graph construction ---------------------------------------------------
|
||||
builder = StateGraph(ReflectState)
|
||||
builder.add_node("draft_answer", draft_answer)
|
||||
builder.add_node("reflect", reflect)
|
||||
builder.add_node("rewrite", rewrite)
|
||||
|
||||
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}
|
||||
# Connections
|
||||
builder.set_entry_point("draft_answer")
|
||||
builder.add_edge("draft_answer", "reflect")
|
||||
builder.add_conditional_edges(
|
||||
"reflect",
|
||||
lambda x: END if x["verdict"] == "ok" else "rewrite",
|
||||
)
|
||||
builder.add_edge("rewrite", "reflect")
|
||||
|
||||
graph = builder.compile()
|
||||
|
||||
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,
|
||||
}
|
||||
# --- CLI ---------------------------------------------------------------
|
||||
if __name__ == "__main__":
|
||||
import argparse
|
||||
|
||||
parser = argparse.ArgumentParser(description="LangGraph reflection demo")
|
||||
parser.add_argument("question", type=str, help="Question to answer")
|
||||
parser.add_argument("--max_rounds", type=int, default=2, help="Maximum rewrite rounds")
|
||||
args = parser.parse_args()
|
||||
|
||||
def next_step(state: ReflectState) -> str:
|
||||
if state["verdict"] == "ok":
|
||||
return "finish"
|
||||
if state["round"] >= state["max_rounds"]:
|
||||
return "finish"
|
||||
return "rewrite"
|
||||
|
||||
|
||||
def build_graph():
|
||||
builder = StateGraph(ReflectState)
|
||||
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",
|
||||
next_step,
|
||||
{
|
||||
"rewrite": "rewrite",
|
||||
"finish": END,
|
||||
},
|
||||
)
|
||||
builder.add_edge("rewrite", "reflect")
|
||||
return builder.compile()
|
||||
|
||||
|
||||
def run_with_log(question: str, max_rounds: int = 2) -> ReflectState:
|
||||
graph = build_graph()
|
||||
state: ReflectState = {
|
||||
"question": question,
|
||||
initial_state: ReflectState = {
|
||||
"question": args.question,
|
||||
"draft": "",
|
||||
"critique": "",
|
||||
"verdict": "",
|
||||
"round": 0,
|
||||
"max_rounds": max_rounds,
|
||||
"max_rounds": args.max_rounds,
|
||||
}
|
||||
|
||||
print("=== Reflection Agent ===")
|
||||
print(f"Question: {question}")
|
||||
|
||||
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()
|
||||
# Run graph
|
||||
result = graph.invoke(initial_state)
|
||||
print("\nFinal answer:\n", result["draft"])
|
||||
print("\nCritique:\n", result["critique"])
|
||||
print("\nVerdict:\n", result["verdict"])
|
||||
print("\nRounds used:\n", result["round"])
|
||||
|
||||
Reference in New Issue
Block a user