import os from langchain.agents import create_agent from langchain_community.utilities import SQLDatabase from langchain_core.messages import HumanMessage from agent import chat_model from config.db import execute_sql from model import ChatRequest from prompt.sql import SQL_PROMPT from prompt.system import BILL_PROMPT from prompt.table import BILL_TABLE_SCHEMA tables = { "bill": ["bill_book", "bill_category", "bill_pay", "bill_record"] } def query(request: ChatRequest): try: db = SQLDatabase.from_uri(os.getenv("DB_URL"), include_tables=tables[request.type]) table_info = db.get_table_info() sql = chat_model.invoke( SQL_PROMPT.format(schema=table_info, option=BILL_TABLE_SCHEMA, question=request.message) ).content.strip() print("sql: ", sql) result = execute_sql(sql) print("result: ", result) agent = create_agent( model=chat_model, system_prompt=BILL_PROMPT.format(question=request.message, result=result), ) response = agent.invoke({"message": HumanMessage(content=request.message)}) print("response: ", response) return response["messages"][-1].content except Exception as e: print(f"\n[错误]: {str(e)}")