diff --git a/main.py b/main.py index 4c9d18b..dcfc8eb 100644 --- a/main.py +++ b/main.py @@ -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"])