diff --git a/main.py b/main.py index 3720b6e..4c9d18b 100644 --- a/main.py +++ b/main.py @@ -1,117 +1,174 @@ -""" -LangGraph Reflection Agent - -Usage: - python main.py "Your question" -""" +import json +import os +import re 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): question: str draft: str critique: str - verdict: str # 'ok' or 'needs_revision' + verdict: str round: 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: - from langchain_openai import ChatOpenAI - llm = ChatOpenAI(model="gpt-3.5-turbo") - 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." +class ReflectionResult(BaseModel): + verdict: Literal["ok", "needs_revision"] = Field( + description="Whether the draft is acceptable or needs one more revision." ) - response = await llm.invoke(prompt) - # Simple parsing - 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)." + critique: list[str] = Field( + description="Two or three short critique bullets in Russian." ) - 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 -start_edge = "draft_answer" -end_edge = None # will be set in condition +def draft_answer(state: ReflectState) -> dict: + llm = build_llm() + 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": - return "END" - if state["round"] >= state.get("max_rounds", 2): - return "END" + return "finish" + if state["round"] >= state["max_rounds"]: + return "finish" return "rewrite" -# Add edges with condition -from langgraph.graph import END -graph_builder.set_entry_point(start_edge) +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() -graph_builder.add_conditional_edges( - "reflect", - should_end, - { - "rewrite": "rewrite", - "END": END, - }, -) -# rewrite -> reflect -graph_builder.add_edge("rewrite", "reflect") - -graph = graph_builder.compile() - -if __name__ == "__main__": - if len(sys.argv) < 2: - print("Usage: python main.py \"Your question\"") - sys.exit(1) - question = sys.argv[1] - initial_state: ReflectState = { +def run_with_log(question: str, max_rounds: int = 2) -> ReflectState: + graph = build_graph() + state: ReflectState = { "question": question, "draft": "", "critique": "", "verdict": "", "round": 0, - "max_rounds": 2, + "max_rounds": max_rounds, } - result = graph.invoke(initial_state) - print("\n--- Final Answer ---") - print(result["draft"]) - print("\n--- Critique ---") - print(result["critique"]) + + 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()