feat:初始化工程
This commit is contained in:
44
services/agent_service.py
Normal file
44
services/agent_service.py
Normal file
@@ -0,0 +1,44 @@
|
||||
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)}")
|
||||
Reference in New Issue
Block a user