/
fatboyslava
/
ffs_assistant
Обзор
Документация
Войти
/
fatboyslava
/
ffs_assistant
Код
Запросы
0
Задачи
Вики
Пакеты
0
Релизы
0
CI/CD
Аналитика
Безопасность
master
src/agent_loop.py
252 строки
7 KB
Viacheslav Kovalev
update easy
01 июл 2026, 23:37
01 июл 2026, 23:37
feb52fd
Код
Авторство
О чём код?
"""Main agent loop with TAO (Think-Act-Observe) pattern.""" import argparse import json import logging import sys from typing import Any from gigachat import GigaChat from gigachat.models import Chat, Messages, MessagesRole from pydantic import BaseModel from src.config import ( DATA_DIR, GIGACHAT_AUTH_URL, GIGACHAT_CREDENTIALS, GIGACHAT_SCOPE, USER_PROMPT, ) from src.prompts import FIXED_PART, USER_PART_TEMPLATE from src.utils import ( build_data_index, build_functions, validate_data_size, ) LOG_WIDTH = 72 TOOL_RESULT_PREVIEW_CHARS = 700 logger = logging.getLogger("ffs_assistant.agent") def configure_logging() -> None: logging.basicConfig( level=logging.INFO, format="%(message)s", stream=sys.stdout, force=True, ) def get_message(response: Any) -> Any: return response.choices[0].message def get_function_call(message: Any) -> Any | None: return getattr(message, "function_call", None) def get_arguments(function_call: Any) -> dict[str, Any]: arguments = getattr(function_call, "arguments", None) if arguments is None and isinstance(function_call, dict): arguments = function_call.get("arguments") if isinstance(arguments, str): try: return json.loads(arguments) except json.JSONDecodeError as error: return {"_invalid_json": arguments, "_error": str(error)} return arguments or {} def get_function_name(function_call: Any) -> str | None: if isinstance(function_call, dict): return function_call.get("name") return getattr(function_call, "name", None) def message_to_history(message: Any) -> Messages: return Messages( role=MessagesRole.ASSISTANT, content=getattr(message, "content", "") or "", function_call=getattr(message, "function_call", None), functions_state_id=getattr(message, "functions_state_id", None), ) def build_system_prompt() -> str: validate_data_size() data_index = build_data_index() user_part = USER_PART_TEMPLATE.format( data_dir=DATA_DIR, data_index=data_index, user_prompt=USER_PROMPT or "Не задан.", ) return "\n\n".join([FIXED_PART, user_part]).strip() def build_gigachat_kwargs() -> dict[str, Any]: giga_kwargs: dict[str, Any] = { "credentials": GIGACHAT_CREDENTIALS, "scope": GIGACHAT_SCOPE, "verify_ssl_certs": False, } if GIGACHAT_AUTH_URL: giga_kwargs["auth_url"] = GIGACHAT_AUTH_URL return giga_kwargs def build_chat( messages: list[Messages], functions: list[Any], ) -> Chat: return Chat( messages=messages, function_call="auto", functions=functions, ) def serialize_tool_result(result: Any) -> str: if isinstance(result, BaseModel): return result.model_dump_json(exclude_none=True) if isinstance(result, str): return result return json.dumps(result, ensure_ascii=False, default=str) def compact_json(value: Any) -> str: return json.dumps(value, ensure_ascii=False, default=str) def preview_text(value: str, limit: int = TOOL_RESULT_PREVIEW_CHARS) -> str: value = value.strip() if len(value) <= limit: return value return f"{value[:limit]}..." def log_turn(step: int, max_steps: int) -> None: logger.info("") logger.info("=" * LOG_WIDTH) logger.info("🔄 Turn %s/%s", step, max_steps) logger.info("=" * LOG_WIDTH) def log_assistant_content(content: str) -> None: if content: logger.info("") logger.info("🤖 %s", content) def log_tool_call(name: str | None, arguments: dict[str, Any]) -> None: logger.info("") logger.info("🔧 Tool: %s(%s)", name, compact_json(arguments)) def log_tool_result(result: str) -> None: logger.info(" → %s", preview_text(result)) def log_finished(content: str) -> None: if not content: logger.info("(no text output)") logger.info("✅ Agent finished") def log_max_steps(max_steps: int) -> None: logger.warning("") logger.warning("⚠️ Max turns (%s) reached. Stopping.", max_steps) def call_tool( name: str | None, arguments: dict[str, Any], ) -> str: from src.tools.grep_tool.grep_tool import grep from src.tools.read_tool.read_tool import file_read_tool tool = {"grep": grep, "read": file_read_tool}.get(name) if tool is None: return f"Error: unknown tool '{name}'" try: result = tool(**arguments) except Exception as error: return f"Error calling {name}: {type(error).__name__}: {error}" return serialize_tool_result(result) def call_llm( giga: GigaChat, messages: list[Messages], functions: list[Any], ) -> Any: response = giga.chat(build_chat(messages, functions)) return get_message(response) def agent_loop(question: str, max_steps: int = 12) -> str: if not GIGACHAT_CREDENTIALS: raise RuntimeError("Set GIGACHAT_CREDENTIALS in .env or environment") functions = build_functions() messages = [ Messages(role=MessagesRole.SYSTEM, content=build_system_prompt()), Messages(role=MessagesRole.USER, content=question), ] with GigaChat(**build_gigachat_kwargs()) as giga: for step in range(1, max_steps + 1): log_turn(step, max_steps) message = call_llm(giga, messages, functions) function_call = get_function_call(message) content = getattr(message, "content", "") or "" log_assistant_content(content) if not function_call: log_finished(content) return content function_name = get_function_name(function_call) arguments = get_arguments(function_call) log_tool_call(function_name, arguments) messages.append(message_to_history(message)) observation = call_tool(function_name, arguments) log_tool_result(observation) function_msg = Messages( role=MessagesRole.FUNCTION, name=function_name, content=observation, ) messages.append(function_msg) log_max_steps(max_steps) raise RuntimeError(f"Model did not finish after {max_steps} TAO steps") def main() -> None: configure_logging() parser = argparse.ArgumentParser(description="Primitive file-first TAO-loop") parser.add_argument("question", nargs="*", help="User question") parser.add_argument("--max-steps", type=int, default=8) args = parser.parse_args() try: validate_data_size() except RuntimeError as error: logger.error("Ошибка: %s", error) raise SystemExit(1) from None question = " ".join(args.question).strip() if not question: question = input("Вопрос: ").strip() agent_loop(question, max_steps=args.max_steps) if __name__ == "__main__": main()