增加qa知识库相关

This commit is contained in:
qinyong@9artedu.com 2026-06-23 13:34:52 +08:00
parent 2d8623685f
commit 35d93bd787
18 changed files with 359 additions and 9 deletions

View File

View File

View File

@ -0,0 +1,30 @@
from typing import Annotated
from fastapi import Depends
from langchain_huggingface import HuggingFaceEndpointEmbeddings
from sqlalchemy.ext.asyncio import AsyncSession
from app.chat_qa_query.service.query_qa_service import QueryQAService
from app.client.mysql_client_manager import db_assistant_mysql_client_manager
from app.repository.archive_messages_repository import ArchiveMessagesRepository
from app.repository.milvus.message_qa_repository import QARepository
"""
普通对象 / 客户端 / 仓库 return不需要关闭不需要释放不需要上下文管理 直接 return 实例
数据库连接 / 会话 / 需要自动释放的资源 yield必须用完自动关闭 / 释放 / 回滚 必须用 yield
"""
async def get_assistant_session()->AsyncSession:
async with db_assistant_mysql_client_manager.session_factory() as db_assistant_session:
yield db_assistant_session
async def get_archive_messages_repository(session: Annotated[AsyncSession, Depends(get_assistant_session)])->ArchiveMessagesRepository:
return ArchiveMessagesRepository(session)
async def get_qa_milvus_repository()->QARepository:
return QARepository()
async def get_qa_query_service(archive_mysql_repository: Annotated[ArchiveMessagesRepository, Depends(get_archive_messages_repository)],
qa_milvus_repository: Annotated[QARepository, Depends(get_qa_milvus_repository)]) -> QueryQAService:
return QueryQAService(archive_mysql_repository=archive_mysql_repository,
qa_milvus_repository=qa_milvus_repository)

View File

View File

@ -0,0 +1,17 @@
from typing import Annotated
from fastapi import APIRouter
from fastapi.params import Depends
from watchfiles import awatch
from app.chat_qa_query.api.qa_dependencies import get_qa_query_service
from app.chat_qa_query.api.schema.query_qa_schema import QueryQASchema
from app.chat_qa_query.service.query_qa_service import QueryQAService
query_qa_router = APIRouter()
@query_qa_router.post("/api/qa/query")
async def query_handler(query: QueryQASchema, query_service: Annotated[QueryQAService, Depends(get_qa_query_service)]):
return await query_service.query(query.external_str, query.messages_top15)

View File

View File

@ -0,0 +1,6 @@
from pydantic import BaseModel
class QueryQASchema(BaseModel):
external_str:str # 用户聊天内容
messages_top15: list[str] # 历史聊天内容

View File

@ -0,0 +1,10 @@
from typing import TypedDict
from app.repository.archive_messages_repository import ArchiveMessagesRepository
from app.repository.milvus.message_qa_repository import QARepository
class QueryQAContext(TypedDict):
archive_mysql_repository: ArchiveMessagesRepository
qa_milvus_repository : QARepository

View File

@ -0,0 +1,58 @@
import asyncio
from datetime import datetime, timedelta
from typing import List, Dict, Any
from langgraph.constants import START, END
from langgraph.graph import StateGraph
from app.chat_qa.nodes.finalize_node import finalize
from app.chat_qa.nodes.generate_qa_node import generate_qa
from app.chat_qa.nodes.quality_check_node import quality_check
from app.chat_qa.nodes.query_message_node import query_message
from app.chat_qa.nodes.route_by_quality_node import route_by_quality
from app.chat_qa_query.query_context import QueryQAContext
from app.chat_qa_query.query_nodes.node_answer_output import answer_output
from app.chat_qa_query.query_nodes.node_rewrite import rewrite_query
from app.chat_qa_query.query_nodes.node_search_embedding import search_embed
from app.chat_qa_query.query_qa_state import QueryQAState
from app.client.embedding_client_manager import embedding_client
from app.client.milvus_client_manager import milvus_client
from app.client.mysql_client_manager import db_assistant_mysql_client_manager
from app.repository.archive_messages_repository import ArchiveMessagesRepository
from app.repository.milvus.message_qa_repository import QARepository
from app.repository.milvus.summary_repository import SummaryRepository
graph_builder = StateGraph(state_schema=QueryQAState, context_schema=QueryQAContext)
graph_builder.add_node("rewrite_query", rewrite_query)
graph_builder.add_node("search_embed", search_embed)
graph_builder.add_node("answer_output", answer_output)
graph_builder.add_edge(START, "rewrite_query")
graph_builder.add_edge("rewrite_query", "search_embed")
graph_builder.add_edge("search_embed", "answer_output")
graph_builder.add_edge("answer_output", END)
qa_graph = graph_builder.compile()
if __name__ == '__main__':
async def test():
db_assistant_mysql_client_manager.init()
embedding_client.init()
milvus_client.init()
async with db_assistant_mysql_client_manager.session_factory() as db_session:
archive_messages_mysql_repository = ArchiveMessagesRepository(db_session)
qa_milvus_repository = QARepository()
state: QueryQAState = QueryQAState(original_query="培训的费用?")
context = QueryQAContext(meta_mysql_repository=archive_messages_mysql_repository,
qa_milvus_repository=qa_milvus_repository
)
async for chunk in qa_graph.astream(input=state, context=context, stream_mode="custom"):
print("=================="+chunk)
milvus_client.close()
await embedding_client.close()
asyncio.run(test())

