下面给你加上多轮对话上下文 功能,前端只要传同一个 session_id,就能支持连续追问(比如先问车队清单,再问「其中维修中的有几辆」),代码保持小白友好,直接替换就能用。
一、实现思路
- 前端每次请求带一个 会话ID(session_id),同一场对话用同一个ID
- 后端用字典保存每个会话的历史对话
- 每次查询时,把历史对话一起发给大模型,让它理解上下文
- 回答完自动把本轮对话存入历史,支持无限追问
二、完整升级后的 main.py
python
from dotenv import load_dotenv
import os
from contextlib import asynccontextmanager
from fastapi import FastAPI, HTTPException
from fastapi.middleware.cors import CORSMiddleware
from pydantic import BaseModel
from langchain_community.utilities import SQLDatabase
from langchain_community.agent_toolkits import create_sql_agent
from langchain_openai import ChatOpenAI
# ===================== 1. 加载配置 =====================
load_dotenv()
# 全局变量
sql_agent = None
# 会话存储器:key=会话ID,value=历史对话列表
session_store = {}
# 单会话最大保留历史轮数,防止token溢出
MAX_HISTORY_ROUNDS = 10
# ===================== 2. 服务启动初始化 =====================
@asynccontextmanager
async def lifespan(app: FastAPI):
global sql_agent
print("正在初始化数据库连接...")
db = SQLDatabase.from_uri(
os.getenv("MYSQL_URI"),
include_tables=["bus_vehicle", "bus_fleet", "repair_order"],
sample_rows_in_table_info=2,
view_support=False
)
print("正在初始化大模型...")
llm = ChatOpenAI(
model=os.getenv("MODEL_NAME"),
api_key=os.getenv("OPENAI_API_KEY"),
base_url=os.getenv("OPENAI_BASE_URL"),
temperature=0
)
custom_prompt = """
你是公交机务数据查询专家,只能根据数据库表结构回答问题。
请结合对话上下文理解用户问题,必要时可以追问用户补充信息。
【业务规则】
- "一公司" = company_name = "第一运营分公司"
- "X车队" = fleet_name = "X车队"
- 车辆状态:1=运营中,2=维修中,3=停运
- "上月"指上一个自然月
【严格禁令】
1. 只能生成SELECT查询,绝对禁止DELETE/UPDATE/INSERT/DROP等修改语句
2. 禁止用SELECT *,必须明确字段名
3. 问题不明确就反问用户,不要猜测
4. 回答先给结论,再附数据清单
"""
print("正在创建SQL Agent...")
sql_agent = create_sql_agent(
llm=llm,
db=db,
agent_type="openai-tools",
verbose=False,
extra_system_message=custom_prompt,
max_iterations=5,
handle_parsing_errors=True
)
print("服务启动完成!")
yield
print("服务已关闭")
# ===================== 3. FastAPI 应用 =====================
app = FastAPI(
title="公交机务数据查询接口",
description="支持多轮对话的Text-to-SQL查询接口",
version="1.1.0",
lifespan=lifespan
)
app.add_middleware(
CORSMiddleware,
allow_origins=["*"],
allow_credentials=True,
allow_methods=["*"],
allow_headers=["*"],
)
# ===================== 4. 请求响应格式 =====================
class QueryRequest(BaseModel):
session_id: str # 会话ID,同一场对话保持一致
question: str # 用户问题
class QueryResponse(BaseModel):
code: int = 200
message: str = "success"
data: str = ""
# ===================== 5. 工具函数:管理会话历史 =====================
def get_session_history(session_id: str) -> list:
"""获取指定会话的历史记录,不存在则新建"""
if session_id not in session_store:
session_store[session_id] = []
return session_store[session_id]
def append_history(session_id: str, user_question: str, ai_answer: str):
"""追加一轮对话到历史,超过最大轮数就删掉最早的"""
history = get_session_history(session_id)
history.append(f"用户:{user_question}")
history.append(f"助手:{ai_answer}")
# 保留最近 N 轮
if len(history) > MAX_HISTORY_ROUNDS * 2:
session_store[session_id] = history[-MAX_HISTORY_ROUNDS * 2:]
def build_prompt_with_history(session_id: str, question: str) -> str:
"""把历史对话和当前问题拼成完整输入"""
history = get_session_history(session_id)
if not history:
return question
history_text = "\n".join(history)
return f"""【对话历史】
{history_text}
【当前问题】
{question}
请结合对话历史回答当前问题。"""
# ===================== 6. 接口 =====================
@app.get("/health", summary="健康检查")
async def health_check():
return {"code": 200, "message": "服务运行正常"}
@app.post("/api/query", summary="自然语言查询(支持多轮)", response_model=QueryResponse)
async def query_data(request: QueryRequest):
if not sql_agent:
raise HTTPException(status_code=500, detail="Agent未初始化")
session_id = request.session_id.strip()
question = request.question.strip()
if not session_id:
raise HTTPException(status_code=400, detail="session_id不能为空")
if not question:
raise HTTPException(status_code=400, detail="问题不能为空")
try:
# 1. 拼接历史上下文
full_prompt = build_prompt_with_history(session_id, question)
# 2. 调用 Agent 查询
result = sql_agent.invoke({"input": full_prompt})
answer = result["output"]
# 3. 保存本轮对话到历史
append_history(session_id, question, answer)
return QueryResponse(
code=200,
message="success",
data=answer
)
except Exception as e:
raise HTTPException(status_code=500, detail=f"查询失败:{str(e)}")
# 可选:清空指定会话的历史
@app.post("/api/clear", summary="清空会话历史")
async def clear_session(session_id: str):
if session_id in session_store:
del session_store[session_id]
return {"code": 200, "message": "会话已清空"}
三、前端怎么调用
1. 调用逻辑
- 第一次对话:前端生成一个唯一ID(比如UUID)作为
session_id - 后续追问:一直用同一个
session_id,后端自动记住上下文 - 开启新对话:生成新的
session_id即可
2. 前端调用示例(JS)
javascript
// 生成唯一会话ID(新对话时生成一次就行)
const sessionId = crypto.randomUUID();
async function queryBusData(question) {
const response = await fetch('[http://localhost:8000/api/query](http://localhost:8000/api/query)', {
method: 'POST',
headers: { 'Content-Type': 'application/json' },
body: JSON.stringify({
session_id: sessionId,
question: question
})
});
const result = await response.json();
return result.data;
}
// 测试多轮对话
async function test() {
console.log(await queryBusData("一公司119车队的车辆清单"));
// 直接追问,不用重复说"119车队"
console.log(await queryBusData("其中维修中的有几辆?"));
console.log(await queryBusData("上个月它们一共有多少条维修工单?"));
}
test();
四、关键说明
-
session_id 怎么生成
- 前端用
crypto.randomUUID()生成,或者用户登录后用用户ID+时间戳 - 同一个聊天窗口/同一场对话用同一个,开新对话就换一个
- 前端用
-
历史记录限制
- 代码里设置了最多保留 10 轮对话,防止token太多导致报错
- 可以根据需要改
MAX_HISTORY_ROUNDS的值
-
生产环境优化建议
- 现在历史存在内存里,服务重启就没了;生产环境可以存 Redis
- 可以加会话过期时间,比如24小时没说话就自动清理
- 敏感场景可以按用户权限隔离会话
-
效果验证
打开
[http://localhost:8000/docs](http://localhost:8000/docs),先传一个问题,再传「其中XX有多少」,能正确理解上下文就算成功。