增加qa知识库相关
This commit is contained in:
parent
2d8623685f
commit
35d93bd787
0
app/chat_qa_query/__init__.py
Normal file
0
app/chat_qa_query/__init__.py
Normal file
0
app/chat_qa_query/api/__init__.py
Normal file
0
app/chat_qa_query/api/__init__.py
Normal file
30
app/chat_qa_query/api/qa_dependencies.py
Normal file
30
app/chat_qa_query/api/qa_dependencies.py
Normal 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)
|
||||
0
app/chat_qa_query/api/router/__init__.py
Normal file
0
app/chat_qa_query/api/router/__init__.py
Normal file
17
app/chat_qa_query/api/router/query_router.py
Normal file
17
app/chat_qa_query/api/router/query_router.py
Normal 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)
|
||||
|
||||
0
app/chat_qa_query/api/schema/__init__.py
Normal file
0
app/chat_qa_query/api/schema/__init__.py
Normal file
6
app/chat_qa_query/api/schema/query_qa_schema.py
Normal file
6
app/chat_qa_query/api/schema/query_qa_schema.py
Normal file
@ -0,0 +1,6 @@
|
||||
from pydantic import BaseModel
|
||||
|
||||
|
||||
class QueryQASchema(BaseModel):
|
||||
external_str:str # 用户聊天内容
|
||||
messages_top15: list[str] # 历史聊天内容
|
||||
10
app/chat_qa_query/query_context.py
Normal file
10
app/chat_qa_query/query_context.py
Normal 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
|
||||
|
||||
58
app/chat_qa_query/query_graph.py
Normal file
58
app/chat_qa_query/query_graph.py
Normal 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())
|
||||
0
app/chat_qa_query/query_nodes/__init__.py
Normal file
0
app/chat_qa_query/query_nodes/__init__.py
Normal file
16
app/chat_qa_query/query_nodes/node_answer_output.py
Normal file
16
app/chat_qa_query/query_nodes/node_answer_output.py
Normal 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("")
|
||||
103
app/chat_qa_query/query_nodes/node_rewrite.py
Normal file
103
app/chat_qa_query/query_nodes/node_rewrite.py
Normal 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} # 兜底
|
||||
29
app/chat_qa_query/query_nodes/node_search_embedding.py
Normal file
29
app/chat_qa_query/query_nodes/node_search_embedding.py
Normal 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 {}
|
||||
8
app/chat_qa_query/query_qa_state.py
Normal file
8
app/chat_qa_query/query_qa_state.py
Normal 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对
|
||||
0
app/chat_qa_query/service/__init__.py
Normal file
0
app/chat_qa_query/service/__init__.py
Normal file
30
app/chat_qa_query/service/query_qa_service.py
Normal file
30
app/chat_qa_query/service/query_qa_service.py
Normal 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 ""
|
||||
@ -98,4 +98,50 @@ class QARepository:
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"QA数据入库失败: {str(e)}", exc_info=True)
|
||||
return []
|
||||
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 {}
|
||||
@ -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
|
||||
Loading…
Reference in New Issue
Block a user