Files
brojs-task-6a1d75d1fd30e81c…/main.py
T

175 lines
7.1 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
import json
import os
import re
import sys
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"
class ReflectState(TypedDict):
question: str
draft: str
critique: str
verdict: str
round: int
max_rounds: int
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."
)
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 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()
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 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 "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,
"draft": "",
"critique": "",
"verdict": "",
"round": 0,
"max_rounds": 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()