/
igormayer
/
langgraph_sql_agent_web
Обзор
Документация
Войти
/
igormayer
/
langgraph_sql_agent_web
Код
Запросы
0
Задачи
Вики
Пакеты
0
Релизы
0
Аналитика
Безопасность
master
backend/sqlagent/graph.py
154 строки
6 KB
Igor Mayer
fixes for running in docker
08 фев 2026, 21:16
08 фев 2026, 21:16
988621e
Код
Авторство
О чём код?
from typing import Literal from langchain_core.callbacks.manager import dispatch_custom_event from langgraph.graph import StateGraph, END from langchain_core.runnables import RunnableConfig from langgraph.types import StreamWriter from strip_markdown import strip_markdown from backend.sqlagent.state import SqlAgentState from backend.sqlagent.factory import get_toolkit async def table_discovery( state: SqlAgentState, writer: StreamWriter, config: RunnableConfig ): """Этап 1: Поиск релевантных таблиц.""" writer({"message": "Изучаю структуру базы данных и подбираю таблицы для выполнения запроса..."}) tools_map, llm = get_toolkit(config) list_tables_tool = tools_map["sql_db_list_tables"] tables = await list_tables_tool.ainvoke("") prompt = f"Вопрос: {state['question']}. Доступные таблицы: {tables}. Напиши через запятую только названия таблиц, которые нужны для ответа." res = await llm.ainvoke(prompt) return {"tables": [t.strip() for t in res.content.split(",")]} async def schema_retrieval( state: SqlAgentState, writer: StreamWriter, config: RunnableConfig ): writer({"message": "Получаю схему таблиц..."}) tools_map, _ = get_toolkit(config) get_schema_tool = tools_map["sql_db_schema"] schema = await get_schema_tool.ainvoke(", ".join(state["tables"])) return {"schema": schema} async def sql_generation( state: SqlAgentState, writer: StreamWriter, config: RunnableConfig ): writer({"message": "Пишу SQL запрос..."}) _, llm = get_toolkit(config) prompt = f"Напиши SQL (Clickhouse) запрос. Схема: {state['schema']}. Вопрос: {state['question']}. Пиши ТОЛЬКО код." # Очистка от markdown-разметки если она есть query = await llm.ainvoke(prompt) return {"sql_query": query.content.replace("```sql", "").replace("```", "").strip()} async def sql_execution( state: SqlAgentState, writer: StreamWriter, config: RunnableConfig ): writer({"message": "Выполняю запрос и проверяю его на ошибки..."}) tools_map, _ = get_toolkit(config) exec_tool = tools_map["sql_db_query"] try: result = await exec_tool.ainvoke(state["sql_query"]) return {"sql_result": str(result), "error_count": 0} except Exception as e: return {"sql_result": str(e), "error_count": state.get("error_count", 0) + 1} async def answer_formulation( state: SqlAgentState, writer: StreamWriter, config: RunnableConfig ): writer({"message": "Формулирую ответ..."}) _, llm = get_toolkit(config) prompt = f"Вопрос: {state['question']}\nSQL: {state['sql_query']}\nРезультат: {state['sql_result']}\nДай ответ пользователю." res = await llm.ainvoke(prompt) return {"final_answer": res.content} # Логика переходов (Conditional Edges) async def should_retry( state: SqlAgentState, writer: StreamWriter ) -> Literal["retry", "continue"]: # Если в результате ошибка и мы не превысили 3 попытки — пробуем перегенерировать if "DB::Exception" in state["sql_result"] and state["error_count"] < 3: writer({"message": "Я ошибся, переделываю запрос..."}) return "retry" return "continue" builder = StateGraph(SqlAgentState) builder.add_node("discovery", table_discovery) builder.add_node("schema", schema_retrieval) builder.add_node("generator", sql_generation) builder.add_node("executor", sql_execution) builder.add_node("summarizer", answer_formulation) builder.set_entry_point("discovery") builder.add_edge("discovery", "schema") builder.add_edge("schema", "generator") builder.add_edge("generator", "executor") builder.add_conditional_edges( "executor", should_retry, { "retry": "generator", # Идем обратно в генератор для исправления SQL "continue": "summarizer" } ) builder.add_edge("summarizer", END) app = builder.compile() async def get_agent_stream(question: str, settings: dict): db_dialect = settings.get("db_dialect", "") db_host = settings.get("db_host", "") db_port = settings.get("db_port", "") db_name = settings.get("db_name", "") db_user = settings.get("db_user", "") db_password = settings.get("db_password", "") model_host_url = settings.get("model_host_url", "") model_host_type = settings.get("model_host_type", "") model_name = settings.get("model_name", "") model_api_key = settings.get("model_api_key", "") thread_id = settings.get("thread_id", "default") config = { 'configurable': { 'thread_id': thread_id, 'db_dialect': db_dialect, 'db_host': db_host, 'db_host': db_port, 'db_name': db_name, 'db_user': db_user, 'db_password': db_password, 'model_host_url': model_host_url, 'model_host_type': model_host_type, 'model_name': model_name, 'model_api_key': model_api_key, **settings } } async for stream_mode, chunk in app.astream( {"question": question, "error_count": 0}, config=config, stream_mode=["custom", "updates"] ): yield stream_mode, chunk