diff --git a/check_tasks.py b/check_tasks.py index b444340..ab0daf9 100644 --- a/check_tasks.py +++ b/check_tasks.py @@ -1,4 +1,6 @@ -import asyncio, json +import asyncio, json, os +# Обходим локальный прокси для всех наших сервисов (иначе SSL-ошибка) +os.environ["NO_PROXY"] = "openrouter.ai,platform.brojs.ru,git.brojs.ru," + os.environ.get("NO_PROXY", "") from dotenv import load_dotenv load_dotenv() from src.agent.mcp_client import load_journal_toolsets, JOURNAL_PREFIX diff --git a/get_task_ids.py b/get_task_ids.py new file mode 100644 index 0000000..18183d5 --- /dev/null +++ b/get_task_ids.py @@ -0,0 +1,24 @@ +import asyncio, json, os +os.environ["NO_PROXY"] = "openrouter.ai,platform.brojs.ru,git.brojs.ru," + os.environ.get("NO_PROXY", "") +from dotenv import load_dotenv +load_dotenv() +from src.agent.mcp_client import load_journal_toolsets, JOURNAL_PREFIX +from src.agent.constants import COURSE_ID + +async def main(): + j = load_journal_toolsets() + tool = next(t for t in j.tasks_submissions_tools if t.name == f"{JOURNAL_PREFIX}tasks_list") + raw = await tool.ainvoke({"courseId": COURSE_ID}) + if isinstance(raw, list): + raw = next((x["text"] for x in raw if x.get("type") == "text"), str(raw)) + data = json.loads(raw) if isinstance(raw, str) else raw + items = data.get("tasks", data) if isinstance(data, dict) else data + for item in items: + t = item.get("task", item) if isinstance(item, dict) else {} + tid = t.get("id", "") + status = item.get("status", "") + title = t.get("title", "") + if status == "todo": + print(f"TODO {tid} {title}") + +asyncio.run(main()) diff --git a/read_tasks.py b/read_tasks.py new file mode 100644 index 0000000..3a9e82f --- /dev/null +++ b/read_tasks.py @@ -0,0 +1,25 @@ +import asyncio, json, os +os.environ["NO_PROXY"] = "openrouter.ai,platform.brojs.ru,git.brojs.ru," + os.environ.get("NO_PROXY", "") +from dotenv import load_dotenv +load_dotenv() +from src.agent.mcp_client import load_journal_toolsets, JOURNAL_PREFIX + +TASK_IDS = [ + "6a1864fd", # Планирующий агент + "6a186500", # Структурированный вывод (Pydantic) + "6a1864f7", # RAG-агент с ChromaDB + "6a1864fa", # Самокорректирующийся агент +] + +async def main(): + j = load_journal_toolsets() + text_tool = next(t for t in j.tasks_submissions_tools if t.name == f"{JOURNAL_PREFIX}task_text") + for tid in TASK_IDS: + raw = await text_tool.ainvoke({"taskId": tid}) + text = raw if isinstance(raw, str) else next((x["text"] for x in raw if x.get("type") == "text"), str(raw)) + print(f"\n{'='*60}") + print(f"ЗАДАНИЕ {tid}") + print('='*60) + print(text) + +asyncio.run(main()) diff --git a/run_pipeline.py b/run_pipeline.py index 2f92913..0c695c9 100644 --- a/run_pipeline.py +++ b/run_pipeline.py @@ -1,28 +1,64 @@ """Запуск пайплайна для выполнения заданий курса.""" import asyncio -import sys import os +import traceback os.environ["PYTHONIOENCODING"] = "utf-8" +# Обходим локальный прокси для OpenRouter (иначе SSL-ошибка) +os.environ["NO_PROXY"] = "openrouter.ai,platform.brojs.ru,git.brojs.ru," + os.environ.get("NO_PROXY", "") from dotenv import load_dotenv load_dotenv() -from langchain_core.messages import HumanMessage -from src.agent.graph.pipeline import pipeline +from src.agent.graph.pipeline import pipeline, TaskInfo, process_one_task, route, PipelineState + +# Целевые задания (None = все todo-задания автоматически) +TARGET_IDS = [ + "6a1864f78a94f887e50d46da", # Экзамен: RAG-агент с ChromaDB и веб-поиском +] async def main(): print("=== Запуск пайплайна BroJS ===") - result = await pipeline.ainvoke( - { - "tasks": [], + + if TARGET_IDS: + tasks = [TaskInfo(id=tid, title="", status="todo") for tid in TARGET_IDS] + print(f"Целевые задания: {TARGET_IDS}") + initial_state = { + "tasks": tasks, "current_index": 0, "results": [], "errors": [], - }, - {"configurable": {"thread_id": "pipeline-main"}}, - ) + } + from langgraph.graph import StateGraph, START + builder = StateGraph(PipelineState) + builder.add_node("process_one_task", process_one_task) + builder.add_edge(START, "process_one_task") + builder.add_conditional_edges( + "process_one_task", route, + {"process_one_task": "process_one_task", "__end__": "__end__"} + ) + targeted_pipeline = builder.compile() + try: + result = await targeted_pipeline.ainvoke( + initial_state, + {"configurable": {"thread_id": "pipeline-targeted"}}, + ) + except Exception as e: + print(f"КРИТИЧЕСКАЯ ОШИБКА: {e}") + traceback.print_exc() + return + else: + try: + result = await pipeline.ainvoke( + {"tasks": [], "current_index": 0, "results": [], "errors": []}, + {"configurable": {"thread_id": "pipeline-main"}}, + ) + except Exception as e: + print(f"КРИТИЧЕСКАЯ ОШИБКА: {e}") + traceback.print_exc() + return + print("\n=== Результат пайплайна ===") for r in result.get("results", []): print(f" Task {r['task_id'][:8]}: {r['status']} (mode={r['mode']}, retries={r['retries']})") diff --git a/solve_task.py b/solve_task.py new file mode 100644 index 0000000..2fa8a9f --- /dev/null +++ b/solve_task.py @@ -0,0 +1,360 @@ +""" +Быстрый решатель заданий BroJS. +Схема: читаем задание (MCP) → 1 LLM-вызов → пушим на Gitea → сабмитим (MCP). +Автономный — не импортирует src.agent, нет двойной загрузки MCP. + +Использование: + python solve_task.py +""" +import asyncio +import base64 +import json +import os +import sys + +# Обходим локальный прокси +os.environ["NO_PROXY"] = "openrouter.ai,platform.brojs.ru,git.brojs.ru," + os.environ.get("NO_PROXY", "") + +from dotenv import load_dotenv +load_dotenv() + +import httpx +from langchain_mcp_adapters.client import MultiServerMCPClient +from langchain_openai import ChatOpenAI + +GITEA_BASE_URL = "https://git.brojs.ru" +GITEA_OWNER = os.getenv("GITEA_OWNER", "glevelll") +GITEA_TOKEN = os.getenv("GITEA_TOKEN", "") +JOURNAL_TOKEN = os.getenv("JOURNAL_TOKEN", "") +OPENAI_API_KEY = os.getenv("OPENAI_API_KEY", "") + +MCP_URL = "https://platform.brojs.ru/jrnl-bh/api/mcp" + +# --------------------------------------------------------------------------- +# LLM +# --------------------------------------------------------------------------- + +llm = ChatOpenAI( + model="openai/gpt-oss-20b:free", + base_url="https://openrouter.ai/api/v1", + api_key=OPENAI_API_KEY, + temperature=0.0, + max_tokens=4096, +) + +# --------------------------------------------------------------------------- +# MCP — один клиент на весь запуск +# --------------------------------------------------------------------------- + +_mcp_tools: dict = {} + + +async def _load_mcp(retries=5, pause=30): + global _mcp_tools + if _mcp_tools: + return + config = { + "journal": { + "transport": "streamable_http", + "url": MCP_URL, + "headers": {"Authorization": f"Bearer {JOURNAL_TOKEN}"}, + } + } + client = MultiServerMCPClient(config) + for attempt in range(1, retries + 1): + try: + tools = await client.get_tools(server_name="journal") + _mcp_tools = {t.name: t for t in tools} + print(f" [mcp] Загружено {len(_mcp_tools)} инструментов") + return + except Exception as e: + if "429" in str(e) and attempt < retries: + print(f" [mcp] 429 при загрузке, жду {pause}с...") + await asyncio.sleep(pause) + else: + raise + + +async def mcp_call(name: str, args: dict, retries=5, pause=30): + await _load_mcp() + tool = _mcp_tools.get(name) + if not tool: + raise RuntimeError(f"MCP tool '{name}' not found. Available: {list(_mcp_tools.keys())}") + for attempt in range(1, retries + 1): + try: + result = await tool.ainvoke(args) + if isinstance(result, list): + return next((x["text"] for x in result if x.get("type") == "text"), str(result)) + return str(result) + except Exception as e: + if "429" in str(e) and attempt < retries: + print(f" [mcp] {name} → 429, жду {pause}с (попытка {attempt}/{retries})...") + await asyncio.sleep(pause) + else: + raise + + +# --------------------------------------------------------------------------- +# Gitea +# --------------------------------------------------------------------------- + +def _gh(): + return {"Authorization": f"token {GITEA_TOKEN}", "Content-Type": "application/json"} + + +def gitea_create_repo(name: str) -> str: + with httpx.Client(timeout=30) as c: + r = c.post(f"{GITEA_BASE_URL}/api/v1/user/repos", headers=_gh(), + json={"name": name, "private": False, "auto_init": False}) + if r.status_code == 409: + return f"{GITEA_BASE_URL}/{GITEA_OWNER}/{name}" + r.raise_for_status() + return r.json().get("html_url", f"{GITEA_BASE_URL}/{GITEA_OWNER}/{name}") + + +def gitea_write(repo: str, path: str, content: str, msg: str): + encoded = base64.b64encode(content.encode()).decode() + url = f"{GITEA_BASE_URL}/api/v1/repos/{GITEA_OWNER}/{repo}/contents/{path}" + with httpx.Client(timeout=30) as c: + r = c.get(url, headers=_gh()) + if r.status_code == 200: + sha = r.json().get("sha", "") + c.put(url, headers=_gh(), json={"message": msg, "content": encoded, "sha": sha}).raise_for_status() + else: + c.post(url, headers=_gh(), json={"message": msg, "content": encoded}).raise_for_status() + + +# --------------------------------------------------------------------------- +# LLM: генерация кода +# --------------------------------------------------------------------------- + +_PROMPT = '''\ +Ты — Python-разработчик. Напиши решение для учебного задания по LLM/AI. +Используй фреймворк deepagents (create_deep_agent) — это обязательное требование курса. + +## Задание +{task_text} + +## ОБЯЗАТЕЛЬНЫЕ ТЕХНИЧЕСКИЕ ПАТТЕРНЫ + +### LLM — всегда OpenRouter: +```python +import os +from langchain_openai import ChatOpenAI +llm = ChatOpenAI( + model="openai/gpt-oss-20b:free", + base_url="https://openrouter.ai/api/v1", + api_key=os.getenv("OPENAI_API_KEY"), + temperature=0.0, +) +``` + +### Базовый агент (deepagents) — ОБЯЗАТЕЛЬНАЯ основа: +```python +import asyncio, os +from langchain_openai import ChatOpenAI +from langchain_core.messages import HumanMessage +from langchain.tools import tool +from deepagents import create_deep_agent +from deepagents.backends import FilesystemBackend, LocalShellBackend, CompositeBackend + +llm = ChatOpenAI(model="openai/gpt-oss-20b:free", base_url="https://openrouter.ai/api/v1", api_key=os.getenv("OPENAI_API_KEY")) + +backend = CompositeBackend([ + LocalShellBackend(workspace_dir="./workspace"), + FilesystemBackend(), +]) + +@tool +def my_tool(query: str) -> str: + """Tool description.""" + return f"result for {{query}}" + +agent = create_deep_agent( + model=llm, + tools=[my_tool], + backend=backend, + system_prompt="You are a helpful agent.", +) + +async def main(): + result = await agent.ainvoke( + {{"messages": [HumanMessage(content="Your task here")]}}, + {{"configurable": {{"thread_id": "session-1"}}}}, + ) + print(result["messages"][-1].content) + +if __name__ == "__main__": + asyncio.run(main()) +``` +requirements.txt: deepagents, langchain-openai>=0.3.0, langchain>=1.2.10, langgraph>=0.2.0 + +### RAG с Qdrant (для RAG-заданий): +```python +from langchain_openai import OpenAIEmbeddings +from langchain_qdrant import QdrantVectorStore +from qdrant_client import QdrantClient +from qdrant_client.models import Distance, VectorParams +from langchain_core.documents import Document + +embeddings = OpenAIEmbeddings(model="text-embedding-3-small", base_url="https://openrouter.ai/api/v1", api_key=os.getenv("OPENAI_API_KEY")) +client = QdrantClient(":memory:") +client.create_collection("knowledge", vectors_config=VectorParams(size=1536, distance=Distance.COSINE)) +vector_store = QdrantVectorStore(client=client, collection_name="knowledge", embedding=embeddings) + +@tool +def search_knowledge(query: str) -> str: + """Search the knowledge base.""" + docs = vector_store.similarity_search(query, k=3) + return "\\n".join(d.page_content for d in docs) if docs else "No results." + +@tool +def add_to_knowledge(content: str, title: str = "doc") -> str: + """Add content to knowledge base.""" + vector_store.add_documents([Document(page_content=content, metadata={{"title": title}})]) + return f"Added: {{title}}" +``` +requirements.txt добавить: langchain-qdrant, qdrant-client + +### Планирующий агент (для planning-заданий): +```python +from langgraph.graph import StateGraph, START, END +from typing import TypedDict, Annotated +from langgraph.graph.message import add_messages + +class PlanState(TypedDict): + messages: Annotated[list, add_messages] + plan: list[str] + current_step: int + +def planner_node(state): + # LLM создаёт план + ... + +def executor_node(state): + # LLM выполняет шаг плана + ... +``` + +### Самокорректирующийся агент: +```python +# Агент проверяет свой вывод и исправляет если нужно +@tool +def validate_output(output: str) -> str: + """Validate the output and return issues if any.""" + issues = [] + if len(output) < 10: + issues.append("Output too short") + return "OK" if not issues else f"Issues: {{', '.join(issues)}}" +``` + +### Структурированный вывод (Pydantic): +```python +from pydantic import BaseModel, Field +from langchain_core.output_parsers import PydanticOutputParser + +class MyOutput(BaseModel): + field1: str = Field(description="...") + field2: int = Field(description="...") + +parser = PydanticOutputParser(pydantic_object=MyOutput) +``` + +## Требования +- Полный рабочий код без заглушек (no pass, TODO, ...) +- ОБЯЗАТЕЛЬНО использовать create_deep_agent из deepagents +- requirements.txt: deepagents, langchain>=1.2.10, langchain-openai>=0.3.0, langgraph>=0.2.0 + нужные доп. зависимости + +## Ответ — ТОЛЬКО JSON без markdown: +{{"main_py": "...", "requirements_txt": "...", "extra_files": {{}}}} + +extra_files — только если нужны доп. файлы, иначе пустой объект. +''' + + +async def generate(task_text: str, retries=5) -> dict: + prompt = _PROMPT.format(task_text=task_text) + for attempt in range(1, retries + 1): + try: + print(f" [llm] Генерирую решение (попытка {attempt})...") + resp = await llm.ainvoke(prompt) + raw = resp.content.strip() + if raw.startswith("```"): + raw = raw.split("```")[1] + if raw.startswith("json"): + raw = raw[4:] + return json.loads(raw.strip()) + except json.JSONDecodeError as e: + print(f" [llm] JSON parse error: {e}. Повтор...") + if attempt == retries: + raise + except Exception as e: + if "429" in str(e) and attempt < retries: + wait = 90 * attempt + print(f" [llm] 429, жду {wait}с (попытка {attempt}/{retries})...") + await asyncio.sleep(wait) + else: + raise + + +# --------------------------------------------------------------------------- +# Основная логика +# --------------------------------------------------------------------------- + +async def solve(task_id: str): + print(f"\n{'='*60}") + print(f"Задание: {task_id}") + print('='*60) + + # 1. Читаем текст задания + print("[1/5] Читаем текст задания...") + task_text = await mcp_call("task_text", {"taskId": task_id}) + print(f" Получено {len(task_text)} символов") + + # 2. Генерируем код + print("[2/5] Генерируем код (1 LLM-вызов)...") + solution = await generate(task_text) + main_py = solution.get("main_py", "") + requirements = solution.get("requirements_txt", "") + extra = solution.get("extra_files", {}) + print(f" main.py: {len(main_py)} символов, requirements.txt: {len(requirements)} символов") + + # 3. Создаём репо + repo = f"task-{task_id}" + print(f"[3/5] Создаём репозиторий {repo}...") + repo_url = gitea_create_repo(repo) + print(f" {repo_url}") + + # 4. Пушим файлы + print("[4/5] Пушим файлы...") + gitea_write(repo, "main.py", main_py, "add main.py") + print(" main.py ✓") + gitea_write(repo, "requirements.txt", requirements, "add requirements.txt") + print(" requirements.txt ✓") + for fname, fcontent in extra.items(): + gitea_write(repo, fname, fcontent, f"add {fname}") + print(f" {fname} ✓") + + # 5. Сабмитим + print("[5/5] Сабмитим...") + await mcp_call("task_update_answer", { + "taskId": task_id, "answerType": "link", "content": repo_url, + }) + print(" task_update_answer ✓") + await asyncio.sleep(3) + await mcp_call("task_submit", {"taskId": task_id, "confirmSubmit": True}) + print(" task_submit ✓") + + print(f"\n✅ Готово! Репозиторий: {repo_url}") + return repo_url + + +async def main(): + if len(sys.argv) < 2: + print("Использование: python solve_task.py ") + sys.exit(1) + await solve(sys.argv[1]) + + +if __name__ == "__main__": + asyncio.run(main()) diff --git a/src/agent/__init__.py b/src/agent/__init__.py index 5fb1392..898ad77 100644 --- a/src/agent/__init__.py +++ b/src/agent/__init__.py @@ -1,3 +1,2 @@ -from src.agent.agent import agent, homework_direct_agent, rework_agent - -__all__ = ["agent", "homework_direct_agent", "rework_agent"] +# Намеренно пустой — предотвращает двойную загрузку MCP при импорте подмодулей. +# Импортируй напрямую: from src.agent.agent import agent diff --git a/src/agent/agent.py b/src/agent/agent.py index 7cfa312..c027b90 100644 --- a/src/agent/agent.py +++ b/src/agent/agent.py @@ -12,7 +12,7 @@ from src.agent.constants import ( from src.agent.gitea_tools import GITEA_TOOLS from src.agent.llm import llm from src.agent.mcp_client import load_journal_toolsets -from src.agent.middlewares import SanitizeToolCallsMiddleware, ValidateJournalWorkflowMiddleware +from src.agent.middlewares import RetryOnRateLimitMiddleware, SanitizeToolCallsMiddleware, ValidateJournalWorkflowMiddleware from src.agent.prompts import ( homework_doing_instructions, main_agent_instructions, @@ -129,6 +129,7 @@ homework_direct_agent = create_deep_agent( system_prompt=homework_doing_instructions, backend=_composite_backend, middleware=[ + RetryOnRateLimitMiddleware(), SanitizeToolCallsMiddleware(known_tools=_subagent_tool_names["homework_doing"]), ValidateJournalWorkflowMiddleware(), ], @@ -144,6 +145,7 @@ rework_agent = create_deep_agent( system_prompt=rework_instructions, backend=_composite_backend, middleware=[ + RetryOnRateLimitMiddleware(), SanitizeToolCallsMiddleware(known_tools=_subagent_tool_names["homework_doing"]), ValidateJournalWorkflowMiddleware(), ], diff --git a/src/agent/graph/pipeline.py b/src/agent/graph/pipeline.py index 7569393..05b1301 100644 --- a/src/agent/graph/pipeline.py +++ b/src/agent/graph/pipeline.py @@ -10,10 +10,10 @@ from typing import TypedDict from langchain_core.messages import HumanMessage from langgraph.graph import START, StateGraph -from src.agent.agent import homework_direct_agent, rework_agent +from src.agent.agent import homework_direct_agent, journal as _journal_toolsets, rework_agent from src.agent.constants import COURSE_ID, GITEA_OWNER from src.agent.gitea_tools import _get as gitea_get -from src.agent.mcp_client import JOURNAL_PREFIX, load_journal_toolsets +from src.agent.mcp_client import JOURNAL_PREFIX # --------------------------------------------------------------------------- # Типы состояния @@ -36,11 +36,8 @@ class PipelineState(TypedDict): # Вспомогательные функции # --------------------------------------------------------------------------- -_journal = load_journal_toolsets() - - def _get_journal_tool(suffix: str): - all_tools = _journal.courses_lessons_tools + _journal.tasks_submissions_tools + all_tools = _journal_toolsets.courses_lessons_tools + _journal_toolsets.tasks_submissions_tools target = f"{JOURNAL_PREFIX}{suffix}" for t in all_tools: if t.name == target: @@ -116,13 +113,13 @@ async def _force_submit(task_id: str) -> bool: print(f"[pipeline] force_submit: инструменты не найдены") return False try: - await update_tool.ainvoke({ + await _mcp_invoke(update_tool, { "taskId": task_id, "answerType": "link", "content": repo_url, "commit": {"repoUrl": repo_url, "branch": "main"}, }) - await submit_tool.ainvoke({"taskId": task_id, "confirmSubmit": True}) + await _mcp_invoke(submit_tool, {"taskId": task_id, "confirmSubmit": True}) print(f"[pipeline] Задание {task_id[:8]} — сабмит выполнен пайплайном ✓") return True except Exception as e: @@ -137,12 +134,25 @@ async def _is_submitted(task_id: str) -> bool: return status in ("ready_for_review", "done") +async def _mcp_invoke(tool, args: dict, retries: int = 5, pause: int = 30): + """Вызывает MCP-инструмент с retry при 429.""" + for attempt in range(1, retries + 1): + try: + return await tool.ainvoke(args) + except Exception as e: + if "429" in str(e) and attempt < retries: + print(f"[pipeline] MCP 429, жду {pause}с (попытка {attempt}/{retries})...") + await asyncio.sleep(pause) + else: + raise + + async def _task_text(task_id: str) -> str: tool = _get_journal_tool("task_text") if not tool: return "" try: - return _parse_text(await tool.ainvoke({"taskId": task_id})) + return _parse_text(await _mcp_invoke(tool, {"taskId": task_id})) except Exception: return "" @@ -152,7 +162,7 @@ async def _task_json(task_id: str) -> dict: if not tool: return {} try: - raw = _parse_text(await tool.ainvoke({"taskId": task_id})) + raw = _parse_text(await _mcp_invoke(tool, {"taskId": task_id})) return json.loads(raw) except Exception: return {} @@ -276,6 +286,10 @@ async def process_one_task(state: PipelineState) -> dict: results = list(state.get("results", [])) errors = list(state.get("errors", [])) + # Пауза перед стартом — даём BroJS MCP сбросить rate limit после загрузки инструментов + print(f"[pipeline] Задание {task_id[:8]} — пауза 10с перед стартом...") + await asyncio.sleep(10) + repo_url = await _existing_repo_url(task_id) is_rework = repo_url is not None @@ -348,8 +362,16 @@ async def process_one_task(state: PipelineState) -> dict: "retries": retries, }) - except Exception as e: - print(f"[pipeline] Задание {task_id[:8]} — ОШИБКА: {e}") + except BaseException as e: + import traceback + # Разворачиваем ExceptionGroup (Python 3.11+) чтобы увидеть реальные ошибки + if isinstance(e, ExceptionGroup): + for i, sub in enumerate(e.exceptions): + print(f"[pipeline] Задание {task_id[:8]} — под-ошибка {i+1}: {type(sub).__name__}: {sub}") + traceback.print_exception(type(sub), sub, sub.__traceback__) + else: + print(f"[pipeline] Задание {task_id[:8]} — ОШИБКА: {type(e).__name__}: {e}") + traceback.print_exc() errors.append(f"Задание {task_id} ({'rework' if is_rework else 'new'}): {e}") # Пауза между заданиями чтобы не перегружать rate limit diff --git a/src/agent/middlewares/__init__.py b/src/agent/middlewares/__init__.py index 382af07..ee7f120 100644 --- a/src/agent/middlewares/__init__.py +++ b/src/agent/middlewares/__init__.py @@ -1,4 +1,5 @@ from src.agent.middlewares.sanitize_tool_calls import SanitizeToolCallsMiddleware from src.agent.middlewares.validate_journal_workflow import ValidateJournalWorkflowMiddleware +from src.agent.middlewares.retry_on_rate_limit import RetryOnRateLimitMiddleware -__all__ = ["SanitizeToolCallsMiddleware", "ValidateJournalWorkflowMiddleware"] +__all__ = ["SanitizeToolCallsMiddleware", "ValidateJournalWorkflowMiddleware", "RetryOnRateLimitMiddleware"] diff --git a/src/agent/middlewares/retry_on_rate_limit.py b/src/agent/middlewares/retry_on_rate_limit.py new file mode 100644 index 0000000..1414570 --- /dev/null +++ b/src/agent/middlewares/retry_on_rate_limit.py @@ -0,0 +1,46 @@ +"""Middleware: повторяет вызов инструмента при 429 Rate Limit.""" +from __future__ import annotations + +import asyncio +from typing import Any + +from langchain.agents.middleware import AgentMiddleware, AgentState +from langchain_core.messages import ToolMessage + + +_PAUSE = 30 # секунд ожидания при 429 +_TRIES = 5 # максимум попыток + + +def _is_429(exc: Exception) -> bool: + msg = str(exc) + return "429" in msg or "rate" in msg.lower() + + +class RetryOnRateLimitMiddleware(AgentMiddleware[AgentState[Any], Any]): + """Перехватывает 429 от любого инструмента и повторяет с паузой.""" + + def wrap_tool_call(self, request, handler): + for attempt in range(1, _TRIES + 1): + try: + return handler(request) + except Exception as e: + if _is_429(e) and attempt < _TRIES: + name = request.tool_call.get("name", "") + print(f"[retry-mw] {name} → 429, жду {_PAUSE}с (попытка {attempt}/{_TRIES})...") + import time + time.sleep(_PAUSE) + else: + raise + + async def awrap_tool_call(self, request, handler): + for attempt in range(1, _TRIES + 1): + try: + return await handler(request) + except Exception as e: + if _is_429(e) and attempt < _TRIES: + name = request.tool_call.get("name", "") + print(f"[retry-mw] {name} → 429, жду {_PAUSE}с (попытка {attempt}/{_TRIES})...") + await asyncio.sleep(_PAUSE) + else: + raise