(五)多轮对话上下文,实现交互。

下面给你加上多轮对话上下文 功能,前端只要传同一个 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();

四、关键说明

  1. session_id 怎么生成

    • 前端用 crypto.randomUUID() 生成,或者用户登录后用用户ID+时间戳
    • 同一个聊天窗口/同一场对话用同一个,开新对话就换一个
  2. 历史记录限制

    • 代码里设置了最多保留 10 轮对话,防止token太多导致报错
    • 可以根据需要改 MAX_HISTORY_ROUNDS 的值
  3. 生产环境优化建议

    • 现在历史存在内存里,服务重启就没了;生产环境可以存 Redis
    • 可以加会话过期时间,比如24小时没说话就自动清理
    • 敏感场景可以按用户权限隔离会话
  4. 效果验证

    打开 [http://localhost:8000/docs](http://localhost:8000/docs),先传一个问题,再传「其中XX有多少」,能正确理解上下文就算成功。

相关推荐
wangchunyu1141 小时前
食韵味简 —— 免费开源的中式菜谱 API 服务
python·django·开源
znx9391 小时前
因子分析:量化交易的底层核心与盈利逻辑基石
人工智能·python·机器学习·期魔方
jason.zeng@15022071 小时前
(十)分层架构的多文件工程
python·架构·prompt·交互·ai编程·llama
小小张说故事1 小时前
Python 装饰器到底是个啥?从 @ 符号到手写一个,5 个例子讲透
后端·python
haerapi1 小时前
KingbaseES 全文检索实战之前:先建立中文搜索的可解释基线
开发语言·python·全文检索
专业程序开发源1 小时前
SSM笔记本在线销售系统32649-计算机课程设计、毕业设计
java·spring boot·后端·python·django·php·课程设计
脉动数据行情11 小时前
Python asyncio 异步实现比特币 BTC 实时行情监听 高并发版
开发语言·python·区块链
ClickHouseDB1 小时前
ClickHouse Terraform Provider 正式支持 ClickStack 资源管理
网络·数据库·python
外收内放1 小时前
Python基础语法练习题(55-56)
开发语言·python