45 lines
1.3 KiB
Python
45 lines
1.3 KiB
Python
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)}")
|