View File

@ -0,0 +1,16 @@
"""
"""
from langgraph.runtime import Runtime
from app.chat_qa_query.query_qa_state import QueryQAState
async def answer_output(state:QueryQAState, runtime: Runtime):
writer = runtime.stream_writer
qa_pairs = state.get("qa_pairs", "")
if qa_pairs:
writer(qa_pairs.get("answer"))
else:
writer("")

View File

@ -0,0 +1,103 @@
"""
问题改写
"""
from langchain_core.output_parsers import JsonOutputParser
from langchain_core.prompts import PromptTemplate
from app.chat_qa_query.query_qa_state import QueryQAState
from app.llm import llm
from app.core.log import logger
problem_rewriting_prompt = """
你是一名智能客服意图理解助手你的任务是根据用户的历史会话和当前问题提取用户的核心诉求并将其改写为一个语义完整信息自包含的独立问题
## 改写规则(必须严格遵守)
1. **核心诉求提取**
- 从当前问题中提取用户明确表达的诉求可能有一个或多个去重
- 如果当前问题包含代词"这个"""""必须结合历史会话进行指代消解用具体名词替换代词
2. **信息补全**
- 如果当前问题缺少主语或关键上下文如只问"费用多少"必须从历史会话中提取相关主题补全
- 改写后的问题必须是一个**不依赖历史会话也能被独立理解**的完整问题
3. **语义保持**
- 不得改变用户的原始意图
- 不得引入历史会话中不存在的新信息
- 不得遗漏当前问题中的任何关键诉求
4. **输出格式**
- 仅输出JSON不要任何解释性文字
- rewritten_query 必须是一个完整的疑问句"是什么""有哪些""怎么样""如何"等结尾或以问号结尾
## 改写示例
### 示例1信息补全
历史会话
- 用户我想了解游戏特效培训
- 助手好的请问您想了解哪方面
当前问题培训的费用
改写结果
{{
"rewritten_query": "游戏特效培训的费用是多少?"
}}
### 示例2指代消解
历史会话
- 用户你们有Python课程吗
- 助手有的我们有Python基础班和进阶班
当前问题这个课程的时长是多久
改写结果
{{
"rewritten_query": "Python课程的时长是多久"
}}
### 示例3多诉求提取
历史会话
- 用户我想报名数据分析培训
当前问题流程费用和就业支持
改写结果
{{
"rewritten_query": "数据分析培训的报名流程、费用以及就业支持分别是什么?"
}}
### 示例4无需改写
历史会话
当前问题Java培训的课程内容有哪些
改写结果
{{
"rewritten_query": "Java培训的课程内容有哪些"
}}
---
## 现在请处理以下输入
历史会话
{history_text}
当前问题{query}
请直接输出JSON
"""
async def rewrite_query(state:QueryQAState):
"""
问题改写
:param query: 用户问题
:param history_text: 历史会话
:return: 改写后的问题
"""
history_text = state.get("history", "")
query = state.get("original_query", "")
prompt = PromptTemplate(template=problem_rewriting_prompt, input_variables=["history_text", "query"])
output = JsonOutputParser()
chain = prompt | llm | output
try:
result = await chain.ainvoke({"history_text": history_text, "query": query})
rewritten_query = result.get("rewritten_query", "")
return {"rewritten_query": rewritten_query}
except Exception as e:
logger.error(f"QA质量检查prompt模板错误{e}")
return {"rewritten_query": query} # 兜底

View File

@ -0,0 +1,29 @@
"""
向量化
"""
from app.chat_qa_query.query_qa_state import QueryQAState
from app.client.embedding_client_manager import embedding_client
from app.repository.milvus.message_qa_repository import QARepository
from app.core.log import logger
async def search_embed(state:QueryQAState) :
logger.info("search_embed")
rewritten_query = state.get("rewritten_query", "")
original_query = state.get("original_query", "")
batch_embeddings = await embedding_client.aembed_query(rewritten_query)
qa_repository = QARepository()
result = await qa_repository.search_qa(query_text=batch_embeddings, top_k=3)
logger.info(f"原始问题:{original_query},改写后的问题:{rewritten_query},匹配结果:{result}")
if result:
question = result[0]["question"]
answer = result[0]["answer"]
qa_pairs={
"question": question,
"answer": answer,
}
return {"qa_pairs": qa_pairs}
return {}

View File

@ -0,0 +1,8 @@
from typing import TypedDict, List, Dict, Optional, Any
class QueryQAState(TypedDict):
original_query: str # 原始问题
rewritten_query: str # 改写后的问题
history: list # 历史对话记录
qa_pairs: Dict[str, Any] # 匹配到的QA对

View File

View File

@ -0,0 +1,30 @@
import json
from typing import List
from langchain_huggingface import HuggingFaceEndpointEmbeddings
from app.chat_qa_query.query_context import QueryQAContext
from app.chat_qa_query.query_graph import qa_graph
from app.chat_qa_query.query_qa_state import QueryQAState
from app.core.log import logger
from app.repository.archive_messages_repository import ArchiveMessagesRepository
from app.repository.milvus.message_qa_repository import QARepository
class QueryQAService:
def __init__(self, archive_mysql_repository: ArchiveMessagesRepository,
qa_milvus_repository : QARepository):
self.archive_mysql_repository = archive_mysql_repository
self.qa_milvus_repository = qa_milvus_repository
async def query(self, query:str, history:List[str]=[])-> str:
state: QueryQAState = QueryQAState(original_query=query, history=history)
context = QueryQAContext(archive_mysql_repository=self.archive_mysql_repository,qa_milvus_repository=self.qa_milvus_repository)
try:
res= await qa_graph.ainvoke(input=state, context=context, stream_mode="custom")
return res
except Exception as e:
logger.info(f"QA服务错误{e}")
return ""

View File

@ -99,3 +99,49 @@ class QARepository:
except Exception as e:
logger.error(f"QA数据入库失败: {str(e)}", exc_info=True)
return []
async def search_qa(self, query_text: List[float], top_k: int = 5) -> List[Dict[str, Any]]:
"""
根据查询文本搜索QA数据
Args:
query_text: 查询文本
top_k: 返回结果数量
Returns:
搜索结果列表格式 [{"question": "...", "answer": "...", "score": 0.0}, ...]
"""
res = milvus_client.client.search(
collection_name=self.collection_name,
anns_field="question_dense_vector",
data=[query_text],
limit=top_k,
search_params={"metric_type": "COSINE"},
output_fields=["question", "answer", "created_at"]
)
score_threshold = 0.6
filter_results = []
# res 是二维结构:[ [单条结果列表] ]
for hits in res:
for hit in hits:
if hit.get("distance",0) >= score_threshold:
# 按需组装数据
item = {
"question": hit.entity.get("question"),
"answer": hit.entity.get("answer"),
"created_at": hit.entity.get("created_at"),
"score": hit.get("distance",0)
}
filter_results.append(item)
# filter_results 就是最终过滤后的结果
print("filter_results:", filter_results)
return filter_results
# 判空再取值
# if filter_results:
# return filter_results[0]
# else:
# # 无匹配结果,返回空字典 / None根据业务选其一
# return {}

View File

@ -4,19 +4,16 @@ from typing import Optional
from apscheduler.schedulers.asyncio import AsyncIOScheduler
from fastapi import FastAPI
from app.client.embedding_client_manager import EmbeddingClientManager, embedding_client
from app.client.milvus_client_manager import milvus_client
from app.client.mysql_client_manager import db_assistant_mysql_client_manager
from app.repository.archive_messages_repository import ArchiveMessagesRepository
from app.repository.milvus.summary_repository import SummaryRepository
from app.summary.service import SummaryService, build
from app.chat_qa_query.api.router.query_router import query_qa_router
from app.core.liftspan import lifespan
from app.summary.service import build
from app.core.log import logger
app = FastAPI(title="消息摘要服务", version="1.0")
app = FastAPI(lifespan=lifespan)
scheduler = AsyncIOScheduler(timezone='Asia/Shanghai')
app.include_router(query_qa_router)
async def generate_daily_summary(date: Optional[str] = None):
logger.info("16:30分定时任务测试")
if not date:
yesterday = datetime.now() - timedelta(days=1)
#date = yesterday.strftime("%Y-%-m-%-d") # 格式: 2026-6-15